diff --git a/CHANGELOG.md b/CHANGELOG.md index d4f1433..2c87ad7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,7 @@ - Attachments sent alongside text no longer incorrectly carry the reply-to relationship - Incoming formatted messages were parsed twice; removed redundant `` pre-strip pass - `matrixSDKLogConfigured` flag no longer latches when `setLevel` is missing from the SDK logger +- Chat SDK now classifies rooms listed in `m.direct` as direct messages ### Changes diff --git a/src/index.test.ts b/src/index.test.ts index cc27915..b15ec3f 100644 --- a/src/index.test.ts +++ b/src/index.test.ts @@ -2,7 +2,13 @@ import { describe, expect, it, vi } from "vitest"; import { Chat, getEmoji, stringifyMarkdown } from "chat"; import type { AdapterPostableMessage, ChatInstance, Logger, StateAdapter } from "chat"; import { createMemoryState } from "@chat-adapter/state-memory"; -import { EventType, MsgType, RelationType, type MatrixClient } from "matrix-js-sdk"; +import { + ClientEvent, + EventType, + MsgType, + RelationType, + type MatrixClient, +} from "matrix-js-sdk"; import { MatrixError } from "matrix-js-sdk/lib/http-api/errors"; import { encodeRecoveryKey } from "matrix-js-sdk/lib/crypto-api/recovery-key"; import { createMatrixAdapter, MatrixAdapter } from "./index"; @@ -225,6 +231,10 @@ function makeClient() { getAccountDataFromServer: vi.fn( async (): Promise | null> => null ), + getAccountData: vi.fn( + (_type: string): { getContent: () => Record } | undefined => + undefined + ), getAccessToken: vi.fn(() => "token"), getCrypto: vi.fn(() => crypto), getEventMapper: vi.fn(() => (raw: Record) => mapRawToEvent(raw)), @@ -457,6 +467,90 @@ describe("MatrixAdapter", () => { }); }); + it("classifies only m.direct rooms as DMs and refreshes the synchronous cache", async () => { + const client = makeClient(); + client.getAccountDataFromServer.mockResolvedValue({ + "@alice:beeper.com": ["!direct:beeper.com"], + }); + const adapter = createMatrixAdapter({ + baseURL: "https://matrix.example.com", + auth: { + type: "accessToken", + accessToken: "token", + userID: "@bot:beeper.com", + }, + createClient: () => asMatrixClient(client), + }); + + await adapter.initialize(makeChatInstance()); + + expect(client.getAccountDataFromServer).toHaveBeenCalledWith(EventType.Direct); + expect(adapter.isDM("matrix:!direct%3Abeeper.com")).toBe(true); + // A two-person room is still a normal room unless m.direct marks it. + expect(adapter.isDM("matrix:!unmarked-two-person%3Abeeper.com")).toBe(false); + expect(adapter.isDM("not-a-matrix-thread")).toBe(false); + + const newDirectThread = await adapter.openDM("@dave:beeper.com"); + expect(adapter.isDM(newDirectThread)).toBe(true); + + client.__handlers.get(ClientEvent.AccountData)?.( + makeEvent({ + getType: () => EventType.Direct, + getContent: () => ({ + "@erin:beeper.com": ["!updated-direct:beeper.com"], + }), + }) + ); + expect(adapter.isDM("matrix:!direct%3Abeeper.com")).toBe(false); + expect(adapter.isDM("matrix:!updated-direct%3Abeeper.com")).toBe(true); + }); + + it("fails closed when cold-start m.direct priming is unavailable", async () => { + const client = makeClient(); + client.getAccountDataFromServer.mockRejectedValue(new Error("homeserver unavailable")); + const adapter = createMatrixAdapter({ + baseURL: "https://matrix.example.com", + auth: { + type: "accessToken", + accessToken: "token", + userID: "@bot:beeper.com", + }, + createClient: () => asMatrixClient(client), + }); + + await expect(adapter.initialize(makeChatInstance())).resolves.toBeUndefined(); + + expect(client.startClient).toHaveBeenCalledOnce(); + expect(adapter.isDM("matrix:!unmarked%3Abeeper.com")).toBe(false); + }); + + it("refreshes stale cached m.direct data during initialization", async () => { + const client = makeClient(); + client.getAccountData.mockReturnValue({ + getContent: () => ({ + "@alice:beeper.com": ["!stale-direct:beeper.com"], + }), + }); + client.getAccountDataFromServer.mockResolvedValue({ + "@alice:beeper.com": ["!fresh-direct:beeper.com"], + }); + const adapter = createMatrixAdapter({ + baseURL: "https://matrix.example.com", + auth: { + type: "accessToken", + accessToken: "token", + userID: "@bot:beeper.com", + }, + createClient: () => asMatrixClient(client), + }); + + await adapter.initialize(makeChatInstance()); + + expect(client.getAccountDataFromServer).toHaveBeenCalledWith(EventType.Direct); + expect(adapter.isDM("matrix:!stale-direct%3Abeeper.com")).toBe(false); + expect(adapter.isDM("matrix:!fresh-direct%3Abeeper.com")).toBe(true); + }); + it("rejects thread IDs with an empty room ID", () => { const adapter = new MatrixAdapter({ baseURL: "https://hs.beeper.com", @@ -2676,12 +2770,7 @@ describe("MatrixAdapter", () => { it("merges fresh m.direct account data before persisting a newly created DM", async () => { const fakeClient = makeClient(); - fakeClient.getAccountDataFromServer - .mockResolvedValueOnce({}) - .mockResolvedValueOnce({ - "@bob:beeper.com": ["!existing-dm:beeper.com"], - "@carol:beeper.com": ["!carol-dm:beeper.com"], - }); + fakeClient.getAccountDataFromServer.mockResolvedValue({}); fakeClient.createRoom.mockResolvedValue({ room_id: "!new-dm:beeper.com" }); const adapter = new MatrixAdapter({ @@ -2691,6 +2780,12 @@ describe("MatrixAdapter", () => { }); await adapter.initialize(makeChatInstance({ state: makeStateAdapter() })); + fakeClient.getAccountDataFromServer + .mockResolvedValueOnce({}) + .mockResolvedValueOnce({ + "@bob:beeper.com": ["!existing-dm:beeper.com"], + "@carol:beeper.com": ["!carol-dm:beeper.com"], + }); await adapter.openDM("@bob:beeper.com"); expect(fakeClient.setAccountData).toHaveBeenCalledWith(EventType.Direct, { diff --git a/src/index.ts b/src/index.ts index 042be70..2a38392 100644 --- a/src/index.ts +++ b/src/index.ts @@ -208,6 +208,7 @@ export class MatrixAdapter implements Adapter { private readonly reactionByEventID = new Map(); private readonly myReactionByKey = new Map(); private readonly processedTimelineEventIDs = new Set(); + private readonly directRoomIDs = new Set(); private lastSecretsBundlePersistAt = 0; private secretsBundleUnavailableLogged = false; private liveSyncReady = false; @@ -294,8 +295,14 @@ export class MatrixAdapter implements Adapter { } this.dispatchTimelineEvent(event, undefined, false); }); + this.client.on(ClientEvent.AccountData, (event) => { + if (event.getType() === EventType.Direct) { + this.replaceDirectRoomIDs(this.normalizeDirectAccountData(event.getContent())); + } + }); await this.maybeInitE2EE(); + await this.primeDirectRoomIDs(); await this.client.startClient(this.syncOptions); this.started = true; @@ -322,6 +329,7 @@ export class MatrixAdapter implements Adapter { this.client.stopClient(); this.reactionByEventID.clear(); this.myReactionByKey.clear(); + this.directRoomIDs.clear(); this.client = null; this.started = false; this.logger.info("Matrix adapter shutdown complete"); @@ -352,6 +360,15 @@ export class MatrixAdapter implements Adapter { return channelIdFromThreadId(threadId); } + isDM(threadId: string): boolean { + try { + const { roomID } = this.decodeThreadId(threadId); + return this.directRoomIDs.has(roomID); + } catch { + return false; + } + } + renderFormatted(content: FormattedContent): string { return stringifyMarkdown(content); } @@ -1042,11 +1059,38 @@ export class MatrixAdapter implements Adapter { private async loadDirectAccountData(): Promise { const cached = this.loadCachedDirectAccountData(); if (Object.keys(cached).length > 0) { + this.replaceDirectRoomIDs(cached); return cached; } const direct = await this.requireClient().getAccountDataFromServer(EventType.Direct); - return this.normalizeDirectAccountData(direct); + const normalized = this.normalizeDirectAccountData(direct); + this.replaceDirectRoomIDs(normalized); + return normalized; + } + + private async primeDirectRoomIDs(): Promise { + // Install cached m.direct data as a fallback, but always refresh it from the + // homeserver so warm starts cannot retain stale room classifications. + this.replaceDirectRoomIDs(this.loadCachedDirectAccountData()); + try { + const direct = await this.requireClient().getAccountDataFromServer(EventType.Direct); + this.replaceDirectRoomIDs(this.normalizeDirectAccountData(direct)); + } catch (error) { + this.logger.warn( + "Failed to refresh Matrix direct rooms; retaining cached m.direct data and treating other rooms as channels", + { error } + ); + } + } + + private replaceDirectRoomIDs(direct: DirectAccountData): void { + this.directRoomIDs.clear(); + for (const roomIDs of Object.values(direct)) { + for (const roomID of roomIDs) { + this.directRoomIDs.add(roomID); + } + } } private loadCachedDirectAccountData(): DirectAccountData { @@ -1140,6 +1184,7 @@ export class MatrixAdapter implements Adapter { [userID]: [...existingRooms, roomID], }; await this.requireClient().setAccountData(EventType.Direct, updated); + this.directRoomIDs.add(roomID); } }