diff --git a/.changeset/calm-lids-align.md b/.changeset/calm-lids-align.md new file mode 100644 index 00000000..820065a1 --- /dev/null +++ b/.changeset/calm-lids-align.md @@ -0,0 +1,9 @@ +--- +'@zapo-js/store-mongo': patch +'@zapo-js/store-mysql': patch +'@zapo-js/store-postgres': patch +'@zapo-js/store-redis': patch +'@zapo-js/store-sqlite': patch +--- + +Persist PN/LID Signal-address mappings alongside configured sessions so both aliases share canonical state. diff --git a/README.md b/README.md index 682d0f19..90126551 100644 --- a/README.md +++ b/README.md @@ -104,6 +104,10 @@ client.on('message', async (event) => { await client.connect() ``` +The Signal `session` provider also owns the internal PN/LID mapping used to +keep ratchets canonical when WhatsApp alternates between phone-number and LID +addressing. Official persistent backends store that mapping automatically. + That's the minimum to pair, listen for messages, and reply. For everything else - sending media, reactions, polls, groups, newsletters, app-state mutations, business profile, events catalog, store providers, the typed diff --git a/packages/store-mongo/src/__tests__/integration.test.ts b/packages/store-mongo/src/__tests__/integration.test.ts index 7481bfc6..725106fb 100644 --- a/packages/store-mongo/src/__tests__/integration.test.ts +++ b/packages/store-mongo/src/__tests__/integration.test.ts @@ -1048,6 +1048,57 @@ describe('store-mongo integration', { timeout: 60_000 }, () => { await senderKey.clear() }) + it('signal: PN/LID mappings are replaceable and session-scoped', async (t) => { + if (!store) return t.skip('ZAPO_TEST_MONGO_* not set') + + const mappingA = store.stores.lidPnMapping(nextSessionId('lid-pn-a')) + const mappingB = store.stores.lidPnMapping(nextSessionId('lid-pn-b')) + await Promise.all([mappingA.clear(), mappingB.clear()]) + + assert.equal(await mappingA.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '111222') + assert.equal(await mappingA.getLidUser('5511999999999'), '111222') + assert.equal(await mappingA.getPnUser('111222'), '5511999999999') + assert.equal(await mappingB.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), '333444') + assert.equal(await mappingA.getPnUser('111222'), null) + await mappingA.setLidUser('5511888888888', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), null) + assert.equal(await mappingA.getPnUser('333444'), '5511888888888') + await mappingA.clear() + assert.equal(await mappingA.getLidUser('5511888888888'), null) + assert.equal(await mappingA.getPnUser('333444'), null) + }) + + it('signal: concurrent PN/LID replacements preserve one owner', async (t) => { + if (!store) return t.skip('ZAPO_TEST_MONGO_* not set') + + const sessionId = nextSessionId('lid-pn-concurrent') + const mappingA = store.stores.lidPnMapping(sessionId) + const mappingB = store.stores.lidPnMapping(sessionId) + await mappingA.clear() + + for (let index = 0; index < 5; index += 1) { + const lidUser = `55566${index}` + const pnUsers = [`55117777777${index}`, `55116666666${index}`] + await Promise.all([ + mappingA.setLidUser(pnUsers[0], lidUser), + mappingB.setLidUser(pnUsers[1], lidUser) + ]) + + const owner = await mappingA.getPnUser(lidUser) + assert.ok(owner === pnUsers[0] || owner === pnUsers[1]) + assert.equal(await mappingA.getLidUser(owner), lidUser) + assert.equal( + await mappingA.getLidUser(owner === pnUsers[0] ? pnUsers[1] : pnUsers[0]), + null + ) + } + + await mappingA.clear() + }) + it('signal: session lifecycle and batch queries', async (t) => { if (!store) return t.skip('ZAPO_TEST_MONGO_* not set') diff --git a/packages/store-mongo/src/createMongoStore.ts b/packages/store-mongo/src/createMongoStore.ts index 21417b0f..46acb425 100644 --- a/packages/store-mongo/src/createMongoStore.ts +++ b/packages/store-mongo/src/createMongoStore.ts @@ -7,6 +7,7 @@ import { WaContactMongoStore } from './contact.store' import { WaDeviceListMongoStore } from './device-list.store' import { WaGroupMetadataMongoStore } from './group-metadata.store' import { WaIdentityMongoStore } from './identity.store' +import { WaLidPnMappingMongoStore } from './lid-pn-mapping.store' import { WaMessageSecretMongoStore } from './message-secret.store' import { WaMessageMongoStore } from './message.store' import { WaPreKeyMongoStore } from './pre-key.store' @@ -69,6 +70,7 @@ export interface WaMongoStoreResult { readonly preKey: (sessionId: string) => WaPreKeyMongoStore readonly session: (sessionId: string) => WaSessionMongoStore readonly identity: (sessionId: string) => WaIdentityMongoStore + readonly lidPnMapping: (sessionId: string) => WaLidPnMappingMongoStore readonly signal: (sessionId: string) => WaSignalMongoStore readonly senderKey: (sessionId: string) => WaSenderKeyMongoStore readonly appState: (sessionId: string) => WaAppStateMongoStore @@ -91,7 +93,7 @@ function isDb(value: WaMongoStoreConfig['db']): value is Db { } /** - * Builds a MongoDB-backed {@link WaStoreBackend} bundle. All 11 persistent + * Builds a MongoDB-backed {@link WaStoreBackend} bundle. All 12 persistent * domains + 4 cache domains live in a single database (split into * collections by `collectionPrefix`). * @@ -165,6 +167,8 @@ export function createMongoStore(config: WaMongoStoreConfig): WaMongoStoreResult preKey: (sessionId) => new WaPreKeyMongoStore(opts(sessionId, 'preKey')), session: (sessionId) => new WaSessionMongoStore(opts(sessionId, 'session')), identity: (sessionId) => new WaIdentityMongoStore(opts(sessionId, 'identity')), + lidPnMapping: (sessionId) => + new WaLidPnMappingMongoStore(opts(sessionId, 'lidPnMapping')), signal: (sessionId) => new WaSignalMongoStore(opts(sessionId, 'signal')), senderKey: (sessionId) => new WaSenderKeyMongoStore(opts(sessionId, 'senderKey')), appState: (sessionId) => new WaAppStateMongoStore(opts(sessionId, 'appState')), diff --git a/packages/store-mongo/src/index.ts b/packages/store-mongo/src/index.ts index 05290e61..ff84a8c6 100644 --- a/packages/store-mongo/src/index.ts +++ b/packages/store-mongo/src/index.ts @@ -4,6 +4,7 @@ export { WaAuthMongoStore } from './auth.store' export { WaPreKeyMongoStore } from './pre-key.store' export { WaSessionMongoStore } from './session.store' export { WaIdentityMongoStore } from './identity.store' +export { WaLidPnMappingMongoStore } from './lid-pn-mapping.store' export { WaSignalMongoStore } from './signal.store' export { WaSenderKeyMongoStore } from './sender-key.store' export { WaAppStateMongoStore } from './appstate.store' diff --git a/packages/store-mongo/src/lid-pn-mapping.store.ts b/packages/store-mongo/src/lid-pn-mapping.store.ts new file mode 100644 index 00000000..9289c358 --- /dev/null +++ b/packages/store-mongo/src/lid-pn-mapping.store.ts @@ -0,0 +1,102 @@ +import type { WaLidPnMappingStore } from 'zapo-js/store' + +import { BaseMongoStore } from './BaseMongoStore' +import type { WaMongoStorageOptions } from './types' + +const COLLECTION = 'signal_lid_pn_mappings' +const REPLACE_MAX_ATTEMPTS = 3 + +interface LidPnMappingDoc { + _id: { session_id: string; pn_user: string } + lid_user: string +} + +function isDuplicateKeyError(error: unknown): boolean { + return ( + typeof error === 'object' && + error !== null && + 'code' in error && + (error as { readonly code?: unknown }).code === 11_000 + ) +} + +/** MongoDB-backed PN/LID mapping store scoped by Zapo session id. */ +export class WaLidPnMappingMongoStore extends BaseMongoStore implements WaLidPnMappingStore { + private writeTail: Promise = Promise.resolve() + + public constructor(options: WaMongoStorageOptions) { + super(options) + } + + protected override async createIndexes(): Promise { + await this.col(COLLECTION).createIndex( + { '_id.session_id': 1, lid_user: 1 }, + { unique: true } + ) + } + + public async getLidUser(pnUser: string): Promise { + await this.ensureIndexes() + const doc = await this.col(COLLECTION).findOne({ + _id: { session_id: this.sessionId, pn_user: pnUser } + }) + return doc?.lid_user ?? null + } + + public async getPnUser(lidUser: string): Promise { + await this.ensureIndexes() + const doc = await this.col(COLLECTION).findOne({ + '_id.session_id': this.sessionId, + lid_user: lidUser + }) + return doc?._id.pn_user ?? null + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + await this.runWriteSerialized(async () => { + for (let attempt = 1; attempt <= REPLACE_MAX_ATTEMPTS; attempt += 1) { + try { + await this.withSession(async (session) => { + const collection = this.col(COLLECTION) + await collection.deleteMany( + { + '_id.session_id': this.sessionId, + '_id.pn_user': { $ne: pnUser }, + lid_user: lidUser + }, + { session } + ) + await collection.updateOne( + { _id: { session_id: this.sessionId, pn_user: pnUser } }, + { $set: { lid_user: lidUser } }, + { upsert: true, session } + ) + }) + return + } catch (error) { + if (attempt === REPLACE_MAX_ATTEMPTS || !isDuplicateKeyError(error)) { + throw error + } + } + } + }) + } + + public async clear(): Promise { + await this.runWriteSerialized(async () => { + await this.ensureIndexes() + await this.col(COLLECTION).deleteMany({ + '_id.session_id': this.sessionId + }) + }) + } + + private runWriteSerialized(task: () => Promise): Promise { + const result = this.writeTail.then(task) + this.writeTail = result.then( + () => undefined, + () => undefined + ) + return result + } +} diff --git a/packages/store-mysql/src/__tests__/integration.test.ts b/packages/store-mysql/src/__tests__/integration.test.ts index 0a87ea38..bed97fbd 100644 --- a/packages/store-mysql/src/__tests__/integration.test.ts +++ b/packages/store-mysql/src/__tests__/integration.test.ts @@ -89,6 +89,7 @@ describe('store-mysql integration', { timeout: 60_000 }, () => { await ensureMysqlMigrations(pool, [ 'auth', 'signal', + 'lidPnMapping', 'senderKey', 'appState', 'retry', @@ -1080,6 +1081,29 @@ describe('store-mysql integration', { timeout: 60_000 }, () => { await senderKey.clear() }) + it('signal: PN/LID mappings are replaceable and session-scoped', async (t) => { + if (!store) return t.skip('ZAPO_TEST_MYSQL_* not set') + + const mappingA = store.stores.lidPnMapping(nextSessionId('lid-pn-a')) + const mappingB = store.stores.lidPnMapping(nextSessionId('lid-pn-b')) + await Promise.all([mappingA.clear(), mappingB.clear()]) + + assert.equal(await mappingA.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '111222') + assert.equal(await mappingA.getLidUser('5511999999999'), '111222') + assert.equal(await mappingA.getPnUser('111222'), '5511999999999') + assert.equal(await mappingB.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), '333444') + assert.equal(await mappingA.getPnUser('111222'), null) + await mappingA.setLidUser('5511888888888', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), null) + assert.equal(await mappingA.getPnUser('333444'), '5511888888888') + await mappingA.clear() + assert.equal(await mappingA.getLidUser('5511888888888'), null) + assert.equal(await mappingA.getPnUser('333444'), null) + }) + it('signal: session lifecycle and batch queries', async (t) => { if (!store) return t.skip('ZAPO_TEST_MYSQL_* not set') diff --git a/packages/store-mysql/src/connection.ts b/packages/store-mysql/src/connection.ts index 166a5a18..074a4014 100644 --- a/packages/store-mysql/src/connection.ts +++ b/packages/store-mysql/src/connection.ts @@ -362,6 +362,19 @@ const MIGRATIONS: readonly Migration[] = [ ALTER TABLE \`__PREFIX__retry_inbound_counters\` DROP COLUMN updated_at_ms ` + }, + { + name: '0018_signal_lid_pn_mapping', + domain: 'lidPnMapping', + sql: ` + CREATE TABLE IF NOT EXISTS \`__PREFIX__signal_lid_pn_mapping\` ( + session_id VARCHAR(128) NOT NULL, + pn_user VARCHAR(128) NOT NULL, + lid_user VARCHAR(128) NOT NULL, + PRIMARY KEY (session_id, pn_user), + UNIQUE KEY uq_signal_lid_pn_mapping_lid (session_id, lid_user) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 + ` } ] diff --git a/packages/store-mysql/src/createMysqlStore.ts b/packages/store-mysql/src/createMysqlStore.ts index a0411a7e..2704cae8 100644 --- a/packages/store-mysql/src/createMysqlStore.ts +++ b/packages/store-mysql/src/createMysqlStore.ts @@ -9,6 +9,7 @@ import { WaContactMysqlStore } from './contact.store' import { WaDeviceListMysqlStore } from './device-list.store' import { WaGroupMetadataMysqlStore } from './group-metadata.store' import { WaIdentityMysqlStore } from './identity.store' +import { WaLidPnMappingMysqlStore } from './lid-pn-mapping.store' import { WaMessageSecretMysqlStore } from './message-secret.store' import { WaMessageMysqlStore } from './message.store' import { WaPreKeyMysqlStore } from './pre-key.store' @@ -90,6 +91,7 @@ export interface WaMysqlStoreResult { readonly preKey: (sessionId: string) => WaPreKeyMysqlStore readonly session: (sessionId: string) => WaSessionMysqlStore readonly identity: (sessionId: string) => WaIdentityMysqlStore + readonly lidPnMapping: (sessionId: string) => WaLidPnMappingMysqlStore readonly signal: (sessionId: string) => WaSignalMysqlStore readonly senderKey: (sessionId: string) => WaSenderKeyMysqlStore readonly appState: (sessionId: string) => WaAppStateMysqlStore @@ -172,6 +174,8 @@ export function createMysqlStore(config: WaMysqlStoreConfig): WaMysqlStoreResult preKey: (sessionId) => new WaPreKeyMysqlStore(opts(sessionId, 'preKey')), session: (sessionId) => new WaSessionMysqlStore(opts(sessionId, 'session')), identity: (sessionId) => new WaIdentityMysqlStore(opts(sessionId, 'identity')), + lidPnMapping: (sessionId) => + new WaLidPnMappingMysqlStore(opts(sessionId, 'lidPnMapping')), signal: (sessionId) => new WaSignalMysqlStore(opts(sessionId, 'signal')), senderKey: (sessionId) => new WaSenderKeyMysqlStore(opts(sessionId, 'senderKey')), appState: (sessionId) => new WaAppStateMysqlStore(opts(sessionId, 'appState')), diff --git a/packages/store-mysql/src/index.ts b/packages/store-mysql/src/index.ts index acfad21a..d8a97c6f 100644 --- a/packages/store-mysql/src/index.ts +++ b/packages/store-mysql/src/index.ts @@ -9,6 +9,7 @@ export { WaAuthMysqlStore } from './auth.store' export { WaPreKeyMysqlStore } from './pre-key.store' export { WaSessionMysqlStore } from './session.store' export { WaIdentityMysqlStore } from './identity.store' +export { WaLidPnMappingMysqlStore } from './lid-pn-mapping.store' export { WaSignalMysqlStore } from './signal.store' export { WaSenderKeyMysqlStore } from './sender-key.store' export { WaAppStateMysqlStore } from './appstate.store' diff --git a/packages/store-mysql/src/lid-pn-mapping.store.ts b/packages/store-mysql/src/lid-pn-mapping.store.ts new file mode 100644 index 00000000..465334e3 --- /dev/null +++ b/packages/store-mysql/src/lid-pn-mapping.store.ts @@ -0,0 +1,61 @@ +import type { WaLidPnMappingStore } from 'zapo-js/store' + +import { BaseMysqlStore } from './BaseMysqlStore' +import { queryFirst } from './helpers' +import type { WaMysqlStorageOptions } from './types' + +/** MySQL-backed PN/LID mapping store scoped by Zapo session id. */ +export class WaLidPnMappingMysqlStore extends BaseMysqlStore implements WaLidPnMappingStore { + public constructor(options: WaMysqlStorageOptions) { + super(options, ['lidPnMapping']) + } + + public async getLidUser(pnUser: string): Promise { + await this.ensureReady() + const row = queryFirst( + await this.pool.execute( + `SELECT lid_user + FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = ? AND pn_user = ?`, + [this.sessionId, pnUser] + ) + ) + return row ? String(row.lid_user) : null + } + + public async getPnUser(lidUser: string): Promise { + await this.ensureReady() + const row = queryFirst( + await this.pool.execute( + `SELECT pn_user + FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = ? AND lid_user = ?`, + [this.sessionId, lidUser] + ) + ) + return row ? String(row.pn_user) : null + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + await this.withTransaction(async (connection) => { + await connection.execute( + `DELETE FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = ? AND (pn_user = ? OR lid_user = ?)`, + [this.sessionId, pnUser, lidUser] + ) + await connection.execute( + `INSERT INTO ${this.t('signal_lid_pn_mapping')} (session_id, pn_user, lid_user) + VALUES (?, ?, ?)`, + [this.sessionId, pnUser, lidUser] + ) + }) + } + + public async clear(): Promise { + await this.ensureReady() + await this.pool.execute( + `DELETE FROM ${this.t('signal_lid_pn_mapping')} WHERE session_id = ?`, + [this.sessionId] + ) + } +} diff --git a/packages/store-mysql/src/types.ts b/packages/store-mysql/src/types.ts index 6fb93333..d3a56a67 100644 --- a/packages/store-mysql/src/types.ts +++ b/packages/store-mysql/src/types.ts @@ -6,6 +6,7 @@ export type MysqlParam = string | number | bigint | Uint8Array | boolean | null export type WaMysqlMigrationDomain = | 'auth' | 'signal' + | 'lidPnMapping' | 'senderKey' | 'appState' | 'retry' diff --git a/packages/store-postgres/src/__tests__/integration.test.ts b/packages/store-postgres/src/__tests__/integration.test.ts index 1dbe75e4..296f23fb 100644 --- a/packages/store-postgres/src/__tests__/integration.test.ts +++ b/packages/store-postgres/src/__tests__/integration.test.ts @@ -89,6 +89,7 @@ describe('store-postgres integration', { timeout: 60_000 }, () => { await ensurePgMigrations(pool, [ 'auth', 'signal', + 'lidPnMapping', 'senderKey', 'appState', 'retry', @@ -1066,6 +1067,29 @@ describe('store-postgres integration', { timeout: 60_000 }, () => { await senderKey.clear() }) + it('signal: PN/LID mappings are replaceable and session-scoped', async (t) => { + if (!store) return t.skip('ZAPO_TEST_PG_* not set') + + const mappingA = store.stores.lidPnMapping(nextSessionId('lid-pn-a')) + const mappingB = store.stores.lidPnMapping(nextSessionId('lid-pn-b')) + await Promise.all([mappingA.clear(), mappingB.clear()]) + + assert.equal(await mappingA.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '111222') + assert.equal(await mappingA.getLidUser('5511999999999'), '111222') + assert.equal(await mappingA.getPnUser('111222'), '5511999999999') + assert.equal(await mappingB.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), '333444') + assert.equal(await mappingA.getPnUser('111222'), null) + await mappingA.setLidUser('5511888888888', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), null) + assert.equal(await mappingA.getPnUser('333444'), '5511888888888') + await mappingA.clear() + assert.equal(await mappingA.getLidUser('5511888888888'), null) + assert.equal(await mappingA.getPnUser('333444'), null) + }) + it('signal: session lifecycle and batch queries', async (t) => { if (!store) return t.skip('ZAPO_TEST_PG_* not set') diff --git a/packages/store-postgres/src/connection.ts b/packages/store-postgres/src/connection.ts index 2f57ea28..cd610349 100644 --- a/packages/store-postgres/src/connection.ts +++ b/packages/store-postgres/src/connection.ts @@ -104,6 +104,19 @@ const MIGRATIONS: readonly Migration[] = [ ) ` }, + { + name: '0017_signal_lid_pn_mapping', + domain: 'lidPnMapping', + sql: ` + CREATE TABLE IF NOT EXISTS "__PREFIX__signal_lid_pn_mapping" ( + session_id TEXT NOT NULL, + pn_user TEXT NOT NULL, + lid_user TEXT NOT NULL, + PRIMARY KEY (session_id, pn_user), + UNIQUE (session_id, lid_user) + ) + ` + }, { name: '0003_sender_key_schema', domain: 'senderKey', diff --git a/packages/store-postgres/src/createPostgresStore.ts b/packages/store-postgres/src/createPostgresStore.ts index 71c4b145..4de6c871 100644 --- a/packages/store-postgres/src/createPostgresStore.ts +++ b/packages/store-postgres/src/createPostgresStore.ts @@ -9,6 +9,7 @@ import { WaContactPgStore } from './contact.store' import { WaDeviceListPgStore } from './device-list.store' import { WaGroupMetadataPgStore } from './group-metadata.store' import { WaIdentityPgStore } from './identity.store' +import { WaLidPnMappingPgStore } from './lid-pn-mapping.store' import { WaMessageSecretPgStore } from './message-secret.store' import { WaMessagePgStore } from './message.store' import { WaPreKeyPgStore } from './pre-key.store' @@ -85,6 +86,7 @@ export interface WaPgStoreResult { readonly preKey: (sessionId: string) => WaPreKeyPgStore readonly session: (sessionId: string) => WaSessionPgStore readonly identity: (sessionId: string) => WaIdentityPgStore + readonly lidPnMapping: (sessionId: string) => WaLidPnMappingPgStore readonly signal: (sessionId: string) => WaSignalPgStore readonly senderKey: (sessionId: string) => WaSenderKeyPgStore readonly appState: (sessionId: string) => WaAppStatePgStore @@ -169,6 +171,7 @@ export function createPostgresStore(config: WaPgStoreConfig): WaPgStoreResult { preKey: (sessionId) => new WaPreKeyPgStore(opts(sessionId, 'preKey')), session: (sessionId) => new WaSessionPgStore(opts(sessionId, 'session')), identity: (sessionId) => new WaIdentityPgStore(opts(sessionId, 'identity')), + lidPnMapping: (sessionId) => new WaLidPnMappingPgStore(opts(sessionId, 'lidPnMapping')), signal: (sessionId) => new WaSignalPgStore(opts(sessionId, 'signal')), senderKey: (sessionId) => new WaSenderKeyPgStore(opts(sessionId, 'senderKey')), appState: (sessionId) => new WaAppStatePgStore(opts(sessionId, 'appState')), diff --git a/packages/store-postgres/src/index.ts b/packages/store-postgres/src/index.ts index 6b116f07..d2298480 100644 --- a/packages/store-postgres/src/index.ts +++ b/packages/store-postgres/src/index.ts @@ -10,6 +10,7 @@ export { WaAuthPgStore } from './auth.store' export { WaPreKeyPgStore } from './pre-key.store' export { WaSessionPgStore } from './session.store' export { WaIdentityPgStore } from './identity.store' +export { WaLidPnMappingPgStore } from './lid-pn-mapping.store' export { WaSignalPgStore } from './signal.store' export { WaSenderKeyPgStore } from './sender-key.store' export { WaAppStatePgStore } from './appstate.store' diff --git a/packages/store-postgres/src/lid-pn-mapping.store.ts b/packages/store-postgres/src/lid-pn-mapping.store.ts new file mode 100644 index 00000000..eee28bbe --- /dev/null +++ b/packages/store-postgres/src/lid-pn-mapping.store.ts @@ -0,0 +1,64 @@ +import type { WaLidPnMappingStore } from 'zapo-js/store' + +import { BasePgStore } from './BasePgStore' +import { queryFirst } from './helpers' +import type { WaPgStorageOptions } from './types' + +/** PostgreSQL-backed PN/LID mapping store scoped by Zapo session id. */ +export class WaLidPnMappingPgStore extends BasePgStore implements WaLidPnMappingStore { + public constructor(options: WaPgStorageOptions) { + super(options, ['lidPnMapping']) + } + + public async getLidUser(pnUser: string): Promise { + await this.ensureReady() + const row = queryFirst( + await this.pool.query({ + name: this.stmtName('lid_pn_mapping_get'), + text: `SELECT lid_user + FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = $1 AND pn_user = $2`, + values: [this.sessionId, pnUser] + }) + ) + return row ? String(row.lid_user) : null + } + + public async getPnUser(lidUser: string): Promise { + await this.ensureReady() + const row = queryFirst( + await this.pool.query({ + name: this.stmtName('lid_pn_mapping_get_pn'), + text: `SELECT pn_user + FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = $1 AND lid_user = $2`, + values: [this.sessionId, lidUser] + }) + ) + return row ? String(row.pn_user) : null + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + await this.ensureReady() + await this.pool.query({ + name: this.stmtName('lid_pn_mapping_set'), + text: `WITH removed AS ( + DELETE FROM ${this.t('signal_lid_pn_mapping')} + WHERE session_id = $1 AND (pn_user = $2 OR lid_user = $3) + RETURNING 1 + ) + INSERT INTO ${this.t('signal_lid_pn_mapping')} (session_id, pn_user, lid_user) + SELECT $1, $2, $3 FROM (SELECT count(*) FROM removed) AS deleted`, + values: [this.sessionId, pnUser, lidUser] + }) + } + + public async clear(): Promise { + await this.ensureReady() + await this.pool.query({ + name: this.stmtName('lid_pn_mapping_clear'), + text: `DELETE FROM ${this.t('signal_lid_pn_mapping')} WHERE session_id = $1`, + values: [this.sessionId] + }) + } +} diff --git a/packages/store-postgres/src/types.ts b/packages/store-postgres/src/types.ts index c4ab52e4..2ecdbb93 100644 --- a/packages/store-postgres/src/types.ts +++ b/packages/store-postgres/src/types.ts @@ -6,6 +6,7 @@ export type PgParam = string | number | bigint | Uint8Array | boolean | null export type WaPgMigrationDomain = | 'auth' | 'signal' + | 'lidPnMapping' | 'senderKey' | 'appState' | 'retry' diff --git a/packages/store-redis/README.md b/packages/store-redis/README.md index 4433f8e9..94fb3377 100644 --- a/packages/store-redis/README.md +++ b/packages/store-redis/README.md @@ -2,7 +2,7 @@ Redis-backed persistent store for [`zapo-js`](https://www.npmjs.com/package/zapo-js). Best fit when you already run Redis for caching and want stateless app instances that can share WhatsApp session state through a network database. -Built on [`ioredis`](https://github.com/redis/ioredis). All 11 persistent domains and 4 cache domains live under a single key namespace (`keyPrefix` controls it). Cache TTLs are enforced natively by Redis - no background cleanup job needed. +Built on [`ioredis`](https://github.com/redis/ioredis). All 12 persistent domains and 4 cache domains live under a single key namespace (`keyPrefix` controls it). Cache TTLs are enforced natively by Redis - no background cleanup job needed. ## Install @@ -85,6 +85,9 @@ createRedisStore({ are reclaimed. - **`auth` has no TTL knob** on purpose: expiring login credentials would log the device out. It always persists. +- **`lidPnMapping` also has no TTL knob**: expiring the address index while a + Signal session remains usable could make the same device resolve to a second + ratchet. The client removes it when session state is cleared. > Set a crypto/session TTL short relative to how often a session connects at > your own risk - an idle window longer than the TTL evicts the Signal / app-state diff --git a/packages/store-redis/src/__tests__/integration.test.ts b/packages/store-redis/src/__tests__/integration.test.ts index 070b02a4..ccca1a96 100644 --- a/packages/store-redis/src/__tests__/integration.test.ts +++ b/packages/store-redis/src/__tests__/integration.test.ts @@ -1046,6 +1046,29 @@ describe('store-redis integration', { timeout: 60_000 }, () => { await senderKey.clear() }) + it('signal: PN/LID mappings are replaceable and session-scoped', async (t) => { + if (!store) return t.skip('ZAPO_TEST_REDIS_* not set') + + const mappingA = store.stores.lidPnMapping(nextSessionId('lid-pn-a')) + const mappingB = store.stores.lidPnMapping(nextSessionId('lid-pn-b')) + await Promise.all([mappingA.clear(), mappingB.clear()]) + + assert.equal(await mappingA.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '111222') + assert.equal(await mappingA.getLidUser('5511999999999'), '111222') + assert.equal(await mappingA.getPnUser('111222'), '5511999999999') + assert.equal(await mappingB.getLidUser('5511999999999'), null) + await mappingA.setLidUser('5511999999999', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), '333444') + assert.equal(await mappingA.getPnUser('111222'), null) + await mappingA.setLidUser('5511888888888', '333444') + assert.equal(await mappingA.getLidUser('5511999999999'), null) + assert.equal(await mappingA.getPnUser('333444'), '5511888888888') + await mappingA.clear() + assert.equal(await mappingA.getLidUser('5511888888888'), null) + assert.equal(await mappingA.getPnUser('333444'), null) + }) + it('signal: session lifecycle and batch queries', async (t) => { if (!store) return t.skip('ZAPO_TEST_REDIS_* not set') @@ -1395,6 +1418,7 @@ describe('store-redis storeTtlMs', { timeout: 60_000 }, () => { redis, storeTtlMs: { messagesMs: TTL, + sessionMs: TTL, signalMs: TTL, appStateMs: TTL } @@ -1492,6 +1516,20 @@ describe('store-redis storeTtlMs', { timeout: 60_000 }, () => { await auth.clear() }) + it('PN/LID mappings stay persistent while session keys use a TTL', async (t) => { + if (!store || !redis) return t.skip('ZAPO_TEST_REDIS_* not set') + + const sessionId = nextSessionId('ttl-lid-pn') + const mapping = store.stores.lidPnMapping(sessionId) + await mapping.clear() + + await mapping.setLidUser('5511999999999', '123456789') + + assert.equal(await redis.pttl(`signal:lid-pn:${sessionId}`), -1) + + await mapping.clear() + }) + it('rejects a non-positive ttlMs at store construction', (t) => { if (!store || !redis) return t.skip('ZAPO_TEST_REDIS_* not set') const reused = redis diff --git a/packages/store-redis/src/createRedisStore.ts b/packages/store-redis/src/createRedisStore.ts index 3740e6c1..f1915697 100644 --- a/packages/store-redis/src/createRedisStore.ts +++ b/packages/store-redis/src/createRedisStore.ts @@ -7,6 +7,7 @@ import { WaContactRedisStore } from './contact.store' import { WaDeviceListRedisStore } from './device-list.store' import { WaGroupMetadataRedisStore } from './group-metadata.store' import { WaIdentityRedisStore } from './identity.store' +import { WaLidPnMappingRedisStore } from './lid-pn-mapping.store' import { WaMessageSecretRedisStore } from './message-secret.store' import { WaMessageRedisStore } from './message.store' import { WaPreKeyRedisStore } from './pre-key.store' @@ -89,6 +90,7 @@ export interface WaRedisStoreResult { readonly preKey: (sessionId: string) => WaPreKeyRedisStore readonly session: (sessionId: string) => WaSessionRedisStore readonly identity: (sessionId: string) => WaIdentityRedisStore + readonly lidPnMapping: (sessionId: string) => WaLidPnMappingRedisStore readonly signal: (sessionId: string) => WaSignalRedisStore readonly senderKey: (sessionId: string) => WaSenderKeyRedisStore readonly appState: (sessionId: string) => WaAppStateRedisStore @@ -113,7 +115,7 @@ function isRedis(value: Redis | RedisOptions): value is Redis { /** * Builds a Redis-backed {@link WaStoreBackend} bundle. Best fit when you * already run Redis for caching and want stateless app instances that can - * share session state - all 11 persistent domains and 4 cache domains live + * share session state - all 12 persistent domains and 4 cache domains live * under a single key namespace (controlled by `keyPrefix`). * * Cache domains use native Redis TTLs (`EX`/`PEXPIRE` on write); no @@ -181,6 +183,8 @@ export function createRedisStore(config: WaRedisStoreConfig): WaRedisStoreResult new WaSessionRedisStore(opts(sessionId, 'session', storeTtl.sessionMs)), identity: (sessionId) => new WaIdentityRedisStore(opts(sessionId, 'identity', storeTtl.identityMs)), + lidPnMapping: (sessionId) => + new WaLidPnMappingRedisStore(opts(sessionId, 'lidPnMapping')), signal: (sessionId) => new WaSignalRedisStore(opts(sessionId, 'signal', storeTtl.signalMs)), senderKey: (sessionId) => diff --git a/packages/store-redis/src/index.ts b/packages/store-redis/src/index.ts index a4a2d622..66a2f83d 100644 --- a/packages/store-redis/src/index.ts +++ b/packages/store-redis/src/index.ts @@ -4,6 +4,7 @@ export { WaAuthRedisStore } from './auth.store' export { WaPreKeyRedisStore } from './pre-key.store' export { WaSessionRedisStore } from './session.store' export { WaIdentityRedisStore } from './identity.store' +export { WaLidPnMappingRedisStore } from './lid-pn-mapping.store' export { WaSignalRedisStore } from './signal.store' export { WaSenderKeyRedisStore } from './sender-key.store' export { WaAppStateRedisStore } from './appstate.store' diff --git a/packages/store-redis/src/lid-pn-mapping.store.ts b/packages/store-redis/src/lid-pn-mapping.store.ts new file mode 100644 index 00000000..f1aa6286 --- /dev/null +++ b/packages/store-redis/src/lid-pn-mapping.store.ts @@ -0,0 +1,45 @@ +import type { WaLidPnMappingStore } from 'zapo-js/store' + +import { BaseRedisStore } from './BaseRedisStore' +import type { WaRedisStorageOptions } from './types' + +const LUA_SET_MAPPING = ` +local pn_field = 'p:' .. ARGV[1] +local lid_field = 'l:' .. ARGV[2] +local old_lid = redis.call('HGET', KEYS[1], pn_field) +if old_lid then + redis.call('HDEL', KEYS[1], 'l:' .. old_lid) +end +local old_pn = redis.call('HGET', KEYS[1], lid_field) +if old_pn then + redis.call('HDEL', KEYS[1], 'p:' .. old_pn) +end +redis.call('HSET', KEYS[1], pn_field, ARGV[2], lid_field, ARGV[1]) +return 1 +` + +/** Redis-backed PN/LID mapping store using one atomic bidirectional hash. */ +export class WaLidPnMappingRedisStore extends BaseRedisStore implements WaLidPnMappingStore { + private readonly mappingKey: string + + public constructor(options: WaRedisStorageOptions) { + super(options) + this.mappingKey = this.k('signal:lid-pn', this.sessionId) + } + + public async getLidUser(pnUser: string): Promise { + return this.redis.hget(this.mappingKey, `p:${pnUser}`) + } + + public async getPnUser(lidUser: string): Promise { + return this.redis.hget(this.mappingKey, `l:${lidUser}`) + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + await this.redis.eval(LUA_SET_MAPPING, 1, this.mappingKey, pnUser, lidUser) + } + + public async clear(): Promise { + await this.redis.del(this.mappingKey) + } +} diff --git a/packages/store-sqlite/src/__tests__/contracts.test.ts b/packages/store-sqlite/src/__tests__/contracts.test.ts index de566848..935a5331 100644 --- a/packages/store-sqlite/src/__tests__/contracts.test.ts +++ b/packages/store-sqlite/src/__tests__/contracts.test.ts @@ -6,12 +6,14 @@ import test from 'node:test' import { WaContactMemoryStore, + WaLidPnMappingMemoryStore, WaMessageMemoryStore, WaPrivacyTokenMemoryStore, WaThreadMemoryStore } from 'zapo-js/store' import { WaContactSqliteStore } from '../contact.store' +import { WaLidPnMappingSqliteStore } from '../lid-pn-mapping.store' import { WaMessageSqliteStore } from '../message.store' import { WaPrivacyTokenSqliteStore } from '../privacy-token.store' import { WaThreadSqliteStore } from '../thread.store' @@ -83,6 +85,55 @@ test('privacy token store contract parity between memory and sqlite providers', } }) +test('LID/PN mapping contract parity between memory and sqlite providers', async () => { + await runLidPnMappingStoreContract(async () => new WaLidPnMappingMemoryStore()) + + const dir = await mkdtemp(join(tmpdir(), 'zapo-lid-pn-contract-')) + try { + await runLidPnMappingStoreContract( + async () => + new WaLidPnMappingSqliteStore({ + path: join(dir, 'state.sqlite'), + sessionId: 'session-mapping', + driver: 'better-sqlite3' + }) + ) + } finally { + await rm(dir, { recursive: true, force: true }) + } +}) + +async function runLidPnMappingStoreContract( + factory: () => Promise< + { + getLidUser: (pnUser: string) => Promise + getPnUser: (lidUser: string) => Promise + setLidUser: (pnUser: string, lidUser: string) => Promise + clear: () => Promise + } & Destroyable + > +): Promise { + const store = await factory() + try { + assert.equal(await store.getLidUser('5511999999999'), null) + assert.equal(await store.getPnUser('111222'), null) + await store.setLidUser('5511999999999', '111222') + assert.equal(await store.getLidUser('5511999999999'), '111222') + assert.equal(await store.getPnUser('111222'), '5511999999999') + await store.setLidUser('5511999999999', '333444') + assert.equal(await store.getLidUser('5511999999999'), '333444') + assert.equal(await store.getPnUser('111222'), null) + await store.setLidUser('5511888888888', '333444') + assert.equal(await store.getLidUser('5511999999999'), null) + assert.equal(await store.getPnUser('333444'), '5511888888888') + await store.clear() + assert.equal(await store.getLidUser('5511888888888'), null) + assert.equal(await store.getPnUser('333444'), null) + } finally { + await store.destroy?.() + } +} + async function runMessageStoreContract( factory: () => Promise< { diff --git a/packages/store-sqlite/src/createSqliteStore.ts b/packages/store-sqlite/src/createSqliteStore.ts index 12f43bec..d93c809d 100644 --- a/packages/store-sqlite/src/createSqliteStore.ts +++ b/packages/store-sqlite/src/createSqliteStore.ts @@ -7,6 +7,7 @@ import { WaContactSqliteStore } from './contact.store' import { WaDeviceListSqliteStore } from './device-list.store' import { WaGroupMetadataSqliteStore } from './group-metadata.store' import { WaIdentitySqliteStore } from './identity.store' +import { WaLidPnMappingSqliteStore } from './lid-pn-mapping.store' import { WaMessageSecretSqliteStore } from './message-secret.store' import { WaMessageSqliteStore } from './message.store' import { WaPreKeySqliteStore } from './pre-key.store' @@ -99,6 +100,7 @@ export interface WaSqliteStoreResult { readonly preKey: (sessionId: string) => WaPreKeySqliteStore readonly session: (sessionId: string) => WaSessionSqliteStore readonly identity: (sessionId: string) => WaIdentitySqliteStore + readonly lidPnMapping: (sessionId: string) => WaLidPnMappingSqliteStore readonly signal: (sessionId: string) => WaSignalSqliteStore readonly senderKey: (sessionId: string) => SenderKeySqliteStore readonly appState: (sessionId: string) => WaAppStateSqliteStore @@ -117,8 +119,8 @@ export interface WaSqliteStoreResult { /** * Builds a SQLite-backed {@link WaStoreBackend} bundle: persistent stores - * for `auth`, `signal`, `preKey`, `session`, `identity`, `senderKey`, - * `appState`, `messages`, `threads`, `contacts`, `privacyToken`, plus + * for `auth`, `signal`, `preKey`, `session`, `identity`, `lidPnMapping`, + * `senderKey`, `appState`, `messages`, `threads`, `contacts`, `privacyToken`, plus * TTL-evicted caches for `retry`, `groupMetadata`, `deviceList`, * `messageSecret`. Feed the result into `createStore({ backends: { sqlite: * createSqliteStore(...) }, providers: { ... } })` from `zapo-js`. @@ -192,6 +194,8 @@ export function createSqliteStore(config: WaSqliteStoreConfig): WaSqliteStoreRes hasSessionBatchSize: batchSizes?.signalHasSession }), identity: (sessionId) => new WaIdentitySqliteStore(opts(sessionId, 'identity')), + lidPnMapping: (sessionId) => + new WaLidPnMappingSqliteStore(opts(sessionId, 'lidPnMapping')), signal: (sessionId) => new WaSignalSqliteStore(opts(sessionId, 'signal')), senderKey: (sessionId) => new SenderKeySqliteStore(opts(sessionId, 'senderKey')), appState: (sessionId) => new WaAppStateSqliteStore(opts(sessionId, 'appState')), diff --git a/packages/store-sqlite/src/index.ts b/packages/store-sqlite/src/index.ts index 1a84a298..9c2444a2 100644 --- a/packages/store-sqlite/src/index.ts +++ b/packages/store-sqlite/src/index.ts @@ -13,6 +13,7 @@ export { WaAuthSqliteStore } from './auth.store' export { WaPreKeySqliteStore } from './pre-key.store' export { WaSessionSqliteStore } from './session.store' export { WaIdentitySqliteStore } from './identity.store' +export { WaLidPnMappingSqliteStore } from './lid-pn-mapping.store' export { WaSignalSqliteStore } from './signal.store' export { SenderKeySqliteStore } from './sender-key.store' export { WaAppStateSqliteStore } from './appstate.store' diff --git a/packages/store-sqlite/src/lid-pn-mapping.store.ts b/packages/store-sqlite/src/lid-pn-mapping.store.ts new file mode 100644 index 00000000..d2dc4eca --- /dev/null +++ b/packages/store-sqlite/src/lid-pn-mapping.store.ts @@ -0,0 +1,54 @@ +import type { WaLidPnMappingStore } from 'zapo-js/store' +import { asOptionalString } from 'zapo-js/util' + +import { BaseSqliteStore } from './BaseSqliteStore' +import type { WaSqliteStorageOptions } from './types' + +/** SQLite-backed PN/LID mapping store scoped by Zapo session id. */ +export class WaLidPnMappingSqliteStore extends BaseSqliteStore implements WaLidPnMappingStore { + public constructor(options: WaSqliteStorageOptions) { + super(options, ['lidPnMapping']) + } + + public async getLidUser(pnUser: string): Promise { + const db = await this.getConnection() + const row = db.get>( + `SELECT lid_user + FROM signal_lid_pn_mapping + WHERE session_id = ? AND pn_user = ?`, + [this.options.sessionId, pnUser] + ) + return row ? (asOptionalString(row.lid_user) ?? null) : null + } + + public async getPnUser(lidUser: string): Promise { + const db = await this.getConnection() + const row = db.get>( + `SELECT pn_user + FROM signal_lid_pn_mapping + WHERE session_id = ? AND lid_user = ?`, + [this.options.sessionId, lidUser] + ) + return row ? (asOptionalString(row.pn_user) ?? null) : null + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + await this.withTransaction((db) => { + db.run( + `DELETE FROM signal_lid_pn_mapping + WHERE session_id = ? AND (pn_user = ? OR lid_user = ?)`, + [this.options.sessionId, pnUser, lidUser] + ) + db.run( + `INSERT INTO signal_lid_pn_mapping (session_id, pn_user, lid_user) + VALUES (?, ?, ?)`, + [this.options.sessionId, pnUser, lidUser] + ) + }) + } + + public async clear(): Promise { + const db = await this.getConnection() + db.run('DELETE FROM signal_lid_pn_mapping WHERE session_id = ?', [this.options.sessionId]) + } +} diff --git a/packages/store-sqlite/src/migrations.ts b/packages/store-sqlite/src/migrations.ts index ae8acf48..c59840ea 100644 --- a/packages/store-sqlite/src/migrations.ts +++ b/packages/store-sqlite/src/migrations.ts @@ -7,6 +7,7 @@ const UNIQUE_CONSTRAINT_ID_RE = /UNIQUE constraint failed: [A-Za-z_][A-Za-z0-9_] export type WaSqliteMigrationDomain = | 'auth' | 'signal' + | 'lidPnMapping' | 'senderKey' | 'appState' | 'retry' @@ -114,6 +115,21 @@ const SQLITE_MIGRATIONS: readonly WaSqliteMigration[] = [ `) } }, + { + id: '0017_signal_lid_pn_mapping', + domain: 'lidPnMapping', + up: (db) => { + db.exec(` + CREATE TABLE IF NOT EXISTS signal_lid_pn_mapping ( + session_id TEXT NOT NULL, + pn_user TEXT NOT NULL, + lid_user TEXT NOT NULL, + PRIMARY KEY (session_id, pn_user), + UNIQUE (session_id, lid_user) + ); + `) + } + }, { id: '0001_sender_key_schema', domain: 'senderKey', diff --git a/packages/store-sqlite/src/table-names.ts b/packages/store-sqlite/src/table-names.ts index 2db24f25..31f10bb5 100644 --- a/packages/store-sqlite/src/table-names.ts +++ b/packages/store-sqlite/src/table-names.ts @@ -11,6 +11,7 @@ const WA_SQLITE_TABLE_NAME_ORDER = Object.freeze([ 'signal_prekey', 'signal_session', 'signal_identity', + 'signal_lid_pn_mapping', 'sender_keys', 'sender_key_distribution', 'appstate_sync_keys', @@ -38,6 +39,7 @@ const WA_SQLITE_DEFAULT_TABLE_NAMES: Readonly> signal_prekey: 'signal_prekey', signal_session: 'signal_session', signal_identity: 'signal_identity', + signal_lid_pn_mapping: 'signal_lid_pn_mapping', sender_keys: 'sender_keys', sender_key_distribution: 'sender_key_distribution', appstate_sync_keys: 'appstate_sync_keys', diff --git a/packages/store-sqlite/src/types.ts b/packages/store-sqlite/src/types.ts index edeed65a..bcf79142 100644 --- a/packages/store-sqlite/src/types.ts +++ b/packages/store-sqlite/src/types.ts @@ -13,6 +13,7 @@ export type WaSqliteTableName = | 'signal_prekey' | 'signal_session' | 'signal_identity' + | 'signal_lid_pn_mapping' | 'sender_keys' | 'sender_key_distribution' | 'appstate_sync_keys' @@ -65,6 +66,7 @@ export interface WaSqliteStorageOptions { export type WaSqliteMigrationDomain = | 'auth' | 'signal' + | 'lidPnMapping' | 'senderKey' | 'appState' | 'retry' diff --git a/src/client/WaClient.ts b/src/client/WaClient.ts index b1725137..6c9306e6 100644 --- a/src/client/WaClient.ts +++ b/src/client/WaClient.ts @@ -678,7 +678,12 @@ class WaClientImpl extends EventEmitter { if (shouldClear('retry')) await this.stores.retry.clear() if (shouldClear('signal')) await this.stores.signal.clear() if (shouldClear('preKey')) await this.stores.preKey.clear() - if (shouldClear('session')) await this.stores.session.clear() + if (shouldClear('session')) { + // Keep mappings until every ratchet is gone. If session clearing fails, + // retaining the mapping cannot split surviving PN/LID session state. + await this.stores.session.clear() + await this.deps.signalAddressResolver.clear() + } if (shouldClear('identity')) await this.stores.identity.clear() if (shouldClear('senderKey')) await this.stores.senderKey.clear() if (shouldClear('threads')) await this.stores.threads.clear() diff --git a/src/client/WaClientFactory.ts b/src/client/WaClientFactory.ts index 30b530a7..e00417af 100644 --- a/src/client/WaClientFactory.ts +++ b/src/client/WaClientFactory.ts @@ -101,6 +101,7 @@ import { WA_PRIVACY_TOKEN_NOTIFICATION_TYPE } from '@protocol/constants' import { + canonicalizeSignalJid, isNewsletterJid, isOwnAccountJid, parseSignalAddressFromJid, @@ -118,8 +119,10 @@ import { SignalRotateKeyApi } from '@signal/api/SignalRotateKeyApi' import { SignalSessionSyncApi } from '@signal/api/SignalSessionSyncApi' import { SenderKeyManager } from '@signal/group/SenderKeyManager' import { createSignalSessionResolver, type SignalSessionResolver } from '@signal/session/resolver' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import { SignalProtocol } from '@signal/session/SignalProtocol' import type { WaStoredContactRecord } from '@store/contracts/contact.store' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { WaKeepAlive } from '@transport/keepalive/WaKeepAlive' import { buildAckNode } from '@transport/node/builders/global' import { buildPresenceNode } from '@transport/node/builders/presence' @@ -191,6 +194,7 @@ export interface WaClientDependencies { readonly messageClient: WaMessageClient readonly senderKeyManager: SenderKeyManager readonly signalProtocol: SignalProtocol + readonly signalAddressResolver: SignalAddressResolver readonly signalDigestSync: SignalDigestSyncApi readonly signalDeviceSync: SignalDeviceSyncApi readonly signalIdentitySync: SignalIdentitySyncApi @@ -539,10 +543,14 @@ export function buildWaClientDependencies(input: { defaultMaxAttempts: options.messageMaxAttempts, defaultRetryDelayMs: options.messageRetryDelayMs }) + const signalAddressResolver = new SignalAddressResolver( + sessionStore.lidPnMapping ?? new WaLidPnMappingMemoryStore() + ) const senderKeyManager = new SenderKeyManager(sessionStore.senderKey, { getFutureMessagesMax: () => abPropsCoordinator.getConfigValue('web_signal_future_messages_max'), - skipSignatureVerification: options.dangerous?.disableSenderKeySignatureVerification + skipSignatureVerification: options.dangerous?.disableSenderKeySignatureVerification, + addressResolver: signalAddressResolver }) const signalProtocol = new SignalProtocol( { @@ -551,7 +559,8 @@ export function buildWaClientDependencies(input: { session: sessionStore.session, identity: sessionStore.identity }, - logger + logger, + signalAddressResolver ) const signalSystemQuery: WaClientBuildRuntime['query'] = (node, timeoutMs) => runtime.query(node, timeoutMs, { useSystemId: true }) @@ -575,6 +584,7 @@ export function buildWaClientDependencies(input: { logger, query: signalSystemQuery, identityStore: sessionStore.identity, + addressResolver: signalAddressResolver, defaultTimeoutMs: options.signalFetchKeyBundlesTimeoutMs }) const signalMissingPreKeysSync = new SignalMissingPreKeysSyncApi({ @@ -680,6 +690,7 @@ export function buildWaClientDependencies(input: { identityStore: sessionStore.identity, signalIdentitySync, signalSessionSync, + addressResolver: signalAddressResolver, logger }) const fanoutResolver = createDeviceFanoutResolver({ @@ -757,6 +768,7 @@ export function buildWaClientDependencies(input: { buildMediaMessageContent(mediaMessageBuildOptions, content, ctx), senderKeyManager, signalProtocol, + signalAddressResolver, signalStore: sessionStore.signal, sessionStore: sessionStore.session, identityStore: sessionStore.identity, @@ -822,6 +834,7 @@ export function buildWaClientDependencies(input: { sessionStore: sessionStore.session, senderKeyStore: sessionStore.senderKey, signalProtocol, + signalAddressResolver, sessionResolver, signalDeviceSync, signalMissingPreKeysSync, @@ -1018,6 +1031,7 @@ export function buildWaClientDependencies(input: { getMeJid: () => getCurrentCredentials()?.meJid, getMeLid: () => getCurrentCredentials()?.meLid, signalProtocol, + signalAddressResolver, senderKeyManager, onDecryptFailure: (context: WaRetryDecryptFailureContext, error: unknown) => retryCoordinator.onDecryptFailure(context, error), @@ -1186,7 +1200,10 @@ export function buildWaClientDependencies(input: { }) await runtime.sendNode(ackNode) - const address = parseSignalAddressFromJid(parsed.fromJid) + const signalJid = canonicalizeSignalJid(parsed.fromJid) + const address = await signalAddressResolver.resolve( + parseSignalAddressFromJid(signalJid) + ) if (address.device !== 0) { logger.debug('identity-change: ignoring companion device', { @@ -1195,21 +1212,17 @@ export function buildWaClientDependencies(input: { return true } - const meJid = getCurrentCredentials()?.meJid - if (meJid) { - const meUser = toUserJid(meJid) - const fromUser = toUserJid(parsed.fromJid) - if (meUser === fromUser) { - logger.error('self primary identity changed, disconnecting') - void connectionManager?.getComms()?.stopComms() - await disconnectWithClientSideEffects( - WA_DISCONNECT_REASONS.PRIMARY_IDENTITY_KEY_CHANGE, - true, - null - ) - await clearStoredCredentialsWithClientSideEffects() - return true - } + const credentials = getCurrentCredentials() + if (isOwnAccountJid(signalJid, credentials?.meJid, credentials?.meLid)) { + logger.error('self primary identity changed, disconnecting') + void connectionManager?.getComms()?.stopComms() + await disconnectWithClientSideEffects( + WA_DISCONNECT_REASONS.PRIMARY_IDENTITY_KEY_CHANGE, + true, + null + ) + await clearStoredCredentialsWithClientSideEffects() + return true } const oldIdentity = await sessionStore.identity.getRemoteIdentity(address) @@ -1220,7 +1233,7 @@ export function buildWaClientDependencies(input: { }) await sessionStore.session.deleteSession(address) - const userJid = toUserJid(parsed.fromJid) + const userJid = toUserJid(signalJid) await trustedContactToken.reissueOnIdentityChange(userJid).catch((error) => { logger.warn('identity-change: reissue tc token failed', { message: toError(error).message @@ -1268,10 +1281,14 @@ export function buildWaClientDependencies(input: { }) await runtime.sendNode(ackNode) - const userJid = toUserJid(parsed.fromJid) + const userJid = toUserJid(parsed.fromJid, { + canonicalizeSignalServer: true + }) if (parsed.action === DEVICE_NOTIFICATION_ACTIONS.REMOVE) { - const baseAddress = parseSignalAddressFromJid(parsed.fromJid) + const baseAddress = await signalAddressResolver.resolve( + parseSignalAddressFromJid(parsed.fromJid) + ) for (const device of parsed.devices) { const address = { user: baseAddress.user, @@ -1418,6 +1435,7 @@ export function buildWaClientDependencies(input: { messageClient, senderKeyManager, signalProtocol, + signalAddressResolver, signalDigestSync, signalDeviceSync, signalIdentitySync, diff --git a/src/client/__tests__/client.test.ts b/src/client/__tests__/client.test.ts index 8834e8cb..cec0421f 100644 --- a/src/client/__tests__/client.test.ts +++ b/src/client/__tests__/client.test.ts @@ -870,6 +870,11 @@ function createClearStoredStateHarness(logoutStoreClear?: { receiptQueue: { take: () => [] }, + signalAddressResolver: { + clear: async () => { + cleared.push('lidPnMapping') + } + }, authClient: { clearStoredCredentials: async () => { cleared.push('auth') @@ -966,12 +971,24 @@ test('clearStoredState clears non-mailbox domains by default and preserves mailb 'signal', 'preKey', 'session', + 'lidPnMapping', 'identity', 'senderKey', 'privacyToken' ]) }) +test('clearStoredState retains PN/LID mappings when ratchet clearing fails', async () => { + const { fakeClient, cleared } = createClearStoredStateHarness() + fakeClient.stores.session.clear = async () => { + cleared.push('session') + throw new Error('session clear failed') + } + + await assert.rejects(getClearStoredStateMethod().call(fakeClient), /session clear failed/) + assert.equal(cleared.includes('lidPnMapping'), false) +}) + test('clearStoredState wipes mailbox when explicitly opted in', async () => { const { fakeClient, cleared } = createClearStoredStateHarness({ messages: true, @@ -1018,6 +1035,7 @@ test('clearStoredState respects logoutStoreClear domain toggles', async () => { 'signal', 'preKey', 'session', + 'lidPnMapping', 'identity', 'senderKey' ]) diff --git a/src/client/coordinators/WaMessageDispatchCoordinator.ts b/src/client/coordinators/WaMessageDispatchCoordinator.ts index 932e9b89..ded44020 100644 --- a/src/client/coordinators/WaMessageDispatchCoordinator.ts +++ b/src/client/coordinators/WaMessageDispatchCoordinator.ts @@ -77,6 +77,7 @@ import type { WaRetryReplayPayload } from '@retry/types' import type { SignalDeviceSyncApi } from '@signal/api/SignalDeviceSyncApi' import type { SenderKeyManager } from '@signal/group/SenderKeyManager' import type { SignalResolvedSessionTarget, SignalSessionResolver } from '@signal/session/resolver' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalProtocol } from '@signal/session/SignalProtocol' import type { SignalAddress } from '@signal/types' import type { WaDeviceListStore } from '@store/contracts/device-list.store' @@ -116,6 +117,7 @@ interface WaMessageDispatchCoordinatorOptions { ) => Promise readonly senderKeyManager: SenderKeyManager readonly signalProtocol: SignalProtocol + readonly signalAddressResolver?: SignalAddressResolver readonly signalStore: WaSignalStore readonly sessionStore: WaSessionStore readonly identityStore: WaIdentityStore @@ -1303,11 +1305,16 @@ export class WaMessageDispatchCoordinator { readonly address: SignalAddress } >() - const fanoutAddresses: SignalAddress[] = new Array(fanoutDeviceJids.length) + const parsedFanoutAddresses: SignalAddress[] = new Array(fanoutDeviceJids.length) + for (let index = 0; index < fanoutDeviceJids.length; index += 1) { + parsedFanoutAddresses[index] = parseSignalAddressFromJid(fanoutDeviceJids[index]) + } + const fanoutAddresses = this.deps.signalAddressResolver + ? await this.deps.signalAddressResolver.resolveMany(parsedFanoutAddresses) + : parsedFanoutAddresses for (let index = 0; index < fanoutDeviceJids.length; index += 1) { const jid = fanoutDeviceJids[index] - const address = parseSignalAddressFromJid(jid) - fanoutAddresses[index] = address + const address = fanoutAddresses[index] fanoutTargetsByAddressKey.set(signalAddressKey(address), { jid, address }) } const pendingAddresses = @@ -1781,7 +1788,8 @@ export class WaMessageDispatchCoordinator { this.deps.identityStore, snapshot.updatedAtMs, localIdentity, - this.deps.getIcdcHashLength?.() + this.deps.getIcdcHashLength?.(), + this.deps.signalAddressResolver ) } catch (error) { this.deps.logger.trace('icdc resolution failed', { diff --git a/src/client/coordinators/WaRetryCoordinator.ts b/src/client/coordinators/WaRetryCoordinator.ts index a1f398d8..bea068b9 100644 --- a/src/client/coordinators/WaRetryCoordinator.ts +++ b/src/client/coordinators/WaRetryCoordinator.ts @@ -15,6 +15,7 @@ import { normalizeDeviceJid, parseJidFull, parseSignalAddressFromJid, + signalAddressKey, toUserJid } from '@protocol/jid' import { @@ -37,6 +38,7 @@ import type { SignalDeviceSyncApi } from '@signal/api/SignalDeviceSyncApi' import type { SignalMissingPreKeysSyncApi } from '@signal/api/SignalMissingPreKeysSyncApi' import { generatePreKeyPair } from '@signal/registration/keygen' import type { SignalSessionResolver } from '@signal/session/resolver' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalProtocol } from '@signal/session/SignalProtocol' import type { SignalPreKeyBundle } from '@signal/types' import type { WaPreKeyStore } from '@store/contracts/pre-key.store' @@ -59,6 +61,7 @@ interface WaRetryCoordinatorOptions { readonly sessionStore: WaSessionStore readonly senderKeyStore: WaSenderKeyStore readonly signalProtocol: SignalProtocol + readonly signalAddressResolver?: SignalAddressResolver readonly sessionResolver: SignalSessionResolver readonly signalDeviceSync: SignalDeviceSyncApi readonly signalMissingPreKeysSync: SignalMissingPreKeysSyncApi @@ -509,6 +512,9 @@ export class WaRetryCoordinator { request, requesterJid, requesterAddress, + this.deps.signalAddressResolver + ? await this.deps.signalAddressResolver.resolve(requesterAddress) + : requesterAddress, requesterNormalizedDeviceJid ) if (!sessionReady) { @@ -660,6 +666,7 @@ export class WaRetryCoordinator { request: WaParsedRetryRequest, requesterJid: string, requesterAddress: ReturnType, + requesterSignalAddress: ReturnType, requesterNormalizedDeviceJid: string ): Promise { const requestLogger = this.deps.logger.child({ @@ -668,13 +675,13 @@ export class WaRetryCoordinator { requester: requesterJid }) const [, currentSession] = await Promise.all([ - this.markRetryRequesterSenderKeyAsStale(request, requesterJid, requesterAddress), - this.deps.sessionStore.getSession(requesterAddress) + this.markRetryRequesterSenderKeyAsStale(request, requesterJid, requesterSignalAddress), + this.deps.sessionStore.getSession(requesterSignalAddress) ]) const regIdMismatch = !!currentSession && request.regId > 0 && currentSession.remote.regId !== request.regId if (regIdMismatch && !request.keyBundle) { - await this.deps.sessionStore.deleteSession(requesterAddress) + await this.deps.sessionStore.deleteSession(requesterSignalAddress) } if (request.keyBundle) { if (!request.keyBundle.key || !request.keyBundle.skey.signature) { @@ -689,7 +696,7 @@ export class WaRetryCoordinator { ...getRemoteRetryReasonLogFields(request.retryReason) } ) - await this.deps.sessionStore.deleteSession(requesterAddress) + await this.deps.sessionStore.deleteSession(requesterSignalAddress) return false } if (regIdMismatch) { @@ -700,14 +707,14 @@ export class WaRetryCoordinator { ...getRemoteRetryReasonLogFields(request.retryReason) } ) - await this.deps.sessionStore.deleteSession(requesterAddress) + await this.deps.sessionStore.deleteSession(requesterSignalAddress) return false } } else if (regIdMismatch) { - await this.deps.sessionStore.deleteSession(requesterAddress) + await this.deps.sessionStore.deleteSession(requesterSignalAddress) } await this.deps.signalProtocol.establishOutgoingSession( - requesterAddress, + requesterSignalAddress, { regId: request.regId, identity: request.keyBundle.identity, @@ -727,6 +734,7 @@ export class WaRetryCoordinator { request, requesterJid, requesterAddress, + requesterSignalAddress, requesterNormalizedDeviceJid ) } @@ -737,6 +745,7 @@ export class WaRetryCoordinator { request, requesterJid, requesterAddress, + requesterSignalAddress, requesterNormalizedDeviceJid ) } @@ -750,11 +759,12 @@ export class WaRetryCoordinator { if (!fetched) { return false } - await this.deps.signalProtocol.establishOutgoingSession(requesterAddress, fetched) + await this.deps.signalProtocol.establishOutgoingSession(requesterSignalAddress, fetched) return this.applySessionBaseKeyPolicy( request, requesterJid, requesterAddress, + requesterSignalAddress, requesterNormalizedDeviceJid ) } @@ -763,36 +773,35 @@ export class WaRetryCoordinator { request: WaParsedRetryRequest, requesterJid: string, requesterAddress: ReturnType, + requesterSignalAddress: ReturnType, requesterNormalizedDeviceJid: string ): Promise { if (request.retryCount < 2) { return true } - const currentSession = await this.deps.sessionStore.getSession(requesterAddress) + const currentSession = await this.deps.sessionStore.getSession(requesterSignalAddress) const sessionBaseKey = currentSession?.aliceBaseKey ?? null if (!sessionBaseKey) { return true } const expiresAtMs = Date.now() + this.retryTtlMs + const requesterSessionKey = signalAddressKey(requesterSignalAddress) if (request.retryCount === 2) { this.setRetrySessionBaseKey( request.originalMsgId, - requesterNormalizedDeviceJid, + requesterSessionKey, sessionBaseKey, expiresAtMs ) return true } - const saved = this.getRetrySessionBaseKey( - request.originalMsgId, - requesterNormalizedDeviceJid - ) + const saved = this.getRetrySessionBaseKey(request.originalMsgId, requesterSessionKey) if (!saved || !uint8Equal(saved.baseKey, sessionBaseKey)) { return true } - await this.deps.sessionStore.deleteSession(requesterAddress) + await this.deps.sessionStore.deleteSession(requesterSignalAddress) this.deps.logger.debug('retry request forcing session refresh due to repeated base key', { id: request.stanzaId, originalMsgId: request.originalMsgId, @@ -809,7 +818,7 @@ export class WaRetryCoordinator { if (!fetched) { return false } - await this.deps.signalProtocol.establishOutgoingSession(requesterAddress, fetched) + await this.deps.signalProtocol.establishOutgoingSession(requesterSignalAddress, fetched) return true } @@ -1032,20 +1041,17 @@ export class WaRetryCoordinator { } } - private retrySessionBaseKeyMapKey( - originalMsgId: string, - requesterNormalizedDeviceJid: string - ): string { - return `${originalMsgId}|${requesterNormalizedDeviceJid}` + private retrySessionBaseKeyMapKey(originalMsgId: string, requesterSessionKey: string): string { + return `${originalMsgId}|${requesterSessionKey}` } private setRetrySessionBaseKey( originalMsgId: string, - requesterNormalizedDeviceJid: string, + requesterSessionKey: string, baseKey: Uint8Array, expiresAtMs: number ): void { - const key = this.retrySessionBaseKeyMapKey(originalMsgId, requesterNormalizedDeviceJid) + const key = this.retrySessionBaseKeyMapKey(originalMsgId, requesterSessionKey) setBoundedMapEntry( this.retrySessionBaseKeys, key, @@ -1059,9 +1065,9 @@ export class WaRetryCoordinator { private getRetrySessionBaseKey( originalMsgId: string, - requesterNormalizedDeviceJid: string + requesterSessionKey: string ): RetrySessionBaseKeySnapshot | null { - const key = this.retrySessionBaseKeyMapKey(originalMsgId, requesterNormalizedDeviceJid) + const key = this.retrySessionBaseKeyMapKey(originalMsgId, requesterSessionKey) const entry = this.retrySessionBaseKeys.get(key) if (!entry) { return null diff --git a/src/client/coordinators/__tests__/retry-coordinator.test.ts b/src/client/coordinators/__tests__/retry-coordinator.test.ts index 9969079a..a95e054d 100644 --- a/src/client/coordinators/__tests__/retry-coordinator.test.ts +++ b/src/client/coordinators/__tests__/retry-coordinator.test.ts @@ -588,6 +588,7 @@ type RetrySessionInternals = { request: WaParsedRetryRequest, requesterJid: string, requesterAddress: ReturnType['address'], + requesterSignalAddress: ReturnType['address'], requesterNormalizedDeviceJid: string ) => Promise } @@ -634,6 +635,7 @@ test('retry session update reuses an existing compatible session instead of re-k buildKeyBundleRequest(2, requesterJid), requesterJid, parsed.address, + parsed.address, parsed.normalizedJid ) @@ -643,7 +645,7 @@ test('retry session update reuses an existing compatible session instead of re-k assert.deepEqual(establishOptions, [{ reuseExisting: true }]) }) -test('retry session update resets the session once the base key repeats at retry 3', async () => { +test('retry session update detects a repeated base key across PN/LID aliases', async () => { const existingSession = { remote: { regId: 555, pubKey: new Uint8Array(33) }, aliceBaseKey: new Uint8Array([7, 7, 7]) @@ -671,7 +673,7 @@ test('retry session update resets the session once the base key repeats at retry signalMissingPreKeysSync: { fetchMissingPreKeys: async () => { fetchCount += 1 - return [{ devices: [{ deviceJid: '551100000000:3@s.whatsapp.net', bundle: {} }] }] + return [{ devices: [{ deviceJid: '778899:3@lid', bundle: {} }] }] } } as never, messageClient: {} as never, @@ -680,24 +682,28 @@ test('retry session update resets the session once the base key repeats at retry }) const internals = coordinator as unknown as RetrySessionInternals - const requesterJid = '551100000000:3@s.whatsapp.net' - const parsed = parseJidFull(requesterJid) - const run = (retryCount: number): Promise => + const pn = parseJidFull('551100000000:3@s.whatsapp.net') + const lid = parseJidFull('778899:3@lid') + const run = ( + retryCount: number, + requester: ReturnType + ): Promise => internals.updateLocalSessionFromRetryRequest( - buildKeyBundleRequest(retryCount, requesterJid), - requesterJid, - parsed.address, - parsed.normalizedJid + buildKeyBundleRequest(retryCount, requester.normalizedJid), + requester.normalizedJid, + requester.address, + lid.address, + requester.normalizedJid ) // Retry 2 only records the session base key. - await run(2) + await run(2, pn) assert.equal(deleteCount, 0) assert.equal(fetchCount, 0) // Retry 3 sees the same (reused) base key and forces a clean session: // delete + fetch fresh prekeys + re-establish. - const ready = await run(3) + const ready = await run(3, lid) assert.equal(ready, true) assert.equal(deleteCount, 1) assert.equal(fetchCount, 1) diff --git a/src/message/crypto/icdc.ts b/src/message/crypto/icdc.ts index a79e32d5..70fd2bcf 100644 --- a/src/message/crypto/icdc.ts +++ b/src/message/crypto/icdc.ts @@ -1,6 +1,7 @@ import { sha256, toRawPubKey } from '@crypto' import type { Proto } from '@proto' import { parseSignalAddressFromJid } from '@protocol/jid' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalAddress } from '@signal/types' import type { WaIdentityStore } from '@store/contracts/identity.store' @@ -35,15 +36,23 @@ export async function resolveIcdcMeta( identityStore: WaIdentityStore, updatedAtMs: number | undefined, localIdentity?: { readonly address: SignalAddress; readonly pubKey: Uint8Array }, - hashLength?: number + hashLength?: number, + addressResolver?: SignalAddressResolver ): Promise { if (deviceJids.length === 0) { return null } - const addresses: SignalAddress[] = new Array(deviceJids.length) + const parsedAddresses: SignalAddress[] = new Array(deviceJids.length) for (let i = 0; i < deviceJids.length; i += 1) { - addresses[i] = parseSignalAddressFromJid(deviceJids[i]) + parsedAddresses[i] = parseSignalAddressFromJid(deviceJids[i]) } + const addresses = addressResolver + ? await addressResolver.resolveMany(parsedAddresses) + : parsedAddresses + const localAddress = + localIdentity && addressResolver + ? await addressResolver.resolve(localIdentity.address) + : localIdentity?.address const remoteKeys = await identityStore.getRemoteIdentities(addresses) const keys: Uint8Array[] = [] for (let i = 0; i < addresses.length; i += 1) { @@ -52,8 +61,10 @@ export async function resolveIcdcMeta( keys.push(key) } else if ( localIdentity && - addresses[i].user === localIdentity.address.user && - addresses[i].device === localIdentity.address.device + localAddress && + addresses[i].user === localAddress.user && + addresses[i].server === localAddress.server && + addresses[i].device === localAddress.device ) { keys.push(localIdentity.pubKey) } diff --git a/src/message/primitives/__tests__/incoming.test.ts b/src/message/primitives/__tests__/incoming.test.ts index 3f394333..aaba037c 100644 --- a/src/message/primitives/__tests__/incoming.test.ts +++ b/src/message/primitives/__tests__/incoming.test.ts @@ -5,6 +5,8 @@ import type { WaIncomingMessageEvent, WaIncomingUnavailableMessageEvent } from ' import { createNoopLogger } from '@infra/log/types' import { buildRecoveredIncomingEvent, handleIncomingMessageAck } from '@message/primitives/incoming' import { proto } from '@proto' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import type { BinaryNode } from '@transport/types' function createEncryptedMessageNode(): BinaryNode { @@ -123,6 +125,281 @@ test('1:1 incoming message strips the device from remoteJid and keeps it in send assert.equal(key.participant, undefined) }) +test('direct sender_lid mapping is learned before Signal decrypt', async () => { + const calls: string[] = [] + const encrypted = createEncryptedMessageNode() + const node: BinaryNode = { + ...encrypted, + attrs: { ...encrypted.attrs, sender_lid: '778899@lid' } + } + + await handleIncomingMessageAck(node, { + logger: createNoopLogger(), + sendNode: async () => undefined, + signalAddressResolver: { + learnMessageJidPair: async (firstJid: string, secondJid: string) => { + calls.push(`learn:${firstJid}:${secondJid}`) + return true + } + } as never, + signalProtocol: { + decryptMessage: async () => { + calls.push('decrypt') + return paddedPlaintext({ conversation: 'hi' }) + } + } as never + }) + + assert.deepEqual(calls, ['learn:551100000000@s.whatsapp.net:778899@lid', 'decrypt']) +}) + +test('mapping-store failures do not interrupt incoming Signal decrypt', async () => { + const calls: string[] = [] + const warnings: Array<{ + readonly message: string + readonly id?: unknown + readonly from?: unknown + readonly error?: unknown + }> = [] + const logger = createNoopLogger() + logger.warn = (message, context) => { + warnings.push({ message, id: context?.id, from: context?.from, error: context?.message }) + } + const encrypted = createEncryptedMessageNode() + + const handled = await handleIncomingMessageAck( + { + ...encrypted, + attrs: { ...encrypted.attrs, sender_lid: '778899@lid' } + }, + { + logger, + sendNode: async () => undefined, + signalAddressResolver: { + learnMessageJidPair: async () => { + calls.push('learn') + throw new Error('mapping store unavailable') + } + } as never, + signalProtocol: { + decryptMessage: async () => { + calls.push('decrypt') + return paddedPlaintext({ conversation: 'hi' }) + } + } as never + } + ) + + assert.equal(handled, true) + assert.deepEqual(calls, ['learn', 'decrypt']) + assert.deepEqual(warnings, [ + { + message: 'failed to learn incoming PN/LID mapping', + id: encrypted.attrs.id, + from: encrypted.attrs.from, + error: 'mapping store unavailable' + } + ]) +}) + +test('recipient_latest_lid becomes the canonical Signal address after peer metadata', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + const encrypted = createEncryptedMessageNode() + + await handleIncomingMessageAck( + { + ...encrypted, + attrs: { + ...encrypted.attrs, + recipient: '5511222222222@s.whatsapp.net', + peer_recipient_lid: '101@lid', + recipient_latest_lid: '202@lid' + } + }, + { + logger: createNoopLogger(), + sendNode: async () => undefined, + getMeJid: () => encrypted.attrs.from, + signalAddressResolver: addressResolver, + signalProtocol: { + decryptMessage: async () => paddedPlaintext({ conversation: 'hi' }) + } as never + } + ) + + assert.deepEqual( + await addressResolver.resolve({ + user: '5511222222222', + server: 's.whatsapp.net', + device: 7 + }), + { user: '202', server: 'lid', device: 7 } + ) +}) + +test('direct recipient metadata takes conservative precedence over peer metadata', async () => { + const cases = [ + { + from: '5511999999999@s.whatsapp.net', + recipient: '5511222222222@s.whatsapp.net', + recipientAttr: { recipient_lid: '101@lid', peer_recipient_lid: '202@lid' }, + getMeJid: () => '5511999999999@s.whatsapp.net', + getMeLid: undefined, + expected: '5511222222222@s.whatsapp.net:101@lid' + }, + { + from: '999@lid', + recipient: '101@lid', + recipientAttr: { + recipient_pn: '5511222222222@s.whatsapp.net', + peer_recipient_pn: '5511333333333@s.whatsapp.net' + }, + getMeJid: undefined, + getMeLid: () => '999@lid', + expected: '101@lid:5511222222222@s.whatsapp.net' + } + ] as const + + for (const current of cases) { + const calls: string[] = [] + await handleIncomingMessageAck( + { + tag: 'message', + attrs: { + id: `msg-recipient-${calls.length}`, + from: current.from, + recipient: current.recipient, + recipient_latest_lid: '303@lid', + ...current.recipientAttr + }, + content: [{ tag: 'enc', attrs: { type: 'msg' }, content: new Uint8Array([1]) }] + }, + { + logger: createNoopLogger(), + sendNode: async () => undefined, + getMeJid: current.getMeJid, + getMeLid: current.getMeLid, + signalAddressResolver: { + learnMessageJidPair: async (firstJid: string, secondJid: string) => { + calls.push(`${firstJid}:${secondJid}`) + return true + }, + learnPeerRecipientJidPair: async () => { + calls.push('unexpected-peer-mapping') + return true + } + } as never, + signalProtocol: { + decryptMessage: async () => paddedPlaintext({ conversation: 'hi' }) + } as never + } + ) + assert.deepEqual(calls, [current.expected]) + } +}) + +test('group participant mapping is learned for another device of this account', async () => { + const calls: string[] = [] + await handleIncomingMessageAck( + { + tag: 'message', + attrs: { + id: 'msg-own-group-participant', + from: '12345@g.us', + participant: '999:2@lid', + participant_pn: '5511999999999@s.whatsapp.net' + }, + content: [{ tag: 'enc', attrs: { type: 'msg' }, content: new Uint8Array([1]) }] + }, + { + logger: createNoopLogger(), + sendNode: async () => undefined, + getMeLid: () => '999@lid', + signalAddressResolver: { + learnMessageJidPair: async (firstJid: string, secondJid: string) => { + calls.push(`learn:${firstJid}:${secondJid}`) + return true + } + } as never, + signalProtocol: { + decryptMessage: async () => { + calls.push('decrypt') + return paddedPlaintext({ conversation: 'hi' }) + } + } as never + } + ) + + assert.deepEqual(calls, ['learn:999:2@lid:5511999999999@s.whatsapp.net', 'decrypt']) +}) + +test('peer-recipient metadata is ignored when the message author is not this account', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + const encrypted = createEncryptedMessageNode() + + await handleIncomingMessageAck( + { + ...encrypted, + attrs: { + ...encrypted.attrs, + recipient: '5511222222222@s.whatsapp.net', + peer_recipient_lid: '101@lid' + } + }, + { + logger: createNoopLogger(), + sendNode: async () => undefined, + getMeJid: () => '5511999999999@s.whatsapp.net', + signalAddressResolver: addressResolver, + signalProtocol: { + decryptMessage: async () => paddedPlaintext({ conversation: 'hi' }) + } as never + } + ) + + const recipient = { user: '5511222222222', server: 's.whatsapp.net', device: 7 } as const + assert.strictEqual(await addressResolver.resolve(recipient), recipient) +}) + +test('self sender_lid metadata takes precedence over peer-recipient metadata', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + const encrypted = createEncryptedMessageNode() + + await handleIncomingMessageAck( + { + ...encrypted, + attrs: { + ...encrypted.attrs, + sender_lid: '909@lid', + recipient: '5511222222222@s.whatsapp.net', + peer_recipient_lid: '101@lid' + } + }, + { + logger: createNoopLogger(), + sendNode: async () => undefined, + getMeJid: () => encrypted.attrs.from, + signalAddressResolver: addressResolver, + signalProtocol: { + decryptMessage: async () => paddedPlaintext({ conversation: 'hi' }) + } as never + } + ) + + assert.equal( + ( + await addressResolver.resolve({ + user: '551100000000', + server: 's.whatsapp.net', + device: 0 + }) + ).user, + '909' + ) + const recipient = { user: '5511222222222', server: 's.whatsapp.net', device: 0 } as const + assert.strictEqual(await addressResolver.resolve(recipient), recipient) +}) + test('1:1 message authored by my own other device is fromMe with the recipient as remoteJid', async () => { const emitted: WaIncomingMessageEvent[] = [] const handled = await handleIncomingMessageAck( diff --git a/src/message/primitives/incoming.ts b/src/message/primitives/incoming.ts index 51057330..a0f60b95 100644 --- a/src/message/primitives/incoming.ts +++ b/src/message/primitives/incoming.ts @@ -12,9 +12,10 @@ import { unwrapDeviceSentMessage } from '@message/encode/device-sent' import { unpadPkcs7 } from '@message/encode/padding' import { processIncomingNewsletterMessage } from '@message/kinds/newsletter' import { proto } from '@proto' -import { WA_MESSAGE_TAGS, WA_MESSAGE_TYPES } from '@protocol/constants' +import { WA_DEFAULTS, WA_MESSAGE_TAGS, WA_MESSAGE_TYPES } from '@protocol/constants' import { canonicalizeOwnAccountJid, + canonicalizeSignalServer, isBroadcastJid, isGroupJid, isNewsletterJid, @@ -25,6 +26,7 @@ import { } from '@protocol/jid' import type { WaRetryDecryptFailureContext } from '@retry/types' import type { SenderKeyManager } from '@signal/group/SenderKeyManager' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalProtocol } from '@signal/session/SignalProtocol' import type { SignalAddress } from '@signal/types' import { buildAckNode, buildReceiptNode } from '@transport/node/builders/global' @@ -38,6 +40,7 @@ interface WaIncomingMessageAckHandlerOptions { readonly getMeJid?: () => string | null | undefined readonly getMeLid?: () => string | null | undefined readonly signalProtocol?: SignalProtocol + readonly signalAddressResolver?: SignalAddressResolver readonly senderKeyManager?: SenderKeyManager readonly onDecryptFailure?: ( context: WaRetryDecryptFailureContext, @@ -58,6 +61,14 @@ interface MessageIdentityAttrs { readonly pushName: string | undefined } +function isSignalLidJid(jid: string): boolean { + const address = parseSignalAddressFromJid(jid) + return ( + canonicalizeSignalServer(address.server ?? WA_DEFAULTS.HOST_DOMAIN) === + WA_DEFAULTS.LID_SERVER + ) +} + // Addressing fields use conditional includes so they're absent (not own-undefined) when // the source attr isn't present — the `...keyIdentity` spread then only injects defined // keys into the event's `key`. `pushName` is destructured straight to the event top-level @@ -65,7 +76,11 @@ interface MessageIdentityAttrs { function extractMessageIdentityAttrs(attrs: BinaryNode['attrs']): MessageIdentityAttrs { const rawRemoteJidAlt = attrs.sender_pn ?? attrs.sender_lid const rawParticipantAlt = attrs.participant_pn ?? attrs.participant_lid - const rawRecipientAlt = attrs.peer_recipient_pn ?? attrs.peer_recipient_lid + const rawRecipientAlt = + attrs.recipient_pn ?? + attrs.recipient_lid ?? + attrs.peer_recipient_pn ?? + attrs.peer_recipient_lid const rawRecipient = attrs.recipient const senderUsername = attrs.participant_username ?? attrs.username return { @@ -78,6 +93,85 @@ function extractMessageIdentityAttrs(attrs: BinaryNode['attrs']): MessageIdentit } } +async function learnMessageLidPnMappings( + node: BinaryNode, + options: WaIncomingMessageAckHandlerOptions +): Promise { + const resolver = options.signalAddressResolver + if (!resolver) return + const attrs = node.attrs + const fromJid = attrs.from + if (!fromJid) return + const hasParticipantAuthor = isGroupJid(fromJid) || isBroadcastJid(fromJid) + const authorJid = hasParticipantAuthor ? attrs.participant : fromJid + if (!authorJid) return + + const authorIsLid = isSignalLidJid(authorJid) + if (hasParticipantAuthor) { + const authorAlt = authorIsLid ? attrs.participant_pn : attrs.participant_lid + if (authorAlt) await resolver.learnMessageJidPair(authorJid, authorAlt) + return + } + + const authorIsMe = isOwnAccountJid( + toUserJid(authorJid, { canonicalizeSignalServer: true }), + options.getMeJid?.(), + options.getMeLid?.() + ) + if (authorIsLid) { + if (!authorIsMe) { + if (attrs.sender_pn) await resolver.learnMessageJidPair(authorJid, attrs.sender_pn) + return + } + if (!attrs.recipient) return + if (attrs.recipient_pn) { + await resolver.learnMessageJidPair(attrs.recipient, attrs.recipient_pn) + return + } + if (attrs.peer_recipient_pn) { + await learnPeerRecipientMapping( + resolver, + attrs.recipient, + attrs.peer_recipient_pn, + attrs.recipient_latest_lid + ) + } + return + } + + // WhatsApp Web applies sender_lid last for PN authors, so it takes precedence + // over every peer-recipient field on a self-authored message. + if (attrs.sender_lid) { + await resolver.learnMessageJidPair(authorJid, attrs.sender_lid) + return + } + if (!authorIsMe || !attrs.recipient) return + if (attrs.recipient_lid) { + await resolver.learnMessageJidPair(attrs.recipient, attrs.recipient_lid) + return + } + if (attrs.peer_recipient_lid) { + await learnPeerRecipientMapping( + resolver, + attrs.recipient, + attrs.peer_recipient_lid, + attrs.recipient_latest_lid + ) + } +} + +async function learnPeerRecipientMapping( + resolver: SignalAddressResolver, + recipientJid: string, + peerRecipientJid: string, + latestLid?: string +): Promise { + await resolver.learnPeerRecipientJidPair(recipientJid, peerRecipientJid) + if (!latestLid) return + const pnJid = isSignalLidJid(recipientJid) ? peerRecipientJid : recipientJid + await resolver.learnPeerRecipientJidPair(pnJid, latestLid) +} + type MessageKeyIdentity = Omit /** @@ -520,6 +614,16 @@ export async function handleIncomingMessageAck( return handleIncomingNewsletterMessage(node, options) } + try { + await learnMessageLidPnMappings(node, options) + } catch (error) { + options.logger.warn('failed to learn incoming PN/LID mapping', { + id: node.attrs.id, + from: node.attrs.from, + message: toError(error).message + }) + } + let shouldSendStandardReceipt = true const nodeContent = node.content if (Array.isArray(nodeContent) && nodeContent.length > 0) { diff --git a/src/protocol/__tests__/protocol.test.ts b/src/protocol/__tests__/protocol.test.ts index 0003e111..b934d898 100644 --- a/src/protocol/__tests__/protocol.test.ts +++ b/src/protocol/__tests__/protocol.test.ts @@ -74,6 +74,15 @@ test('canonicalizeOwnAccountJid maps own PN device JIDs to LID', () => { ) }) +test('isOwnAccountJid recognizes hosted aliases for the account identities', () => { + const meJid = '5511@s.whatsapp.net' + const meLid = '1330@lid' + + assert.equal(isOwnAccountJid('5511:99@hosted', meJid, meLid), true) + assert.equal(isOwnAccountJid('1330@hosted.lid', meJid, meLid), true) + assert.equal(isOwnAccountJid('5599@hosted', meJid, meLid), false) +}) + test('jid split and normalization helpers', () => { assert.deepEqual(splitJid('123@s.whatsapp.net'), { user: '123', diff --git a/src/protocol/jid.ts b/src/protocol/jid.ts index 9b1c65ad..8af3f537 100644 --- a/src/protocol/jid.ts +++ b/src/protocol/jid.ts @@ -14,6 +14,8 @@ const KNOWN_SERVERS: Record = { [WA_DEFAULTS.BOT_SERVER]: WA_DEFAULTS.BOT_SERVER } +const WA_CANONICAL_SIGNAL_USER_JID_OPTIONS = { canonicalizeSignalServer: true } as const + /** * Returns the canonical reference for known server strings, avoiding * thousands of duplicate sliced copies (e.g. "lid", "s.whatsapp.net") @@ -223,18 +225,18 @@ export function toUserJid( /** * True when `jid` is the account's own user, matching the `meJid` (pn) or - * `meLid` (lid) identity device-insensitively. Mirrors WhatsApp Web's - * `isMeAccount`. + * `meLid` (lid) identity device-insensitively, including hosted Signal server + * aliases. Mirrors WhatsApp Web's `isMeAccount`. */ export function isOwnAccountJid( jid: string, meJid: string | null | undefined, meLid: string | null | undefined ): boolean { - const candidateUser = toUserJid(jid) + const candidateUser = toUserJid(jid, WA_CANONICAL_SIGNAL_USER_JID_OPTIONS) return ( - (!!meJid && toUserJid(meJid) === candidateUser) || - (!!meLid && toUserJid(meLid) === candidateUser) + (!!meJid && toUserJid(meJid, WA_CANONICAL_SIGNAL_USER_JID_OPTIONS) === candidateUser) || + (!!meLid && toUserJid(meLid, WA_CANONICAL_SIGNAL_USER_JID_OPTIONS) === candidateUser) ) } diff --git a/src/signal/api/SignalIdentitySyncApi.ts b/src/signal/api/SignalIdentitySyncApi.ts index adf2b07e..0669f3b1 100644 --- a/src/signal/api/SignalIdentitySyncApi.ts +++ b/src/signal/api/SignalIdentitySyncApi.ts @@ -2,9 +2,10 @@ import { toSerializedPubKey } from '@crypto/core/keys' import type { Logger } from '@infra/log/types' import { PromiseDedup } from '@infra/perf/PromiseDedup' import { WA_DEFAULTS, WA_IQ_TYPES, WA_NODE_TAGS, WA_XMLNS } from '@protocol/constants' -import { canonicalizeSignalJid, parseSignalAddressFromJid } from '@protocol/jid' +import { canonicalizeSignalJid, parseSignalAddressFromJid, signalAddressKey } from '@protocol/jid' import { decodeExactLength, parseUint } from '@signal/api/codec' import { SIGNAL_KEY_BUNDLE_TYPE_LENGTH, SIGNAL_KEY_DATA_LENGTH } from '@signal/api/constants' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { WaIdentityStore } from '@store/contracts/identity.store' import { findNodeChild, getNodeChildrenByTag } from '@transport/node/helpers' import { assertIqResult } from '@transport/node/query' @@ -20,6 +21,7 @@ interface SignalIdentitySyncApiOptions { readonly logger: Logger readonly query: (node: BinaryNode, timeoutMs?: number) => Promise readonly identityStore?: WaIdentityStore + readonly addressResolver?: SignalAddressResolver readonly defaultTimeoutMs?: number readonly hostDomain?: string } @@ -33,6 +35,7 @@ export class SignalIdentitySyncApi { private readonly logger: SignalIdentitySyncApiOptions['logger'] private readonly query: SignalIdentitySyncApiOptions['query'] private readonly identityStore?: WaIdentityStore + private readonly addressResolver?: SignalAddressResolver private readonly defaultTimeoutMs: number private readonly hostDomain: string private readonly syncDedup = new PromiseDedup() @@ -41,6 +44,7 @@ export class SignalIdentitySyncApi { this.logger = options.logger this.query = options.query this.identityStore = options.identityStore + this.addressResolver = options.addressResolver this.defaultTimeoutMs = options.defaultTimeoutMs ?? WA_DEFAULTS.SIGNAL_FETCH_KEY_BUNDLES_TIMEOUT_MS this.hostDomain = options.hostDomain ?? WA_DEFAULTS.HOST_DOMAIN @@ -120,18 +124,31 @@ export class SignalIdentitySyncApi { const entries = this.parseIdentitySyncResponse(response, normalizedTargets) const { identityStore } = this if (identityStore && entries.length > 0) { - const identities = new Array<{ - readonly address: ReturnType - readonly identityKey: Uint8Array - }>(entries.length) + const addresses = new Array>( + entries.length + ) + for (let index = 0; index < entries.length; index += 1) { + addresses[index] = parseSignalAddressFromJid(entries[index].jid) + } + const resolvedAddresses = this.addressResolver + ? await this.addressResolver.resolveMany(addresses) + : addresses + const identitiesByAddress = new Map< + string, + { + readonly address: ReturnType + readonly identityKey: Uint8Array + } + >() for (let index = 0; index < entries.length; index += 1) { const entry = entries[index] - identities[index] = { - address: parseSignalAddressFromJid(entry.jid), + const address = resolvedAddresses[index] + identitiesByAddress.set(signalAddressKey(address), { + address, identityKey: toSerializedPubKey(entry.identity) - } + }) } - await identityStore.setRemoteIdentities(identities) + await identityStore.setRemoteIdentities(Array.from(identitiesByAddress.values())) } this.logger.debug('signal identity sync success', { requested: normalizedTargets.length, diff --git a/src/signal/api/__tests__/api.test.ts b/src/signal/api/__tests__/api.test.ts index 3a90d076..de97cb2e 100644 --- a/src/signal/api/__tests__/api.test.ts +++ b/src/signal/api/__tests__/api.test.ts @@ -24,8 +24,10 @@ import { generateRegistrationInfo, generateSignedPreKey } from '@signal/registration/keygen' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import { WaDeviceListMemoryStore } from '@store/memory/device-list.store' import { WaIdentityMemoryStore } from '@store/memory/identity.store' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { WaPreKeyMemoryStore } from '@store/memory/pre-key.store' import { WaSignalMemoryStore } from '@store/memory/signal.store' import type { BinaryNode } from '@transport/types' @@ -1115,11 +1117,14 @@ test('signal device sync api maps hosted.lid user response to requested lid user ]) }) -test('signal identity sync api parses result list and stores remote identities', async () => { +test('signal identity sync api stores remote identities under the canonical LID address', async () => { const identityStore = new WaIdentityMemoryStore() + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + await addressResolver.learnMessageJidPair('5511999999999@s.whatsapp.net', '778899@lid') const api = new SignalIdentitySyncApi({ logger: createNoopLogger(), identityStore, + addressResolver, query: async () => iqResult([ { @@ -1166,13 +1171,73 @@ test('signal identity sync api parses result list and stores remote identities', assert.equal(result[0].identity.length, 32) assert.equal(result[0].type, SIGNAL_KEY_BUNDLE_TYPE_BYTES[0]) - const persisted = await identityStore.getRemoteIdentity( + const pnIdentity = await identityStore.getRemoteIdentity( parseSignalAddressFromJid('5511999999999:1@s.whatsapp.net') ) + assert.equal(pnIdentity, null) + const persisted = await identityStore.getRemoteIdentity( + parseSignalAddressFromJid('778899:1@lid') + ) assert.ok(persisted) assert.equal(persisted.length, 33) }) +test('signal identity sync api deduplicates PN/LID aliases before batch persistence', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + await addressResolver.learnMessageJidPair('5511999999999@s.whatsapp.net', '778899@lid') + let persisted: readonly { + readonly address: ReturnType + readonly identityKey: Uint8Array + }[] = [] + const api = new SignalIdentitySyncApi({ + logger: createNoopLogger(), + identityStore: { + setRemoteIdentities: async (entries: typeof persisted) => { + persisted = entries + } + } as never, + addressResolver, + query: async () => + iqResult([ + { + tag: WA_NODE_TAGS.LIST, + attrs: {}, + content: [ + { + tag: WA_NODE_TAGS.USER, + attrs: { jid: '5511999999999:1@s.whatsapp.net' }, + content: [ + { + tag: WA_NODE_TAGS.IDENTITY, + attrs: {}, + content: makeBytes(SIGNAL_KEY_DATA_LENGTH, 1) + } + ] + }, + { + tag: WA_NODE_TAGS.USER, + attrs: { jid: '778899:1@lid' }, + content: [ + { + tag: WA_NODE_TAGS.IDENTITY, + attrs: {}, + content: makeBytes(SIGNAL_KEY_DATA_LENGTH, 2) + } + ] + } + ] + } + ]) + }) + + const result = await api.syncIdentityKeys(['5511999999999:1@s.whatsapp.net', '778899:1@lid']) + + assert.equal(result.length, 2) + assert.equal(persisted.length, 1) + assert.deepEqual(persisted[0].address, { user: '778899', server: 'lid', device: 1 }) + assert.deepEqual(persisted[0].identityKey.subarray(1), makeBytes(SIGNAL_KEY_DATA_LENGTH, 2)) +}) + test('signal identity sync api maps hosted.lid response jid to requested lid jid', async () => { const api = new SignalIdentitySyncApi({ logger: createNoopLogger(), diff --git a/src/signal/group/SenderKeyManager.ts b/src/signal/group/SenderKeyManager.ts index 4134f914..4a5eeb96 100644 --- a/src/signal/group/SenderKeyManager.ts +++ b/src/signal/group/SenderKeyManager.ts @@ -17,7 +17,8 @@ import { SIGNAL_SIGNATURE_LENGTH } from '@signal/api/constants' import { SIGNAL_GROUP_VERSION } from '@signal/constants' import { deriveSenderKeyMsgKey, selectMessageKey } from '@signal/group/SenderKeyChain' import { parseDistributionPayload, parseSenderKeyMessage } from '@signal/group/SenderKeyCodec' -import type { SenderKeyRecord, SignalAddress } from '@signal/types' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' +import type { SenderKeyDistributionRecord, SenderKeyRecord, SignalAddress } from '@signal/types' import type { WaSenderKeyStore } from '@store/contracts/sender-key.store' import { concatBytes } from '@util/bytes' @@ -63,17 +64,20 @@ export class SenderKeyManager { private readonly senderLock = new StoreLock() private readonly getFutureMessagesMax: (() => number) | undefined private readonly skipSignatureVerification: boolean + private readonly addressResolver: SignalAddressResolver | undefined public constructor( store: WaSenderKeyStore, options?: { readonly getFutureMessagesMax?: () => number readonly skipSignatureVerification?: boolean + readonly addressResolver?: SignalAddressResolver } ) { this.store = store this.getFutureMessagesMax = options?.getFutureMessagesMax this.skipSignatureVerification = options?.skipSignatureVerification === true + this.addressResolver = options?.addressResolver } /** @@ -89,6 +93,7 @@ export class SenderKeyManager { readonly ciphertext: GroupSenderKeyCiphertext readonly keyId: number }> { + sender = await this.resolveAddress(sender) return this.runWithSenderLock(groupId, sender, async () => { const senderKey = await this.ensureSenderKeyInternal(groupId, sender) if (!senderKey.signingPrivateKey) { @@ -158,13 +163,17 @@ export class SenderKeyManager { if (participants.length === 0) { return [] } - const distributed = await this.store.getDeviceSenderKeyDistributions(groupId, participants) + const resolvedParticipants = await this.resolveAddresses(participants) + const distributed = await this.store.getDeviceSenderKeyDistributions( + groupId, + resolvedParticipants + ) const pendingParticipants = new Array(participants.length) let pendingCount = 0 for (let index = 0; index < participants.length; index += 1) { const record = distributed[index] if (!record || record.keyId !== senderKeyId) { - pendingParticipants[pendingCount] = participants[index] + pendingParticipants[pendingCount] = resolvedParticipants[index] pendingCount += 1 } } @@ -181,17 +190,19 @@ export class SenderKeyManager { if (participants.length === 0) { return } + const resolvedParticipants = await this.resolveAddresses(participants) const timestampMs = Date.now() - const distributions = new Array(participants.length) - for (let index = 0; index < participants.length; index += 1) { - distributions[index] = { + const distributionsBySender = new Map() + for (let index = 0; index < resolvedParticipants.length; index += 1) { + const sender = resolvedParticipants[index] + distributionsBySender.set(signalAddressKey(sender), { groupId, - sender: participants[index], + sender, keyId: senderKeyId, timestampMs - } + }) } - await this.store.upsertSenderKeyDistributions(distributions) + await this.store.upsertSenderKeyDistributions(Array.from(distributionsBySender.values())) } /** @@ -204,6 +215,7 @@ export class SenderKeyManager { sender: SignalAddress, payload: Uint8Array ): Promise { + sender = await this.resolveAddress(sender) return this.runWithSenderLock(groupId, sender, async () => { if (groupId.length === 0) { throw new Error('sender key distribution missing groupId') @@ -234,10 +246,16 @@ export class SenderKeyManager { /** Decrypts an incoming sender-key group ciphertext into plaintext. */ public async decryptGroupMessage(payload: GroupSenderKeyCiphertext): Promise { - return this.runWithSenderLock(payload.groupId, payload.sender, async () => { + const originalSender = payload.sender + const sender = await this.resolveAddress(originalSender) + return this.runWithSenderLock(payload.groupId, sender, async () => { const parsed = parseSenderKeyMessage(payload.ciphertext) - const senderKey = await this.store.getDeviceSenderKey(payload.groupId, payload.sender) + const senderKey = await this.getSenderKeyWithLegacyFallback( + payload.groupId, + originalSender, + sender + ) if (!senderKey) { throw new Error('missing sender key') } @@ -318,6 +336,39 @@ export class SenderKeyManager { return created } + private async getSenderKeyWithLegacyFallback( + groupId: string, + originalSender: SignalAddress, + sender: SignalAddress + ): Promise { + const current = await this.store.getDeviceSenderKey(groupId, sender) + if (current || !this.addressResolver) return current + + const senderAddressKey = signalAddressKey(sender) + const originalSenderAddressKey = signalAddressKey(originalSender) + const legacySender = + originalSenderAddressKey === senderAddressKey + ? await this.addressResolver.resolvePhoneNumberAlias(sender) + : originalSender + if (!legacySender || signalAddressKey(legacySender) === senderAddressKey) { + return null + } + + const legacyKey = await this.store.getDeviceSenderKey(groupId, legacySender) + if (!legacyKey) return null + + const migratedKey = { ...legacyKey, sender } + const [legacyDistribution] = await this.store.getDeviceSenderKeyDistributions(groupId, [ + legacySender + ]) + await this.store.upsertSenderKey(migratedKey) + if (legacyDistribution) { + await this.store.upsertSenderKeyDistribution({ ...legacyDistribution, sender }) + } + await this.store.deleteDeviceSenderKey(legacySender, groupId) + return migratedKey + } + private runWithSenderLock( groupId: string, sender: SignalAddress, @@ -325,4 +376,14 @@ export class SenderKeyManager { ): Promise { return this.senderLock.run(`senderKey:${groupId}:${signalAddressKey(sender)}`, task) } + + private resolveAddress(address: SignalAddress): SignalAddress | Promise { + return this.addressResolver ? this.addressResolver.resolve(address) : address + } + + private resolveAddresses( + addresses: readonly SignalAddress[] + ): readonly SignalAddress[] | Promise { + return this.addressResolver ? this.addressResolver.resolveMany(addresses) : addresses + } } diff --git a/src/signal/group/__tests__/sender-key.test.ts b/src/signal/group/__tests__/sender-key.test.ts index 65dd9514..8a4d20a9 100644 --- a/src/signal/group/__tests__/sender-key.test.ts +++ b/src/signal/group/__tests__/sender-key.test.ts @@ -8,10 +8,23 @@ import { SIGNAL_GROUP_VERSION } from '@signal/constants' import { deriveSenderKeyMsgKey, selectMessageKey } from '@signal/group/SenderKeyChain' import { parseDistributionPayload, parseSenderKeyMessage } from '@signal/group/SenderKeyCodec' import { SenderKeyManager } from '@signal/group/SenderKeyManager' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SenderKeyRecord, SignalAddress } from '@signal/types' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { SenderKeyMemoryStore } from '@store/memory/sender-key.store' import { concatBytes } from '@util/bytes' +class CapturingSenderKeyStore extends SenderKeyMemoryStore { + public distributionBatchSizes: number[] = [] + + public override async upsertSenderKeyDistributions( + records: Parameters[0] + ): Promise { + this.distributionBatchSizes.push(records.length) + await super.upsertSenderKeyDistributions(records) + } +} + function makeBytes(length: number, seed = 0): Uint8Array { const out = new Uint8Array(length) for (let index = 0; index < out.length; index += 1) { @@ -218,3 +231,95 @@ test('sender key manager bypasses signature check when skipSignatureVerification }) assert.deepEqual(decrypted, plaintext) }) + +test('sender key manager shares one sender-key chain across mapped PN and LID addresses', async () => { + const senderManager = new SenderKeyManager(new SenderKeyMemoryStore()) + const receiverStore = new SenderKeyMemoryStore() + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + const receiverManager = new SenderKeyManager(receiverStore, { addressResolver }) + const groupId = '120363000000000001@g.us' + const pn = makeAddress('5511666666666', 3) + const lid = { user: '445566', server: 'lid', device: 3 } as const + + await addressResolver.learnMessageJidPair('5511666666666:3@s.whatsapp.net', '445566:3@lid') + const first = await senderManager.prepareGroupEncryption(groupId, lid, makeBytes(31, 21)) + await receiverManager.processSenderKeyDistributionPayload( + groupId, + pn, + first.distributionMessage.axolotlSenderKeyDistributionMessage! + ) + + assert.equal(await receiverStore.getDeviceSenderKey(groupId, pn), null) + assert.ok(await receiverStore.getDeviceSenderKey(groupId, lid)) + assert.deepEqual( + await receiverManager.decryptGroupMessage({ + groupId, + sender: lid, + ciphertext: first.ciphertext.ciphertext + }), + makeBytes(31, 21) + ) + + const second = await senderManager.prepareGroupEncryption(groupId, lid, makeBytes(33, 44)) + assert.deepEqual( + await receiverManager.decryptGroupMessage({ + groupId, + sender: pn, + ciphertext: second.ciphertext.ciphertext + }), + makeBytes(33, 44) + ) +}) + +test('sender key manager deduplicates mapped PN/LID distribution recipients', async () => { + const store = new CapturingSenderKeyStore() + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + const manager = new SenderKeyManager(store, { addressResolver }) + const groupId = '120363000000000002@g.us' + const pn = makeAddress('5511555555555', 2) + const lid = { user: '112233', server: 'lid', device: 2 } as const + await addressResolver.learnMessageJidPair('5511555555555@s.whatsapp.net', '112233@lid') + + await manager.markSenderKeyDistributed(groupId, 42, [pn, lid]) + + assert.deepEqual(store.distributionBatchSizes, [1]) + const [distribution] = await store.getDeviceSenderKeyDistributions(groupId, [lid]) + assert.ok(distribution) + assert.equal(distribution.groupId, groupId) + assert.deepEqual(distribution.sender, lid) + assert.equal(distribution.keyId, 42) +}) + +test('sender key manager migrates legacy PN state when a sender switches to LID', async () => { + const senderStore = new SenderKeyMemoryStore() + const receiverStore = new SenderKeyMemoryStore() + const senderManager = new SenderKeyManager(senderStore) + const legacyReceiverManager = new SenderKeyManager(receiverStore) + const groupId = '120363000000000003@g.us' + const pn = makeAddress('5511444444444', 3) + const lid = { user: '998877', server: 'lid', device: 3 } as const + const plaintext = makeBytes(29, 61) + const prepared = await senderManager.prepareGroupEncryption(groupId, pn, plaintext) + await legacyReceiverManager.processSenderKeyDistributionPayload( + groupId, + pn, + prepared.distributionMessage.axolotlSenderKeyDistributionMessage! + ) + + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + await addressResolver.learnMessageJidPair('5511444444444@s.whatsapp.net', '998877@lid') + const upgradedReceiverManager = new SenderKeyManager(receiverStore, { addressResolver }) + + assert.deepEqual( + await upgradedReceiverManager.decryptGroupMessage({ + groupId, + sender: lid, + ciphertext: prepared.ciphertext.ciphertext + }), + plaintext + ) + assert.equal(await receiverStore.getDeviceSenderKey(groupId, pn), null) + assert.ok(await receiverStore.getDeviceSenderKey(groupId, lid)) + assert.deepEqual(await receiverStore.getDeviceSenderKeyDistributions(groupId, [pn]), [null]) + assert.ok((await receiverStore.getDeviceSenderKeyDistributions(groupId, [lid]))[0]) +}) diff --git a/src/signal/index.ts b/src/signal/index.ts index 89d8b878..b5ea662a 100644 --- a/src/signal/index.ts +++ b/src/signal/index.ts @@ -53,6 +53,7 @@ export { SignalSessionSyncApi } from '@signal/api/SignalSessionSyncApi' export { SenderKeyManager } from '@signal/group/SenderKeyManager' export { createAndStoreInitialKeys } from '@signal/registration/utils' export { SignalProtocol } from '@signal/session/SignalProtocol' +export { SignalAddressResolver } from '@signal/session/SignalAddressResolver' export { createSignalSessionResolver, type SignalSessionResolver } from '@signal/session/resolver' export { ADV_PREFIX_ACCOUNT_KEY_INDEX, diff --git a/src/signal/session/SignalAddressResolver.ts b/src/signal/session/SignalAddressResolver.ts new file mode 100644 index 00000000..864fb90e --- /dev/null +++ b/src/signal/session/SignalAddressResolver.ts @@ -0,0 +1,237 @@ +import { SharedExclusiveGate } from '@infra/perf/SharedExclusiveGate' +import { StoreLock } from '@infra/perf/StoreLock' +import { WA_DEFAULTS } from '@protocol/constants' +import { canonicalizeSignalServer, parseSignalAddressFromJid } from '@protocol/jid' +import type { SignalAddress } from '@signal/types' +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' +import { resolvePositive } from '@util/coercion' +import { setBoundedMapEntry } from '@util/collections' + +const DEFAULT_MAX_CACHE_ENTRIES = 8_192 + +interface LidPnUsers { + readonly pnUser: string + readonly lidUser: string +} + +function addressKind(address: SignalAddress): 'pn' | 'lid' | null { + const server = canonicalizeSignalServer(address.server ?? WA_DEFAULTS.HOST_DOMAIN) + if (server === WA_DEFAULTS.HOST_DOMAIN) return 'pn' + if (server === WA_DEFAULTS.LID_SERVER) return 'lid' + return null +} + +function parseLidPnUsers(firstJid: string, secondJid: string): LidPnUsers | null { + const first = parseSignalAddressFromJid(firstJid) + const second = parseSignalAddressFromJid(secondJid) + const firstKind = addressKind(first) + const secondKind = addressKind(second) + if (firstKind === 'pn' && secondKind === 'lid') { + return { pnUser: first.user, lidUser: second.user } + } + if (firstKind === 'lid' && secondKind === 'pn') { + return { pnUser: second.user, lidUser: first.user } + } + return null +} + +/** + * Resolves PN Signal addresses through the persistent PN/LID mapping learned + * from authenticated stanza metadata. Positive and negative lookups are kept + * in-process, so the steady-state path is one Map lookup with no store I/O. + */ +export class SignalAddressResolver { + private readonly store: WaLidPnMappingStore + private readonly cache = new Map() + private readonly lookupLock = new StoreLock() + private readonly mappingGate = new SharedExclusiveGate() + private readonly maxCacheEntries: number + + public constructor( + store: WaLidPnMappingStore, + options: { readonly maxCacheEntries?: number } = {} + ) { + this.store = store + this.maxCacheEntries = resolvePositive( + options.maxCacheEntries, + DEFAULT_MAX_CACHE_ENTRIES, + 'SignalAddressResolver.maxCacheEntries' + ) + } + + /** + * Learns a conservative mapping from ordinary message metadata. The target + * LID is not replaced when another PN already owns it. + */ + public learnMessageJidPair(firstJid: string, secondJid: string): Promise { + return this.learnJidPairInternal(firstJid, secondJid, false) + } + + /** Learns an authoritative cross-reference carried for a peer recipient. */ + public learnPeerRecipientJidPair(firstJid: string, secondJid: string): Promise { + return this.learnJidPairInternal(firstJid, secondJid, true) + } + + private async learnJidPairInternal( + firstJid: string, + secondJid: string, + replaceExisting: boolean + ): Promise { + const mapping = parseLidPnUsers(firstJid, secondJid) + if (!mapping) return false + // Ordinary message metadata is conservative, so a matching hot-cache + // entry cannot authorize a replacement. Authoritative peer metadata + // must still verify the reverse owner: another resolver may have + // replaced the persisted one-to-one mapping since this entry was cached. + if (!replaceExisting && this.cache.get(mapping.pnUser) === mapping.lidUser) return false + // A replacement can evict both the previous LID for this PN and the + // previous PN for this LID. Serialize the rare mutations globally so + // every cache invalidation reflects the same one-to-one store state. + return this.mappingGate.runExclusive(async () => { + let current: string | null + if (this.cache.has(mapping.pnUser)) { + current = this.cache.get(mapping.pnUser) ?? null + } else { + current = await this.store.getLidUser(mapping.pnUser) + } + const currentPn = await this.store.getPnUser(mapping.lidUser) + if (current === mapping.lidUser && currentPn === mapping.pnUser) { + this.cacheMapping(mapping.pnUser, mapping.lidUser) + return false + } + if (!replaceExisting && currentPn !== null) { + this.cacheMapping(mapping.pnUser, current) + return false + } + await this.store.setLidUser(mapping.pnUser, mapping.lidUser) + if (currentPn && currentPn !== mapping.pnUser) { + this.cacheMapping(currentPn, null) + } + this.cacheMapping(mapping.pnUser, mapping.lidUser) + return true + }) + } + + /** Resolves a PN address to its current LID while preserving the device id. */ + public resolve(address: SignalAddress): SignalAddress | Promise { + if (addressKind(address) !== 'pn') return address + const cached = this.cache.get(address.user) + if (cached !== undefined) return this.applyMapping(address, cached) + return this.resolveUncached(address) + } + + /** Batch variant of {@link resolve}; returns the input array when nothing changes. */ + public resolveMany( + addresses: readonly SignalAddress[] + ): readonly SignalAddress[] | Promise { + if (addresses.length === 0) return addresses + let resolved: SignalAddress[] | null = null + let missingUsers: Set | null = null + for (let index = 0; index < addresses.length; index += 1) { + const address = addresses[index] + if (addressKind(address) !== 'pn') continue + const cached = this.cache.get(address.user) + if (cached === undefined) { + if (!missingUsers) missingUsers = new Set() + missingUsers.add(address.user) + continue + } + if (!cached) continue + resolved ??= addresses.slice() + resolved[index] = this.applyMapping(address, cached) + } + if (!missingUsers) return resolved ?? addresses + return this.resolveManyUncached(addresses, resolved, missingUsers) + } + + /** Resolves a LID back to its PN alias for one-time legacy state migration. */ + public resolvePhoneNumberAlias( + address: SignalAddress + ): SignalAddress | null | Promise { + if (addressKind(address) !== 'lid') return null + return this.resolvePhoneNumberAliasUncached(address) + } + + /** Clears both the persistent mapping and its in-process lookup cache. */ + public async clear(): Promise { + await this.mappingGate.runExclusive(async () => { + try { + await this.store.clear() + } finally { + this.cache.clear() + } + }) + } + + private async resolveUncached(address: SignalAddress): Promise { + const lidUser = await this.resolveLidUser(address.user) + return this.applyMapping(address, lidUser) + } + + private async resolveManyUncached( + addresses: readonly SignalAddress[], + resolved: SignalAddress[] | null, + missingUsers: ReadonlySet + ): Promise { + const resolvedLidUsers = new Map() + await Promise.all( + Array.from(missingUsers, async (pnUser) => { + resolvedLidUsers.set(pnUser, await this.resolveLidUser(pnUser)) + }) + ) + for (let index = 0; index < addresses.length; index += 1) { + const address = addresses[index] + if (addressKind(address) !== 'pn' || !missingUsers.has(address.user)) continue + const lidUser = resolvedLidUsers.get(address.user) + if (!lidUser) continue + resolved ??= addresses.slice() + resolved[index] = this.applyMapping(address, lidUser) + } + return resolved ?? addresses + } + + private async resolvePhoneNumberAliasUncached( + address: SignalAddress + ): Promise { + const pnUser = await this.lookupLock.run(`lid:${address.user}`, () => + this.mappingGate.runShared(() => this.store.getPnUser(address.user)) + ) + if (!pnUser) return null + return { + user: pnUser, + server: + address.server === WA_DEFAULTS.HOSTED_LID_SERVER + ? WA_DEFAULTS.HOSTED_SERVER + : WA_DEFAULTS.HOST_DOMAIN, + device: address.device + } + } + + private applyMapping(address: SignalAddress, lidUser: string | null): SignalAddress { + if (!lidUser) return address + return { + user: lidUser, + server: + address.server === WA_DEFAULTS.HOSTED_SERVER + ? WA_DEFAULTS.HOSTED_LID_SERVER + : WA_DEFAULTS.LID_SERVER, + device: address.device + } + } + + private async resolveLidUser(pnUser: string): Promise { + if (this.cache.has(pnUser)) return this.cache.get(pnUser) ?? null + return this.lookupLock.run(pnUser, () => + this.mappingGate.runShared(async () => { + if (this.cache.has(pnUser)) return this.cache.get(pnUser) ?? null + const lidUser = await this.store.getLidUser(pnUser) + this.cacheMapping(pnUser, lidUser) + return lidUser + }) + ) + } + + private cacheMapping(pnUser: string, lidUser: string | null): void { + setBoundedMapEntry(this.cache, pnUser, lidUser, this.maxCacheEntries) + } +} diff --git a/src/signal/session/SignalProtocol.ts b/src/signal/session/SignalProtocol.ts index 7694d0c2..a99477b0 100644 --- a/src/signal/session/SignalProtocol.ts +++ b/src/signal/session/SignalProtocol.ts @@ -5,6 +5,7 @@ import { StoreLock } from '@infra/perf/StoreLock' import { signalAddressKey } from '@protocol/jid' import { MAX_PREV_SESSIONS } from '@signal/constants' import { encodeSignalSessionSnapshot } from '@signal/session/encoding' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import { decryptMsg, decryptMsgFromSession, @@ -46,6 +47,23 @@ interface EstablishOutgoingSessionOptions { readonly knownAbsent?: boolean } +interface SignalEncryptRequest { + readonly address: SignalAddress + readonly plaintext: Uint8Array + readonly expectedIdentity?: Uint8Array +} + +interface SignalPrefetchedSession { + readonly address: SignalAddress + readonly session: SignalSessionRecord +} + +interface SignalEncryptResult { + readonly type: 'msg' | 'pkmsg' + readonly ciphertext: Uint8Array + readonly baseKey: Uint8Array | null +} + export interface SignalProtocolStores { readonly signal: WaSignalStore readonly preKey: WaPreKeyStore @@ -62,11 +80,17 @@ export class SignalProtocol { private readonly stores: SignalProtocolStores private readonly logger: Logger private readonly sessionMutationLock: StoreLock + private readonly addressResolver: SignalAddressResolver | undefined - public constructor(stores: SignalProtocolStores, logger: Logger = new ConsoleLogger('info')) { + public constructor( + stores: SignalProtocolStores, + logger: Logger = new ConsoleLogger('info'), + addressResolver?: SignalAddressResolver + ) { this.stores = stores this.logger = logger this.sessionMutationLock = new StoreLock() + this.addressResolver = addressResolver } /** @@ -81,6 +105,7 @@ export class SignalProtocol { remoteBundle: SignalPreKeyBundle, options: EstablishOutgoingSessionOptions = {} ): Promise { + address = await this.resolveAddress(address) return this.runWithAddressLock(address, async () => { if (options.reuseExisting && !options.knownAbsent) { const existing = await this.stores.session.getSession(address) @@ -119,6 +144,7 @@ export class SignalProtocol { readonly remoteIdentity: Uint8Array readonly reusedExisting: boolean }> { + address = await this.resolveAddress(address) return this.runWithAddressLock(address, async () => { if (options.reuseExisting && !options.knownAbsent) { const existing = await this.stores.session.getSession(address) @@ -171,6 +197,24 @@ export class SignalProtocol { }> }> { if (entries.length === 0) return { resolved: [], skipped: [] } + entries = await this.resolveAddressEntries(entries) + const entryIndexByAddress = new Map() + let uniqueEntries: Array<(typeof entries)[number]> | null = null + for (let index = 0; index < entries.length; index += 1) { + const entry = entries[index] + const key = signalAddressKey(entry.address) + const previousIndex = entryIndexByAddress.get(key) + if (previousIndex !== undefined) { + if (!uint8Equal(entries[previousIndex].remoteIdentity, entry.remoteIdentity)) { + throw new Error('identity mismatch') + } + uniqueEntries ??= entries.slice(0, index) + continue + } + entryIndexByAddress.set(key, index) + if (uniqueEntries) uniqueEntries.push(entry) + } + if (uniqueEntries) entries = uniqueEntries const lockKeys = new Array(entries.length) for (let i = 0; i < entries.length; i += 1) { lockKeys[i] = signalAddressLockKey(entries[i].address) @@ -219,12 +263,9 @@ export class SignalProtocol { address: SignalAddress, plaintext: Uint8Array, expectedIdentity?: Uint8Array - ): Promise<{ - readonly type: 'msg' | 'pkmsg' - readonly ciphertext: Uint8Array - readonly baseKey: Uint8Array | null - }> { - const [encrypted] = await this.encryptMessagesBatch([ + ): Promise { + address = await this.resolveAddress(address) + const [encrypted] = await this.encryptMessagesBatchResolved([ { address, plaintext, expectedIdentity } ]) return encrypted @@ -232,25 +273,23 @@ export class SignalProtocol { /** Batch variant of {@link encryptMessage} that shares per-address locks. */ public async encryptMessagesBatch( - requests: readonly { - readonly address: SignalAddress - readonly plaintext: Uint8Array - readonly expectedIdentity?: Uint8Array - }[], - prefetchedSessions?: readonly { - readonly address: SignalAddress - readonly session: SignalSessionRecord - }[] - ): Promise< - readonly { - readonly type: 'msg' | 'pkmsg' - readonly ciphertext: Uint8Array - readonly baseKey: Uint8Array | null - }[] - > { + requests: readonly SignalEncryptRequest[], + prefetchedSessions?: readonly SignalPrefetchedSession[] + ): Promise { if (requests.length === 0) { return [] } + requests = await this.resolveAddressEntries(requests) + if (prefetchedSessions && prefetchedSessions.length > 0) { + prefetchedSessions = await this.resolveAddressEntries(prefetchedSessions) + } + return this.encryptMessagesBatchResolved(requests, prefetchedSessions) + } + + private async encryptMessagesBatchResolved( + requests: readonly SignalEncryptRequest[], + prefetchedSessions?: readonly SignalPrefetchedSession[] + ): Promise { const lockKeySet = new Set() for (let i = 0; i < requests.length; i += 1) lockKeySet.add(signalAddressLockKey(requests[i].address)) @@ -309,11 +348,7 @@ export class SignalProtocol { string, { readonly address: SignalAddress; readonly identityKey: Uint8Array } >() - const results = new Array<{ - readonly type: 'msg' | 'pkmsg' - readonly ciphertext: Uint8Array - readonly baseKey: Uint8Array | null - }>(requests.length) + const results = new Array(requests.length) for (let index = 0; index < requests.length; index += 1) { const request = requests[index] @@ -386,6 +421,7 @@ export class SignalProtocol { readonly ciphertext: Uint8Array } ): Promise { + address = await this.resolveAddress(address) return this.runWithAddressLock(address, async () => { const currentSession = await this.stores.session.getSession(address) @@ -424,6 +460,36 @@ export class SignalProtocol { return this.sessionMutationLock.run(signalAddressLockKey(address), task) } + private resolveAddress(address: SignalAddress): SignalAddress | Promise { + return this.addressResolver ? this.addressResolver.resolve(address) : address + } + + private async resolveAddressEntries( + entries: readonly T[] + ): Promise { + if (!this.addressResolver || entries.length === 0) return entries + const addresses = new Array(entries.length) + for (let index = 0; index < entries.length; index += 1) { + addresses[index] = entries[index].address + } + const resolvedAddresses = await this.addressResolver.resolveMany(addresses) + if (resolvedAddresses === addresses) return entries + + let resolvedEntries: T[] | null = null + for (let index = 0; index < entries.length; index += 1) { + if (resolvedAddresses[index] === addresses[index]) { + if (resolvedEntries) resolvedEntries.push(entries[index]) + continue + } + resolvedEntries ??= entries.slice(0, index) + resolvedEntries.push({ + ...entries[index], + address: resolvedAddresses[index] + }) + } + return resolvedEntries ?? entries + } + private async decryptPkMsg( currentSession: SignalSessionRecord | null, parsed: ParsedPreKeySignalMessage diff --git a/src/signal/session/__tests__/address-resolver.test.ts b/src/signal/session/__tests__/address-resolver.test.ts new file mode 100644 index 00000000..fe6915fd --- /dev/null +++ b/src/signal/session/__tests__/address-resolver.test.ts @@ -0,0 +1,297 @@ +import assert from 'node:assert/strict' +import test from 'node:test' + +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' +import type { SignalAddress } from '@signal/types' +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' + +class CountingMappingStore extends WaLidPnMappingMemoryStore { + public reads = 0 + public writes = 0 + + public override async getLidUser(pnUser: string): Promise { + this.reads += 1 + return super.getLidUser(pnUser) + } + + public override async setLidUser(pnUser: string, lidUser: string): Promise { + this.writes += 1 + await super.setLidUser(pnUser, lidUser) + } +} + +class DelayedFirstWriteStore implements WaLidPnMappingStore { + private readonly delegate = new WaLidPnMappingMemoryStore() + private firstWrite = true + private readonly startedPromise: Promise + private resolveStarted: (() => void) | null = null + private readonly releasePromise: Promise + private resolveRelease: (() => void) | null = null + public writesStarted = 0 + + public constructor() { + this.startedPromise = new Promise((resolve) => { + this.resolveStarted = resolve + }) + this.releasePromise = new Promise((resolve) => { + this.resolveRelease = resolve + }) + } + + public getLidUser(pnUser: string): Promise { + return this.delegate.getLidUser(pnUser) + } + + public getPnUser(lidUser: string): Promise { + return this.delegate.getPnUser(lidUser) + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + this.writesStarted += 1 + if (this.firstWrite) { + this.firstWrite = false + this.resolveStarted?.() + this.resolveStarted = null + await this.releasePromise + } + await this.delegate.setLidUser(pnUser, lidUser) + } + + public seed(pnUser: string, lidUser: string): Promise { + return this.delegate.setLidUser(pnUser, lidUser) + } + + public clear(): Promise { + return this.delegate.clear() + } + + public waitStarted(): Promise { + return this.startedPromise + } + + public release(): void { + this.resolveRelease?.() + this.resolveRelease = null + } +} + +function pnAddress(device = 0): SignalAddress { + return { user: '5511999999999', server: 's.whatsapp.net', device } +} + +test('SignalAddressResolver caches misses and replaces them when a mapping is learned', async () => { + const store = new CountingMappingStore() + const resolver = new SignalAddressResolver(store) + const pn = pnAddress(3) + + assert.deepEqual(await resolver.resolve(pn), pn) + assert.deepEqual(await resolver.resolve(pn), pn) + assert.equal(store.reads, 1) + + assert.equal( + await resolver.learnMessageJidPair('5511999999999@s.whatsapp.net', '778899@lid'), + true + ) + assert.deepEqual(await resolver.resolve(pn), { + user: '778899', + server: 'lid', + device: 3 + }) + assert.equal(store.reads, 1) + assert.equal(store.writes, 1) + + assert.equal(await resolver.learnMessageJidPair('778899@lid', '5511999999999:3@hosted'), false) + assert.deepEqual( + await resolver.resolve({ user: '5511999999999', server: 'hosted', device: 99 }), + { user: '778899', server: 'hosted.lid', device: 99 } + ) + assert.equal(store.writes, 1) +}) + +test('SignalAddressResolver reloads a persisted mapping in a fresh instance', async () => { + const store = new WaLidPnMappingMemoryStore() + const first = new SignalAddressResolver(store) + await first.learnMessageJidPair('5511888888888@s.whatsapp.net', '112233@lid') + + const restarted = new SignalAddressResolver(store) + assert.deepEqual(await restarted.resolve(pnAddress()), pnAddress()) + assert.deepEqual( + await restarted.resolve({ user: '5511888888888', server: 's.whatsapp.net', device: 7 }), + { user: '112233', server: 'lid', device: 7 } + ) +}) + +test('SignalAddressResolver serializes concurrent remaps for the same PN', async () => { + const store = new DelayedFirstWriteStore() + const resolver = new SignalAddressResolver(store) + + const first = resolver.learnMessageJidPair('5511777777777@s.whatsapp.net', '111@lid') + await store.waitStarted() + const second = resolver.learnMessageJidPair('5511777777777@s.whatsapp.net', '222@lid') + store.release() + + assert.deepEqual(await Promise.all([first, second]), [true, true]) + assert.deepEqual( + await resolver.resolve({ user: '5511777777777', server: 'hosted', device: 99 }), + { user: '222', server: 'hosted.lid', device: 99 } + ) +}) + +test('SignalAddressResolver serializes disjoint remaps that cross reverse entries', async () => { + const store = new DelayedFirstWriteStore() + await store.seed('5511000000001', '101') + await store.seed('5511000000002', '202') + const resolver = new SignalAddressResolver(store) + + const first = resolver.learnPeerRecipientJidPair('5511000000001@s.whatsapp.net', '202@lid') + await store.waitStarted() + const second = resolver.learnPeerRecipientJidPair('5511000000002@s.whatsapp.net', '101@lid') + + await Promise.resolve() + assert.equal(store.writesStarted, 1) + store.release() + + assert.deepEqual(await Promise.all([first, second]), [true, true]) + assert.equal(await store.getLidUser('5511000000001'), '202') + assert.equal(await store.getLidUser('5511000000002'), '101') + assert.equal(await store.getPnUser('101'), '5511000000002') + assert.equal(await store.getPnUser('202'), '5511000000001') +}) + +test('SignalAddressResolver rejects conflicting author metadata without stealing a LID', async () => { + const store = new WaLidPnMappingMemoryStore() + const resolver = new SignalAddressResolver(store) + await resolver.learnMessageJidPair('5511000000001@s.whatsapp.net', '101@lid') + + assert.equal( + await resolver.learnMessageJidPair('5511000000002@s.whatsapp.net', '101@lid'), + false + ) + assert.equal(await store.getPnUser('101'), '5511000000001') + assert.equal(await store.getLidUser('5511000000002'), null) +}) + +test('SignalAddressResolver replaces stale mappings from authoritative peer metadata', async () => { + const store = new WaLidPnMappingMemoryStore() + const resolver = new SignalAddressResolver(store) + await resolver.learnMessageJidPair('5511000000001@s.whatsapp.net', '101@lid') + + assert.equal( + await resolver.learnPeerRecipientJidPair('5511000000002@s.whatsapp.net', '101@lid'), + true + ) + assert.equal(await store.getLidUser('5511000000001'), null) + assert.equal(await store.getLidUser('5511000000002'), '101') + assert.equal(await store.getPnUser('101'), '5511000000002') +}) + +test('SignalAddressResolver reasserts an authoritative mapping after its cache becomes stale', async () => { + const store = new WaLidPnMappingMemoryStore() + const resolver = new SignalAddressResolver(store) + await resolver.learnMessageJidPair('5511000000001@s.whatsapp.net', '101@lid') + + // Simulate another resolver replacing the persisted owner while this + // resolver still has the original PN -> LID pair cached. + await store.setLidUser('5511000000002', '101') + + assert.equal( + await resolver.learnPeerRecipientJidPair('5511000000001@s.whatsapp.net', '101@lid'), + true + ) + assert.equal(await store.getLidUser('5511000000001'), '101') + assert.equal(await store.getLidUser('5511000000002'), null) + assert.equal(await store.getPnUser('101'), '5511000000001') +}) + +test('SignalAddressResolver ignores pairs that are not PN/LID alternates', async () => { + const store = new CountingMappingStore() + const resolver = new SignalAddressResolver(store) + + assert.equal( + await resolver.learnMessageJidPair('5511000000001@s.whatsapp.net', '5511000000002@hosted'), + false + ) + assert.equal(await resolver.learnMessageJidPair('111@lid', '222@hosted.lid'), false) + assert.equal(store.reads, 0) + assert.equal(store.writes, 0) +}) + +test('SignalAddressResolver bounds its positive and negative lookup cache', async () => { + const store = new CountingMappingStore() + await store.setLidUser('5511000000001', '101') + await store.setLidUser('5511000000002', '202') + store.reads = 0 + store.writes = 0 + const resolver = new SignalAddressResolver(store, { maxCacheEntries: 1 }) + + await resolver.resolve({ user: '5511000000001', server: 's.whatsapp.net', device: 0 }) + await resolver.resolve({ user: '5511000000002', server: 's.whatsapp.net', device: 0 }) + await resolver.resolve({ user: '5511000000001', server: 's.whatsapp.net', device: 0 }) + + assert.equal(store.reads, 3) +}) + +test('SignalAddressResolver resolveMany retains results larger than its cache', async () => { + const store = new CountingMappingStore() + await store.setLidUser('5511000000001', '101') + await store.setLidUser('5511000000002', '202') + await store.setLidUser('5511000000003', '303') + store.reads = 0 + store.writes = 0 + const resolver = new SignalAddressResolver(store, { maxCacheEntries: 1 }) + + const resolved = await resolver.resolveMany([ + { user: '5511000000001', server: 's.whatsapp.net', device: 1 }, + { user: '5511000000002', server: 's.whatsapp.net', device: 2 }, + { user: '5511000000003', server: 's.whatsapp.net', device: 3 } + ]) + + assert.deepEqual(resolved, [ + { user: '101', server: 'lid', device: 1 }, + { user: '202', server: 'lid', device: 2 }, + { user: '303', server: 'lid', device: 3 } + ]) + assert.equal(store.reads, 3) +}) + +test('SignalAddressResolver resolves hosted and regular LIDs back to their PN alias', async () => { + const store = new WaLidPnMappingMemoryStore() + const resolver = new SignalAddressResolver(store) + await resolver.learnMessageJidPair('5511000000004@s.whatsapp.net', '404@lid') + + assert.deepEqual( + await resolver.resolvePhoneNumberAlias({ user: '404', server: 'lid', device: 4 }), + { user: '5511000000004', server: 's.whatsapp.net', device: 4 } + ) + assert.deepEqual( + await resolver.resolvePhoneNumberAlias({ user: '404', server: 'hosted.lid', device: 99 }), + { user: '5511000000004', server: 'hosted', device: 99 } + ) +}) + +test('SignalAddressResolver clears its persistent mapping and hot cache together', async () => { + const store = new WaLidPnMappingMemoryStore() + const resolver = new SignalAddressResolver(store) + const pn = { user: '5511000000004', server: 's.whatsapp.net', device: 4 } as const + await resolver.learnMessageJidPair('5511000000004@s.whatsapp.net', '404@lid') + assert.equal((await resolver.resolve(pn)).user, '404') + + await resolver.clear() + + assert.deepEqual(await resolver.resolve(pn), pn) + assert.equal(await store.getLidUser('5511000000004'), null) +}) + +test('WaLidPnMappingMemoryStore evicts its oldest mapping at capacity', async () => { + const store = new WaLidPnMappingMemoryStore({ maxMappings: 2 }) + await store.setLidUser('5511000000001', '101') + await store.setLidUser('5511000000002', '202') + await store.setLidUser('5511000000003', '303') + + assert.equal(await store.getLidUser('5511000000001'), null) + assert.equal(await store.getLidUser('5511000000002'), '202') + assert.equal(await store.getLidUser('5511000000003'), '303') + assert.equal(await store.getPnUser('101'), null) + assert.equal(await store.getPnUser('202'), '5511000000002') +}) diff --git a/src/signal/session/__tests__/resolver.test.ts b/src/signal/session/__tests__/resolver.test.ts index 6fd1972e..77b704f6 100644 --- a/src/signal/session/__tests__/resolver.test.ts +++ b/src/signal/session/__tests__/resolver.test.ts @@ -3,7 +3,9 @@ import test from 'node:test' import { createNoopLogger } from '@infra/log/types' import { createSignalSessionResolver } from '@signal/session/resolver' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalPreKeyBundle } from '@signal/types' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { delay } from '@util/async' async function flushMicrotasks(turns = 3): Promise { @@ -28,6 +30,108 @@ function buildBundle(seed: number): SignalPreKeyBundle { } } +test('signal session resolver treats mapped PN and LID targets as one existing session', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + await addressResolver.learnMessageJidPair('5511999999999@s.whatsapp.net', '778899@lid') + + const existingSession = {} as never + const checkedAddresses: { readonly user: string; readonly server?: string }[] = [] + let singleFetchCalls = 0 + let batchFetchCalls = 0 + const sessionResolver = createSignalSessionResolver({ + signalProtocol: { + establishOutgoingSession: async () => { + throw new Error('must not establish a replacement session') + } + } as never, + sessionStore: { + hasSession: async (address: { readonly user: string; readonly server?: string }) => { + checkedAddresses.push(address) + return address.user === '778899' && address.server === 'lid' + }, + getSessionsBatch: async ( + addresses: readonly { readonly user: string; readonly server?: string }[] + ) => { + checkedAddresses.push(...addresses) + return addresses.map((address) => + address.user === '778899' && address.server === 'lid' ? existingSession : null + ) + } + } as never, + identityStore: { + getRemoteIdentity: async () => null + } as never, + signalIdentitySync: { + syncIdentityKeys: async () => undefined + } as never, + signalSessionSync: { + fetchKeyBundle: async () => { + singleFetchCalls += 1 + throw new Error('must not fetch a replacement key bundle') + }, + fetchKeyBundles: async () => { + batchFetchCalls += 1 + throw new Error('must not fetch replacement key bundles') + } + } as never, + addressResolver, + logger: createNoopLogger() + }) + + await sessionResolver.ensureSession( + { user: '5511999999999', server: 's.whatsapp.net', device: 2 }, + '5511999999999:2@s.whatsapp.net' + ) + const batch = await sessionResolver.ensureSessionsBatch([ + '5511999999999:2@s.whatsapp.net', + '778899:2@lid' + ]) + + assert.equal(singleFetchCalls, 0) + assert.equal(batchFetchCalls, 0) + assert.equal(checkedAddresses.length, 2) + assert.ok(checkedAddresses.every((address) => address.user === '778899')) + assert.deepEqual(batch, [ + { + jid: '5511999999999:2@s.whatsapp.net', + address: { user: '778899', server: 'lid', device: 2 }, + session: existingSession + } + ]) +}) + +test('signal session resolver rejects conflicting identities for PN/LID aliases', async () => { + const addressResolver = new SignalAddressResolver(new WaLidPnMappingMemoryStore()) + await addressResolver.learnMessageJidPair('5511999999999@s.whatsapp.net', '778899@lid') + let storeReads = 0 + const sessionResolver = createSignalSessionResolver({ + signalProtocol: {} as never, + sessionStore: { + getSessionsBatch: async () => { + storeReads += 1 + return [] + } + } as never, + identityStore: {} as never, + signalIdentitySync: {} as never, + signalSessionSync: {} as never, + addressResolver, + logger: createNoopLogger() + }) + + await assert.rejects( + sessionResolver.ensureSessionsBatch( + ['5511999999999:2@s.whatsapp.net', '778899:2@lid'], + new Map([ + ['5511999999999:2@s.whatsapp.net', new Uint8Array(32).fill(1)], + ['778899:2@lid', new Uint8Array(32).fill(2)] + ]) + ), + /identity mismatch/ + ) + assert.equal(storeReads, 0) +}) + test('signal session resolver rejects identity mismatch on reasonIdentity sync', async () => { let syncedIdentityKeys = 0 diff --git a/src/signal/session/__tests__/session.test.ts b/src/signal/session/__tests__/session.test.ts index 5a221d91..11955ddc 100644 --- a/src/signal/session/__tests__/session.test.ts +++ b/src/signal/session/__tests__/session.test.ts @@ -11,6 +11,7 @@ import { generateSignedPreKey } from '@signal/registration/keygen' import { encodeSignalSessionSnapshot } from '@signal/session/encoding' +import { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import { SignalProtocol } from '@signal/session/SignalProtocol' import { decryptMsg, deriveMsgKey, selectMessageKey } from '@signal/session/SignalRatchet' import { @@ -26,6 +27,7 @@ import type { SignalSessionSnapshot } from '@signal/types' import { WaIdentityMemoryStore } from '@store/memory/identity.store' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { WaPreKeyMemoryStore } from '@store/memory/pre-key.store' import { WaSessionMemoryStore } from '@store/memory/session.store' import { WaSignalMemoryStore } from '@store/memory/signal.store' @@ -257,6 +259,167 @@ test('signal protocol establishes outgoing session and decrypts prekey message o ) }) +test('PN sender metadata keeps replies and later PN messages on one LID session', async () => { + const logger = createNoopLogger() + const aliceSignal = new WaSignalMemoryStore() + const alicePreKeys = new WaPreKeyMemoryStore() + const aliceSessions = new WaSessionMemoryStore() + const aliceIdentities = new WaIdentityMemoryStore() + const bobSignal = new WaSignalMemoryStore() + const bobPreKeys = new WaPreKeyMemoryStore() + const bobSessions = new WaSessionMemoryStore() + const bobIdentities = new WaIdentityMemoryStore() + const mappings = new WaLidPnMappingMemoryStore() + + const [aliceRegistration, bobRegistration] = await Promise.all([ + generateRegistrationInfo(), + generateRegistrationInfo() + ]) + await Promise.all([ + aliceSignal.setRegistrationInfo(aliceRegistration), + bobSignal.setRegistrationInfo(bobRegistration) + ]) + const bobSignedPreKey = await generateSignedPreKey(1, bobRegistration.identityKeyPair.privKey) + const bobOneTimePreKey = await generatePreKeyPair(9) + await Promise.all([ + bobSignal.setSignedPreKey(bobSignedPreKey), + bobPreKeys.putPreKey(bobOneTimePreKey) + ]) + + const aliceProtocol = new SignalProtocol( + { + signal: aliceSignal, + preKey: alicePreKeys, + session: aliceSessions, + identity: aliceIdentities + }, + logger + ) + const bobStores = { + signal: bobSignal, + preKey: bobPreKeys, + session: bobSessions, + identity: bobIdentities + } + const bobResolver = new SignalAddressResolver(mappings) + const bobProtocol = new SignalProtocol(bobStores, logger, bobResolver) + const alicePn = makeAddress('5511000000101') + const aliceLid: SignalAddress = { user: '99112233', server: 'lid', device: 0 } + const bobAddress = makeAddress('5511000000202') + + await aliceProtocol.establishOutgoingSession(bobAddress, { + regId: bobRegistration.registrationId, + identity: bobRegistration.identityKeyPair.pubKey, + signedKey: { + id: bobSignedPreKey.keyId, + publicKey: bobSignedPreKey.keyPair.pubKey, + signature: bobSignedPreKey.signature + }, + oneTimeKey: { + id: bobOneTimePreKey.keyId, + publicKey: bobOneTimePreKey.keyPair.pubKey + } + }) + + const firstPlaintext = makeBytes(24, 41) + const first = await aliceProtocol.encryptMessage(bobAddress, firstPlaintext) + assert.equal(first.type, 'pkmsg') + + // The outer direct-message stanza is PN-addressed and carries sender_lid. + await bobResolver.learnMessageJidPair(`${alicePn.user}@s.whatsapp.net`, `${aliceLid.user}@lid`) + assert.deepEqual( + await bobProtocol.decryptMessage(alicePn, { + type: first.type, + ciphertext: first.ciphertext + }), + firstPlaintext + ) + assert.equal(await bobSessions.getSession(alicePn), null) + assert.ok(await bobSessions.getSession(aliceLid)) + + const replyPlaintext = makeBytes(19, 73) + const reply = await bobProtocol.encryptMessage(aliceLid, replyPlaintext) + assert.deepEqual( + await aliceProtocol.decryptMessage(bobAddress, { + type: reply.type, + ciphertext: reply.ciphertext + }), + replyPlaintext + ) + + // A later PN stanza need not repeat sender_lid: a fresh resolver reloads + // the persisted mapping and selects the already-advanced LID ratchet. + const restartedBobProtocol = new SignalProtocol( + bobStores, + logger, + new SignalAddressResolver(mappings) + ) + const secondPlaintext = makeBytes(21, 109) + const second = await aliceProtocol.encryptMessage(bobAddress, secondPlaintext) + assert.equal(second.type, 'msg') + assert.deepEqual( + await restartedBobProtocol.decryptMessage(alicePn, { + type: second.type, + ciphertext: second.ciphertext + }), + secondPlaintext + ) +}) + +test('signal protocol persists one session when PN/LID batch entries converge', async () => { + const mappings = new WaLidPnMappingMemoryStore() + const addressResolver = new SignalAddressResolver(mappings) + await addressResolver.learnMessageJidPair('5511444444444@s.whatsapp.net', '445566@lid') + const persistedSessions: Array<{ readonly address: SignalAddress }> = [] + const persistedIdentities: Array<{ readonly address: SignalAddress }> = [] + const protocol = new SignalProtocol( + { + signal: {} as never, + preKey: {} as never, + session: { + getSessionsBatch: async (addresses: readonly SignalAddress[]) => + addresses.map(() => null), + setSessionsBatch: async ( + entries: readonly { readonly address: SignalAddress }[] + ) => { + persistedSessions.push(...entries) + } + } as never, + identity: { + setRemoteIdentities: async ( + entries: readonly { readonly address: SignalAddress }[] + ) => { + persistedIdentities.push(...entries) + } + } as never + }, + createNoopLogger(), + addressResolver + ) + const remoteIdentity = new Uint8Array(33).fill(7) + const firstSession = { remote: { pubKey: remoteIdentity } } as SignalSessionRecord + const secondSession = { remote: { pubKey: remoteIdentity } } as SignalSessionRecord + + const result = await protocol.persistOutgoingSessionsBatch([ + { + address: { user: '5511444444444', server: 's.whatsapp.net', device: 2 }, + session: firstSession, + remoteIdentity + }, + { + address: { user: '445566', server: 'lid', device: 2 }, + session: secondSession, + remoteIdentity + } + ]) + + assert.equal(result.resolved.length, 1) + assert.strictEqual(result.resolved[0].session, firstSession) + assert.deepEqual(result.resolved[0].address, { user: '445566', server: 'lid', device: 2 }) + assert.equal(persistedSessions.length, 1) + assert.equal(persistedIdentities.length, 1) +}) + test('signal protocol throws when decrypting msg without an existing session', async () => { const store = new WaSignalMemoryStore() const registration = await generateRegistrationInfo() diff --git a/src/signal/session/resolver.ts b/src/signal/session/resolver.ts index 4c8fda5e..22bd3f4b 100644 --- a/src/signal/session/resolver.ts +++ b/src/signal/session/resolver.ts @@ -4,6 +4,7 @@ import { PromiseDedup } from '@infra/perf/PromiseDedup' import { normalizeDeviceJid, parseSignalAddressFromJid, signalAddressKey } from '@protocol/jid' import type { SignalIdentitySyncApi } from '@signal/api/SignalIdentitySyncApi' import type { SignalSessionSyncApi } from '@signal/api/SignalSessionSyncApi' +import type { SignalAddressResolver } from '@signal/session/SignalAddressResolver' import type { SignalProtocol } from '@signal/session/SignalProtocol' import type { SignalAddress, SignalPreKeyBundle, SignalSessionRecord } from '@signal/types' import type { WaIdentityStore } from '@store/contracts/identity.store' @@ -51,6 +52,7 @@ export function createSignalSessionResolver(options: { readonly identityStore: WaIdentityStore readonly signalIdentitySync: SignalIdentitySyncApi readonly signalSessionSync: SignalSessionSyncApi + readonly addressResolver?: SignalAddressResolver readonly logger: Logger }): SignalSessionResolver { const { @@ -59,6 +61,7 @@ export function createSignalSessionResolver(options: { identityStore, signalIdentitySync, signalSessionSync, + addressResolver, logger } = options const dedup = new PromiseDedup() @@ -162,13 +165,15 @@ export function createSignalSessionResolver(options: { ) } - const ensureSession = ( + const ensureSession = async ( address: SignalAddress, jid: string, expectedIdentity?: Uint8Array, reasonIdentity = false - ): Promise => - ensureSessionWithDedup(address, jid, expectedIdentity, reasonIdentity).then(() => {}) + ): Promise => { + const resolvedAddress = addressResolver ? await addressResolver.resolve(address) : address + await ensureSessionWithDedup(resolvedAddress, jid, expectedIdentity, reasonIdentity) + } const ensureSessionsBatch = async ( targetJids: readonly string[], @@ -194,15 +199,17 @@ export function createSignalSessionResolver(options: { normalizedTargetJids.length = normalizedTargetCount normalizedTargetAddresses.length = normalizedTargetCount - const normalizedExpectedIdentityByJid = + const serializedExpectedIdentityByJid = expectedIdentityByJid && expectedIdentityByJid.size > 0 ? new Map() : undefined - if (normalizedExpectedIdentityByJid && expectedIdentityByJid) { + if (serializedExpectedIdentityByJid && expectedIdentityByJid) { for (const [jid, identity] of expectedIdentityByJid.entries()) { try { - toSerializedPubKey(identity) - normalizedExpectedIdentityByJid.set(normalizeDeviceJid(jid), identity) + serializedExpectedIdentityByJid.set( + normalizeDeviceJid(jid), + toSerializedPubKey(identity) + ) } catch (error) { logger.trace( 'ignoring malformed expected identity jid during batch normalization', @@ -211,6 +218,50 @@ export function createSignalSessionResolver(options: { } } } + const expectedIdentityByAddressKey = serializedExpectedIdentityByJid + ? new Map() + : undefined + if (addressResolver) { + const resolvedAddresses = await addressResolver.resolveMany(normalizedTargetAddresses) + const seenAddresses = new Set() + let canonicalCount = 0 + for (let index = 0; index < normalizedTargetCount; index += 1) { + const address = resolvedAddresses[index] + const key = signalAddressKey(address) + const expectedIdentity = serializedExpectedIdentityByJid?.get( + normalizedTargetJids[index] + ) + const previousExpectedIdentity = expectedIdentityByAddressKey?.get(key) + if ( + expectedIdentity && + previousExpectedIdentity && + !uint8Equal(expectedIdentity, previousExpectedIdentity) + ) { + throw new Error('identity mismatch') + } + if (expectedIdentity) expectedIdentityByAddressKey?.set(key, expectedIdentity) + if (seenAddresses.has(key)) continue + seenAddresses.add(key) + normalizedTargetJids[canonicalCount] = normalizedTargetJids[index] + normalizedTargetAddresses[canonicalCount] = address + canonicalCount += 1 + } + normalizedTargetJids.length = canonicalCount + normalizedTargetAddresses.length = canonicalCount + normalizedTargetCount = canonicalCount + } else if (expectedIdentityByAddressKey) { + for (let index = 0; index < normalizedTargetCount; index += 1) { + const expectedIdentity = serializedExpectedIdentityByJid?.get( + normalizedTargetJids[index] + ) + if (expectedIdentity) { + expectedIdentityByAddressKey.set( + signalAddressKey(normalizedTargetAddresses[index]), + expectedIdentity + ) + } + } + } const resolvedByIndex = (await sessionStore.getSessionsBatch( normalizedTargetAddresses @@ -239,17 +290,16 @@ export function createSignalSessionResolver(options: { const missingIndices: number[] = [] for (let index = 0; index < normalizedTargetJids.length; index += 1) { const session = resolvedByIndex[index] - const expectedIdentity = normalizedExpectedIdentityByJid?.get( - normalizedTargetJids[index] + const expectedIdentity = expectedIdentityByAddressKey?.get( + signalAddressKey(normalizedTargetAddresses[index]) ) if (session && expectedIdentity) { - const expectedSerialized = toSerializedPubKey(expectedIdentity) - if (!uint8Equal(session.remote.pubKey, expectedSerialized)) { + if (!uint8Equal(session.remote.pubKey, expectedIdentity)) { logger.warn('signal identity mismatch on existing session vs expected', { jid: normalizedTargetJids[index], source: 'session_vs_expected', session: bytesToHex(session.remote.pubKey), - expected: bytesToHex(expectedSerialized) + expected: bytesToHex(expectedIdentity) }) throw new Error('identity mismatch') } @@ -301,10 +351,9 @@ export function createSignalSessionResolver(options: { }) continue } - const expectedIdentity = normalizedExpectedIdentityByJid?.get(targetJid) - const expectedSerializedIdentity = expectedIdentity - ? toSerializedPubKey(expectedIdentity) - : null + const expectedSerializedIdentity = expectedIdentityByAddressKey?.get( + signalAddressKey(normalizedTargetAddresses[targetIndex]) + ) const bundleIdentity = toSerializedPubKey(batchResult.bundle.identity) if ( expectedSerializedIdentity && diff --git a/src/store/__tests__/create-store.test.ts b/src/store/__tests__/create-store.test.ts index b9a73cac..791404de 100644 --- a/src/store/__tests__/create-store.test.ts +++ b/src/store/__tests__/create-store.test.ts @@ -2,6 +2,8 @@ import assert from 'node:assert/strict' import test from 'node:test' import { createStore } from '@store/createStore' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' +import { WaSessionMemoryStore } from '@store/memory/session.store' const mockAuthBackend = { stores: { @@ -181,3 +183,36 @@ test('createStore allows omitting cacheProviders when backends is set (caches de assert.ok(session.deviceList) assert.ok(session.messageSecret) }) + +test('createStore couples the PN/LID mapping provider to the session backend', async () => { + const mappings = new WaLidPnMappingMemoryStore() + const backend = { + ...mockAuthBackend, + stores: { + ...mockAuthBackend.stores, + session: () => new WaSessionMemoryStore(), + lidPnMapping: () => mappings + } + } + const store = createStore({ + backends: { mock: backend }, + providers: { + auth: 'memory', + signal: 'memory', + preKey: 'memory', + session: 'mock', + identity: 'memory', + senderKey: 'memory', + appState: 'memory', + privacyToken: 'memory', + messages: 'none', + threads: 'none', + contacts: 'none' + } + }) + + const session = store.session('mapping-session') + await session.lidPnMapping!.setLidUser('5511999999999', '778899') + assert.equal(await mappings.getLidUser('5511999999999'), '778899') + await store.destroy() +}) diff --git a/src/store/contracts/lid-pn-mapping.store.ts b/src/store/contracts/lid-pn-mapping.store.ts new file mode 100644 index 00000000..759a5ec0 --- /dev/null +++ b/src/store/contracts/lid-pn-mapping.store.ts @@ -0,0 +1,15 @@ +/** + * Persistent one-to-one phone-number/LID mapping used to canonicalize Signal + * addresses. Values are bare user components; one mapping applies to every + * device of the account. + */ +export interface WaLidPnMappingStore { + /** Returns the current LID user component for a PN user component. */ + getLidUser(pnUser: string): Promise + /** Returns the PN user component that currently owns a LID user component. */ + getPnUser(lidUser: string): Promise + /** Replaces any mapping that currently owns either user component. */ + setLidUser(pnUser: string, lidUser: string): Promise + /** Removes every mapping in the current store session. */ + clear(): Promise +} diff --git a/src/store/createStore.ts b/src/store/createStore.ts index 1192ed11..71887b89 100644 --- a/src/store/createStore.ts +++ b/src/store/createStore.ts @@ -8,6 +8,7 @@ import type { WaContactStore } from '@store/contracts/contact.store' import type { WaDeviceListStore } from '@store/contracts/device-list.store' import type { WaGroupMetadataStore } from '@store/contracts/group-metadata.store' import type { WaIdentityStore } from '@store/contracts/identity.store' +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' import type { WaMessageSecretStore } from '@store/contracts/message-secret.store' import type { WaMessageStore } from '@store/contracts/message.store' import type { WaPreKeyStore } from '@store/contracts/pre-key.store' @@ -23,6 +24,7 @@ import { withContactLock } from '@store/locks/contact.lock' import { withDeviceListLock } from '@store/locks/device-list.lock' import { withGroupMetadataLock } from '@store/locks/group-metadata.lock' import { withIdentityLock } from '@store/locks/identity.lock' +import { withLidPnMappingLock } from '@store/locks/lid-pn-mapping.lock' import { withMessageSecretLock } from '@store/locks/message-secret.lock' import { withMessageLock } from '@store/locks/message.lock' import { withPreKeyLock } from '@store/locks/pre-key.lock' @@ -38,6 +40,7 @@ import { WaContactMemoryStore } from '@store/memory/contact.store' import { WaDeviceListMemoryStore } from '@store/memory/device-list.store' import { WaGroupMetadataMemoryStore } from '@store/memory/group-metadata.store' import { WaIdentityMemoryStore } from '@store/memory/identity.store' +import { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' import { WaMessageSecretMemoryStore } from '@store/memory/message-secret.store' import { WaMessageMemoryStore } from '@store/memory/message.store' import { WaPreKeyMemoryStore } from '@store/memory/pre-key.store' @@ -294,6 +297,15 @@ export function createStore(options?: WaCreateStoreOptions) maxRemoteIdentities: ml.signalRemoteIdentities }) ) + const sessionProvider = providers.session ?? 'memory' + const mappingFactory = usesBackend(sessionProvider) + ? backends[sessionProvider]?.stores.lidPnMapping + : undefined + const rawLidPnMapping: WaLidPnMappingStore = mappingFactory + ? mappingFactory(id) + : new WaLidPnMappingMemoryStore({ + maxMappings: ml.signalLidPnMappings + }) const rawSenderKey = resolveStore( id, backends, @@ -436,6 +448,7 @@ export function createStore(options?: WaCreateStoreOptions) ? withIdentityCache(rawIdentity, cacheLayer.limits?.identity) : rawIdentity ) + const lidPnMappingStore = withLidPnMappingLock(rawLidPnMapping) const senderKeyStore = withSenderKeyLock( cacheLayer.senderKey && usesBackend(providers.senderKey) ? withSenderKeyCache(rawSenderKey, cacheLayer.limits?.senderKey) @@ -485,6 +498,7 @@ export function createStore(options?: WaCreateStoreOptions) destroyIfSupported(preKeyStore), destroyIfSupported(sessionStore), destroyIfSupported(identityStore), + destroyIfSupported(lidPnMappingStore), destroyIfSupported(senderKeyStore), destroyIfSupported(appStateStore), destroyIfSupported(messageStore), @@ -500,6 +514,7 @@ export function createStore(options?: WaCreateStoreOptions) preKey: preKeyStore, session: sessionStore, identity: identityStore, + lidPnMapping: lidPnMappingStore, senderKey: senderKeyStore, appState: appStateStore, retry: retryStore, diff --git a/src/store/index.ts b/src/store/index.ts index 816e81ac..972d95d2 100644 --- a/src/store/index.ts +++ b/src/store/index.ts @@ -27,6 +27,7 @@ export type { WaAppStateStore } from '@store/contracts/appstate.store' export type { WaIdentityStore } from '@store/contracts/identity.store' +export type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' export type { WaPreKeyStore } from '@store/contracts/pre-key.store' export type { WaSenderKeyStore } from '@store/contracts/sender-key.store' export type { WaSessionStore } from '@store/contracts/session.store' @@ -43,6 +44,7 @@ export { WaSignalMemoryStore } from '@store/memory/signal.store' export { WaPreKeyMemoryStore } from '@store/memory/pre-key.store' export { WaSessionMemoryStore } from '@store/memory/session.store' export { WaIdentityMemoryStore } from '@store/memory/identity.store' +export { WaLidPnMappingMemoryStore } from '@store/memory/lid-pn-mapping.store' export { SenderKeyMemoryStore } from '@store/memory/sender-key.store' export { WaRetryMemoryStore } from '@store/memory/retry.store' export { WaGroupMetadataMemoryStore } from '@store/memory/group-metadata.store' diff --git a/src/store/locks/lid-pn-mapping.lock.ts b/src/store/locks/lid-pn-mapping.lock.ts new file mode 100644 index 00000000..1efcd19d --- /dev/null +++ b/src/store/locks/lid-pn-mapping.lock.ts @@ -0,0 +1,22 @@ +import { SharedExclusiveGate } from '@infra/perf/SharedExclusiveGate' +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' +import type { WithDestroyLifecycle } from '@store/types' + +export function withLidPnMappingLock( + store: WaLidPnMappingStore +): WithDestroyLifecycle { + const gate = new SharedExclusiveGate() + const destroyStore = store as { destroy?: () => Promise } + return { + getLidUser: (pnUser) => gate.runShared(() => store.getLidUser(pnUser)), + getPnUser: (lidUser) => gate.runShared(() => store.getPnUser(lidUser)), + // Replacing either side can evict a third-party pair, so every mapping write + // shares one exclusive boundary. Writes are rare; reads stay concurrent. + setLidUser: (pnUser, lidUser) => gate.runExclusive(() => store.setLidUser(pnUser, lidUser)), + clear: () => gate.runExclusive(() => store.clear()), + destroy: async () => { + await gate.close() + await destroyStore.destroy?.() + } + } +} diff --git a/src/store/memory/lid-pn-mapping.store.ts b/src/store/memory/lid-pn-mapping.store.ts new file mode 100644 index 00000000..0f0e1947 --- /dev/null +++ b/src/store/memory/lid-pn-mapping.store.ts @@ -0,0 +1,57 @@ +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' +import { resolvePositive } from '@util/coercion' +import { setBoundedMapEntry } from '@util/collections' + +const DEFAULT_MAX_MAPPINGS = 8_192 + +export interface WaLidPnMappingMemoryStoreOptions { + /** Maximum mappings retained by the bounded in-process store. Default: `8_192`. */ + readonly maxMappings?: number +} + +/** Bounded in-process implementation of {@link WaLidPnMappingStore}. */ +export class WaLidPnMappingMemoryStore implements WaLidPnMappingStore { + private readonly mappings = new Map() + private readonly reverseMappings = new Map() + private readonly maxMappings: number + + public constructor(options: WaLidPnMappingMemoryStoreOptions = {}) { + this.maxMappings = resolvePositive( + options.maxMappings, + DEFAULT_MAX_MAPPINGS, + 'WaLidPnMappingMemoryStore.maxMappings' + ) + } + + public async getLidUser(pnUser: string): Promise { + return this.mappings.get(pnUser) ?? null + } + + public async getPnUser(lidUser: string): Promise { + return this.reverseMappings.get(lidUser) ?? null + } + + public async setLidUser(pnUser: string, lidUser: string): Promise { + const previousLid = this.mappings.get(pnUser) + if (previousLid !== undefined) this.reverseMappings.delete(previousLid) + const previousPn = this.reverseMappings.get(lidUser) + if (previousPn !== undefined) this.mappings.delete(previousPn) + setBoundedMapEntry( + this.mappings, + pnUser, + lidUser, + this.maxMappings, + (evictedPn, evictedLid) => { + if (this.reverseMappings.get(evictedLid) === evictedPn) { + this.reverseMappings.delete(evictedLid) + } + } + ) + this.reverseMappings.set(lidUser, pnUser) + } + + public async clear(): Promise { + this.mappings.clear() + this.reverseMappings.clear() + } +} diff --git a/src/store/types.ts b/src/store/types.ts index 2160d169..78701c48 100644 --- a/src/store/types.ts +++ b/src/store/types.ts @@ -5,6 +5,7 @@ import type { WaContactStore } from '@store/contracts/contact.store' import type { WaDeviceListStore } from '@store/contracts/device-list.store' import type { WaGroupMetadataStore } from '@store/contracts/group-metadata.store' import type { WaIdentityStore } from '@store/contracts/identity.store' +import type { WaLidPnMappingStore } from '@store/contracts/lid-pn-mapping.store' import type { WaMessageSecretStore } from '@store/contracts/message-secret.store' import type { WaMessageStore } from '@store/contracts/message.store' import type { WaPreKeyStore } from '@store/contracts/pre-key.store' @@ -23,6 +24,8 @@ export interface WaStoreSession { readonly preKey: WaPreKeyStore readonly session: WaSessionStore readonly identity: WaIdentityStore + /** Internal companion for canonical Signal addressing. Older custom stores may omit it. */ + readonly lidPnMapping?: WaLidPnMappingStore readonly senderKey: WaSenderKeyStore readonly appState: WaAppStateStore readonly retry: WaRetryStore @@ -50,6 +53,11 @@ export interface WaStoreBackend { readonly preKey: (sessionId: string) => WaPreKeyStore readonly session: (sessionId: string) => WaSessionStore readonly identity: (sessionId: string) => WaIdentityStore + /** + * Optional companion to `session`. First-party backends persist it; + * older third-party backends fall back to the in-process provider. + */ + readonly lidPnMapping?: (sessionId: string) => WaLidPnMappingStore readonly senderKey: (sessionId: string) => WaSenderKeyStore readonly appState: (sessionId: string) => WaAppStateStore readonly messages: (sessionId: string) => WaMessageStore @@ -113,7 +121,9 @@ export interface WaCreateStoreOptions { * Signal sessions (the per-peer Double Ratchet state). Losing this * forces a transparent re-handshake on the next message to/from that * peer – your identity key doesn't change, so no "security code - * changed" notice fires on the peer. Default: `'memory'`. + * changed" notice fires on the peer. Official backends also persist + * the internal PN/LID address mapping with this provider so the same + * ratchet key survives a restart. Default: `'memory'`. */ readonly session?: B | 'memory' /** @@ -352,6 +362,7 @@ export interface WaStoreMemoryLimitSelection { readonly appStateCollectionEntries?: number readonly signalPreKeys?: number readonly signalSessions?: number + readonly signalLidPnMappings?: number readonly signalRemoteIdentities?: number readonly senderKeys?: number readonly senderDistributions?: number