From 316eaae653bd1f85a55542982ad8a72c6ae445e9 Mon Sep 17 00:00:00 2001 From: louzt <179385168+louzt@users.noreply.github.com> Date: Thu, 23 Apr 2026 20:22:43 -0600 Subject: [PATCH 1/5] Harden IRC client lifecycle behavior --- src/lib/irc/IRCClient.ts | 38 ++++++++++++++++++++++--- tests/lib/ircClient.test.ts | 57 +++++++++++++++++++++++++++++++++++++ 2 files changed, 91 insertions(+), 4 deletions(-) diff --git a/src/lib/irc/IRCClient.ts b/src/lib/irc/IRCClient.ts index 5b1d500e..6860a605 100644 --- a/src/lib/irc/IRCClient.ts +++ b/src/lib/irc/IRCClient.ts @@ -391,6 +391,7 @@ type EventCallback = (data: EventMap[K]) => void; export class IRCClient implements IRCClientContext { private sockets: Map = new Map(); + private intentionalDisconnects: Set = new Set(); servers: Map = new Map(); nicks: Map = new Map(); currentUsers: Map = new Map(); // Per-server current users @@ -491,6 +492,10 @@ export class IRCClient implements IRCClientContext { ): Promise { const connectionKey = `${host}:${port}`; + if (serverId) { + this.intentionalDisconnects.delete(serverId); + } + // Check if there's already a pending connection to this server const existingConnection = this.pendingConnections.get(connectionKey); if (existingConnection) { @@ -614,6 +619,7 @@ export class IRCClient implements IRCClientContext { status: "online", }); this.nicks.set(server.id, nickname); + this.intentionalDisconnects.delete(server.id); socket.onopen = () => { //registerAllProtocolHandlers(this); @@ -649,6 +655,13 @@ export class IRCClient implements IRCClientContext { // to ensure connection is fully established before sending PINGs socket.onclose = () => { + if (this.intentionalDisconnects.has(server.id)) { + this.stopWebSocketPing(server.id); + this.sockets.delete(server.id); + this.pendingConnections.delete(connectionKey); + return; + } + if (!this.servers.has(server.id)) { return; } @@ -712,13 +725,29 @@ export class IRCClient implements IRCClientContext { disconnect(serverId: string, quitMessage?: string): void { const socket = this.sockets.get(serverId); + const server = this.servers.get(serverId); + if (socket) { - const message = quitMessage || "ObsidianIRC - Bringing IRC to the future"; - socket.send(`QUIT :${message}`); - socket.close(); + const CONNECTING = 0; + const OPEN = 1; + + if (socket.readyState === CONNECTING || socket.readyState === OPEN) { + this.intentionalDisconnects.add(serverId); + } + + if (socket.readyState === OPEN && server?.isConnected) { + const message = + quitMessage || "ObsidianIRC - Bringing IRC to the future"; + socket.send(`QUIT :${message}`); + } + + if (socket.readyState === CONNECTING || socket.readyState === OPEN) { + socket.close(); + } + this.sockets.delete(serverId); } - const server = this.servers.get(serverId); + if (server) { server.isConnected = false; server.connectionState = "disconnected"; @@ -742,6 +771,7 @@ export class IRCClient implements IRCClientContext { removeServer(serverId: string): void { this.disconnect(serverId); + this.intentionalDisconnects.delete(serverId); this.servers.delete(serverId); this.capNegotiationComplete.delete(serverId); this.pendingCapReqs.delete(serverId); diff --git a/tests/lib/ircClient.test.ts b/tests/lib/ircClient.test.ts index ec7b0a57..6e0e98db 100644 --- a/tests/lib/ircClient.test.ts +++ b/tests/lib/ircClient.test.ts @@ -166,6 +166,63 @@ describe("IRCClient", () => { }); }); + describe("disconnect", () => { + test("does not send QUIT while socket is still connecting", () => { + const mockSocket = new MockWebSocket("wss://irc.example.com:443"); + MockWebSocketSpy.mockReturnValue(mockSocket); + + client.connect( + "Test Server", + "irc.example.com", + 443, + "testuser", + undefined, + undefined, + undefined, + "server-1", + ); + + client.disconnect("server-1"); + + expect(mockSocket.sentMessages).toEqual([]); + expect(mockSocket.readyState).toBe(WebSocket.CLOSED); + }); + + test("does not start reconnection after an intentional disconnect", async () => { + vi.useFakeTimers(); + + const mockSocket = new MockWebSocket("wss://irc.example.com:443"); + MockWebSocketSpy.mockReturnValue(mockSocket); + + const states: string[] = []; + client.on("connectionStateChange", ({ connectionState }) => { + states.push(connectionState); + }); + + const connectionPromise = client.connect( + "Test Server", + "irc.example.com", + 443, + "testuser", + undefined, + undefined, + undefined, + "server-2", + ); + + mockSocket.simulateOpen(); + await connectionPromise; + + client.disconnect("server-2"); + vi.runOnlyPendingTimers(); + + expect(states).toEqual(["connected", "disconnected"]); + expect(MockWebSocketSpy).toHaveBeenCalledTimes(1); + + vi.useRealTimers(); + }); + }); + describe("message handling", () => { test("should handle PRIVMSG correctly", async () => { const mockSocket = new MockWebSocket("ws://irc.example.com:443"); From 0b981b651e4651c514e84d1c06f5e913ad973151 Mon Sep 17 00:00:00 2001 From: louzt <179385168+louzt@users.noreply.github.com> Date: Thu, 23 Apr 2026 20:27:22 -0600 Subject: [PATCH 2/5] docs(architecture): describe native transport backend --- ARCHITECTURE.md | 22 +++++++++++++++++++--- 1 file changed, 19 insertions(+), 3 deletions(-) diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index c9771a7d..1e97e624 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -1,7 +1,7 @@ # ObsidianIRC Architecture > **Modern IRC Client** - React + TypeScript + TailwindCSS + Tauri -> Next-generation IRC client supporting websockets only +> Next-generation IRC client with WebSocket support in web builds and native TCP/TLS transport in Tauri builds ## 🏗️ Project Structure @@ -75,15 +75,31 @@ interface AppState { - Optimistic UI updates ### IRC Protocol Layer -**Location:** `src/lib/ircClient.ts` +**Location:** `src/lib/irc/IRCClient.ts` Event-driven IRC client supporting: -- WebSocket-only connections (no raw TCP) +- WebSocket connections for browser-compatible paths +- Native TCP/TLS connections through the Tauri backend on desktop builds - SASL authentication - IRC v3 message tags - Capability negotiation - Multi-server management +### Transport Layer +**Locations:** `src/lib/socket.ts`, `src-tauri/src/socket.rs` + +ObsidianIRC has a split transport model: + +- `wss://` routes through `WebSocketWrapper` in the frontend. +- `irc://` and `ircs://` route through `TCPSocket` in the frontend. +- `TCPSocket` bridges to Tauri commands (`connect`, `listen`, `send`, `disconnect`) implemented in Rust. +- The Rust backend owns the real TCP/TLS socket lifecycle, reads and writes IRC lines, and emits `tcp-message` events back to the frontend. + +This means transport hardening work usually belongs in two places: + +- frontend lifecycle and state transitions in `IRCClient` / `socket.ts` +- native socket behavior and shutdown semantics in `src-tauri/src/socket.rs` + **Event System:** ```typescript interface EventMap { From 03005a142d26864dd7f649ed8f206fc91b050b1b Mon Sep 17 00:00:00 2001 From: louzt <179385168+louzt@users.noreply.github.com> Date: Thu, 23 Apr 2026 20:47:37 -0600 Subject: [PATCH 3/5] harden socket lifecycle and disconnect intent --- src/lib/irc/IRCClient.ts | 6 +- src/lib/socket.ts | 73 ++++++++++++++++---- tests/lib/ircClient.test.ts | 39 +++++++++++ tests/lib/socket.test.ts | 129 ++++++++++++++++++++++++++++++++++++ 4 files changed, 231 insertions(+), 16 deletions(-) create mode 100644 tests/lib/socket.test.ts diff --git a/src/lib/irc/IRCClient.ts b/src/lib/irc/IRCClient.ts index 6860a605..f0f5dc43 100644 --- a/src/lib/irc/IRCClient.ts +++ b/src/lib/irc/IRCClient.ts @@ -619,7 +619,6 @@ export class IRCClient implements IRCClientContext { status: "online", }); this.nicks.set(server.id, nickname); - this.intentionalDisconnects.delete(server.id); socket.onopen = () => { //registerAllProtocolHandlers(this); @@ -656,6 +655,7 @@ export class IRCClient implements IRCClientContext { socket.onclose = () => { if (this.intentionalDisconnects.has(server.id)) { + this.intentionalDisconnects.delete(server.id); this.stopWebSocketPing(server.id); this.sockets.delete(server.id); this.pendingConnections.delete(connectionKey); @@ -731,9 +731,7 @@ export class IRCClient implements IRCClientContext { const CONNECTING = 0; const OPEN = 1; - if (socket.readyState === CONNECTING || socket.readyState === OPEN) { - this.intentionalDisconnects.add(serverId); - } + this.intentionalDisconnects.add(serverId); if (socket.readyState === OPEN && server?.isConnected) { const message = diff --git a/src/lib/socket.ts b/src/lib/socket.ts index aedc8110..582e028f 100644 --- a/src/lib/socket.ts +++ b/src/lib/socket.ts @@ -19,6 +19,8 @@ export class TCPSocket implements ISocket { private isConnected = false; private _readyState = 0; // 0: CONNECTING, 1: OPEN, 2: CLOSING, 3: CLOSED private unlisten?: () => void; + private closeRequested = false; + private closeFinalized = false; public onopen: (() => void) | null = null; public onmessage: ((event: { data: string }) => void) | null = null; @@ -43,6 +45,8 @@ export class TCPSocket implements ISocket { // Only handle messages for this client if (payload.id !== this.clientId) return; + if (this.closeFinalized) return; + if (payload.event.message) { // Convert byte array to string const data = new TextDecoder().decode( @@ -56,14 +60,15 @@ export class TCPSocket implements ISocket { } if (payload.event.connected === false) { - this.isConnected = false; - this._readyState = 3; // CLOSED - this.onclose?.(); - this.unlisten?.(); - this.unlisten = undefined; + this.finalizeClose(); } }) .then((unlistenFn) => { + if (this.closeFinalized) { + unlistenFn(); + return; + } + this.unlisten = unlistenFn; }) .catch((error: unknown) => { @@ -74,6 +79,23 @@ export class TCPSocket implements ISocket { invoke("connect", { clientId: this.clientId, address }) .then(() => { + if (this.closeFinalized) { + return; + } + + if (this.closeRequested) { + this.isConnected = true; + + return invoke("disconnect", { clientId: this.clientId }) + .then(() => { + this.finalizeClose(); + }) + .catch((error: unknown) => { + this.onerror?.(new Error(`Failed to disconnect: ${error}`)); + this.finalizeClose(); + }); + } + this.isConnected = true; this._readyState = 1; // OPEN // Start listening for messages @@ -82,6 +104,11 @@ export class TCPSocket implements ISocket { this.onopen?.(); }) .catch((error: unknown) => { + if (this.closeRequested) { + this.finalizeClose(); + return; + } + this._readyState = 3; // CLOSED this.onerror?.(new Error(`Failed to connect: ${error}`)); }); @@ -92,7 +119,7 @@ export class TCPSocket implements ISocket { } send(data: string): void { - if (!this.isConnected) { + if (!this.isConnected || this._readyState !== 1) { throw new Error("Socket is not connected"); } @@ -104,21 +131,43 @@ export class TCPSocket implements ISocket { } close(): void { - if (this.isConnected) { + if (this.closeRequested || this.closeFinalized) { + return; + } + + this.closeRequested = true; + + if (this._readyState !== 3) { this._readyState = 2; // CLOSING + } + + if (this.isConnected) { + this.isConnected = false; + invoke("disconnect", { clientId: this.clientId }) .then(() => { - this.isConnected = false; - this._readyState = 3; // CLOSED - this.onclose?.(); - this.unlisten?.(); - this.unlisten = undefined; + this.finalizeClose(); }) .catch((error: unknown) => { this.onerror?.(new Error(`Failed to disconnect: ${error}`)); + this.finalizeClose(); }); } } + + private finalizeClose(): void { + if (this.closeFinalized) { + return; + } + + this.closeFinalized = true; + this.closeRequested = false; + this.isConnected = false; + this._readyState = 3; + this.onclose?.(); + this.unlisten?.(); + this.unlisten = undefined; + } } export class WebSocketWrapper implements ISocket { diff --git a/tests/lib/ircClient.test.ts b/tests/lib/ircClient.test.ts index 6e0e98db..1809a2a3 100644 --- a/tests/lib/ircClient.test.ts +++ b/tests/lib/ircClient.test.ts @@ -221,6 +221,45 @@ describe("IRCClient", () => { vi.useRealTimers(); }); + + test("treats closing sockets as intentional disconnects", async () => { + vi.useFakeTimers(); + + const mockSocket = new MockWebSocket("wss://irc.example.com:443"); + MockWebSocketSpy.mockReturnValue(mockSocket); + + const states: string[] = []; + client.on("connectionStateChange", ({ connectionState }) => { + states.push(connectionState); + }); + + const connectionPromise = client.connect( + "Test Server", + "irc.example.com", + 443, + "testuser", + undefined, + undefined, + undefined, + "server-3", + ); + + mockSocket.simulateOpen(); + await connectionPromise; + + mockSocket.readyState = WebSocket.CLOSING; + client.disconnect("server-3"); + mockSocket.onclose?.(new CloseEvent("close")); + vi.runOnlyPendingTimers(); + + expect(states).toEqual(["connected", "disconnected"]); + expect(mockSocket.sentMessages).not.toContain( + "QUIT :ObsidianIRC - Bringing IRC to the future", + ); + expect(MockWebSocketSpy).toHaveBeenCalledTimes(1); + + vi.useRealTimers(); + }); }); describe("message handling", () => { diff --git a/tests/lib/socket.test.ts b/tests/lib/socket.test.ts new file mode 100644 index 00000000..123b53cd --- /dev/null +++ b/tests/lib/socket.test.ts @@ -0,0 +1,129 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; + +const { invokeMock, listenMock } = vi.hoisted(() => ({ + invokeMock: vi.fn(), + listenMock: vi.fn(), +})); + +vi.mock("@tauri-apps/api/core", () => ({ + invoke: invokeMock, +})); + +vi.mock("@tauri-apps/api/event", () => ({ + listen: listenMock, +})); + +import { TCPSocket } from "../../src/lib/socket"; + +function createDeferred() { + let resolve!: (value: T | PromiseLike) => void; + let reject!: (reason?: unknown) => void; + const promise = new Promise((res, rej) => { + resolve = res; + reject = rej; + }); + + return { promise, resolve, reject }; +} + +async function flushMicrotasks(iterations = 5): Promise { + for (let index = 0; index < iterations; index += 1) { + await Promise.resolve(); + } +} + +function getClientId(socket: TCPSocket): string { + return (socket as unknown as { clientId: string }).clientId; +} + +describe("TCPSocket", () => { + let connectDeferred: ReturnType>; + let unlistenMock: ReturnType; + let tcpMessageListener: + | ((event: { + payload: { + id: string; + event: { + message?: { data: number[] }; + error?: string; + connected?: boolean; + }; + }; + }) => void) + | undefined; + + beforeEach(() => { + connectDeferred = createDeferred(); + unlistenMock = vi.fn(); + tcpMessageListener = undefined; + + listenMock.mockImplementation(async (_eventName, handler) => { + tcpMessageListener = handler; + return unlistenMock; + }); + + invokeMock.mockImplementation((command: string, payload?: unknown) => { + switch (command) { + case "connect": + return connectDeferred.promise; + case "listen": + return Promise.resolve(); + case "disconnect": + return Promise.resolve(payload); + case "send": + return Promise.resolve(payload); + default: + return Promise.resolve(); + } + }); + }); + + afterEach(() => { + vi.clearAllMocks(); + }); + + test("close during connect suppresses a later onopen", async () => { + const socket = new TCPSocket("irc://irc.example.com:6667"); + const onopen = vi.fn(); + const onclose = vi.fn(); + socket.onopen = onopen; + socket.onclose = onclose; + + socket.close(); + + expect(socket.readyState).toBe(2); + + connectDeferred.resolve(); + await flushMicrotasks(); + + expect(invokeMock).toHaveBeenCalledWith("disconnect", { + clientId: getClientId(socket), + }); + expect(onopen).not.toHaveBeenCalled(); + expect(onclose).toHaveBeenCalledTimes(1); + expect(socket.readyState).toBe(3); + }); + + test("local close only emits onclose once when backend close arrives late", async () => { + const socket = new TCPSocket("irc://irc.example.com:6667"); + const onclose = vi.fn(); + socket.onclose = onclose; + + connectDeferred.resolve(); + await flushMicrotasks(); + + socket.close(); + await flushMicrotasks(); + + tcpMessageListener?.({ + payload: { + id: getClientId(socket), + event: { connected: false }, + }, + }); + + expect(onclose).toHaveBeenCalledTimes(1); + expect(unlistenMock).toHaveBeenCalledTimes(1); + expect(socket.readyState).toBe(3); + }); +}); From 54591bb4caad0b5552dc078c352f6be9d12dc878 Mon Sep 17 00:00:00 2001 From: louzt <179385168+louzt@users.noreply.github.com> Date: Thu, 23 Apr 2026 21:04:53 -0600 Subject: [PATCH 4/5] refactor: harden Rust bridge idempotency and socket seam --- src-tauri/src/socket.rs | 62 +++++++++++++++++++++++------- src/lib/socket.ts | 47 +++++++++++++++++++++-- tests/lib/socket.test.ts | 82 +++++++++++++++++++++++++++++++++++++++- 3 files changed, 173 insertions(+), 18 deletions(-) diff --git a/src-tauri/src/socket.rs b/src-tauri/src/socket.rs index f39ad5c8..b184522c 100644 --- a/src-tauri/src/socket.rs +++ b/src-tauri/src/socket.rs @@ -51,6 +51,14 @@ struct MessageData { data: Vec, } +async fn take_connection( + state: &Arc>>, + client_id: &str, +) -> Option { + let mut connections = state.lock().await; + connections.remove(client_id) +} + /// Read task for handling incoming data from the socket async fn read_task( client_id: String, @@ -66,6 +74,11 @@ async fn read_task( loop { match reader.read(&mut read_buf).await { Ok(0) => { + let was_tracked = take_connection(&state, &client_id).await.is_some(); + if !was_tracked { + break; + } + // Connection closed by server // Emit any remaining partial data as a final message if !line_buffer.is_empty() { @@ -87,10 +100,6 @@ async fn read_task( connected: Some(false), }, }); - - // Remove connection from state - let mut connections = state.lock().await; - connections.remove(&client_id); break; } Ok(n) => { @@ -122,6 +131,11 @@ async fn read_task( } } Err(e) => { + let was_tracked = take_connection(&state, &client_id).await.is_some(); + if !was_tracked { + break; + } + // Read error - emit error event and stop let _ = app_handle.emit("tcp-message", ReceivedPayload { id: client_id.clone(), @@ -131,10 +145,6 @@ async fn read_task( connected: Some(false), }, }); - - // Remove connection from state - let mut connections = state.lock().await; - connections.remove(&client_id); break; } } @@ -340,16 +350,14 @@ fn parse_host_port(host_port: &str, default_port: u16) -> Result<(String, u16), /// Disconnect a specific client connection #[tauri::command] pub async fn disconnect(client_id: String, state: State<'_, SocketState>) -> Result<(), String> { - let mut connections = state.0.lock().await; - if let Some(mut handle) = connections.remove(&client_id) { + if let Some(mut handle) = take_connection(&state.0, &client_id).await { // Send shutdown signal if available if let Some(shutdown_tx) = handle.shutdown_tx.take() { let _ = shutdown_tx.send(()); } - Ok(()) - } else { - Err(format!("No connection found for client_id: {}", client_id)) } + + Ok(()) } /// Start listening for messages from all active connections @@ -388,3 +396,31 @@ pub async fn send( Err(format!("No connection found for client_id: {}", client_id)) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn make_connection_handle() -> ConnectionHandle { + let (write_tx, _write_rx) = mpsc::channel(1); + let (shutdown_tx, _shutdown_rx) = oneshot::channel(); + + ConnectionHandle { + write_tx, + shutdown_tx: Some(shutdown_tx), + } + } + + #[tokio::test] + async fn take_connection_is_idempotent() { + let state = Arc::new(Mutex::new(HashMap::new())); + + { + let mut connections = state.lock().await; + connections.insert("client-1".to_string(), make_connection_handle()); + } + + assert!(take_connection(&state, "client-1").await.is_some()); + assert!(take_connection(&state, "client-1").await.is_none()); + } +} diff --git a/src/lib/socket.ts b/src/lib/socket.ts index 582e028f..efb49116 100644 --- a/src/lib/socket.ts +++ b/src/lib/socket.ts @@ -3,6 +3,8 @@ import { invoke } from "@tauri-apps/api/core"; import { listen } from "@tauri-apps/api/event"; +type SocketProtocol = "wss" | "irc" | "ircs"; + export interface ISocket { onopen: (() => void) | null; onmessage: ((event: { data: string }) => void) | null; @@ -14,6 +16,13 @@ export interface ISocket { readyState: number; } +export interface SocketFactoryContext { + url: string; + protocol: SocketProtocol; +} + +export type SocketFactory = (context: SocketFactoryContext) => ISocket; + export class TCPSocket implements ISocket { private clientId: string; private isConnected = false; @@ -210,12 +219,42 @@ export class WebSocketWrapper implements ISocket { } } -export function createSocket(url: string): ISocket { +export function resolveSocketProtocol(url: string): SocketProtocol { if (url.startsWith("wss://")) { - return new WebSocketWrapper(url); + return "wss"; } - if (url.startsWith("irc://") || url.startsWith("ircs://")) { - return new TCPSocket(url); + + if (url.startsWith("irc://")) { + return "irc"; + } + + if (url.startsWith("ircs://")) { + return "ircs"; } + throw new Error("Unsupported socket protocol"); } + +const defaultSocketFactory: SocketFactory = ({ url, protocol }) => { + if (protocol === "wss") { + return new WebSocketWrapper(url); + } + + return new TCPSocket(url); +}; + +let socketFactory: SocketFactory = defaultSocketFactory; + +export function setSocketFactory(factory: SocketFactory): void { + socketFactory = factory; +} + +export function resetSocketFactory(): void { + socketFactory = defaultSocketFactory; +} + +export function createSocket(url: string): ISocket { + const protocol = resolveSocketProtocol(url); + + return socketFactory({ url, protocol }); +} diff --git a/tests/lib/socket.test.ts b/tests/lib/socket.test.ts index 123b53cd..6f40dddb 100644 --- a/tests/lib/socket.test.ts +++ b/tests/lib/socket.test.ts @@ -13,7 +13,13 @@ vi.mock("@tauri-apps/api/event", () => ({ listen: listenMock, })); -import { TCPSocket } from "../../src/lib/socket"; +import { + createSocket, + resetSocketFactory, + resolveSocketProtocol, + setSocketFactory, + TCPSocket, +} from "../../src/lib/socket"; function createDeferred() { let resolve!: (value: T | PromiseLike) => void; @@ -80,6 +86,80 @@ describe("TCPSocket", () => { afterEach(() => { vi.clearAllMocks(); + resetSocketFactory(); + }); + + test("allows a custom socket factory to be injected", () => { + const fakeSocket = { + onopen: null, + onmessage: null, + onerror: null, + onclose: null, + send: vi.fn(), + close: vi.fn(), + readyState: 1, + }; + + const factory = vi.fn(() => fakeSocket); + setSocketFactory(factory); + + const socket = createSocket("ircs://irc.example.com:6697"); + + expect(factory).toHaveBeenCalledWith({ + url: "ircs://irc.example.com:6697", + protocol: "ircs", + }); + expect(socket).toBe(fakeSocket); + }); + + test("reset restores the default websocket routing", () => { + const sentinelFactory = vi.fn(() => { + throw new Error("custom factory should have been reset"); + }); + + class MockSocket extends EventTarget { + static readonly CONNECTING = 0; + static readonly OPEN = 1; + static readonly CLOSING = 2; + static readonly CLOSED = 3; + + readyState = MockSocket.CONNECTING; + onopen = null; + onmessage = null; + onerror = null; + onclose = null; + + constructor(public url: string) { + super(); + } + + send() {} + close() {} + } + + const originalWebSocket = globalThis.WebSocket; + globalThis.WebSocket = MockSocket as unknown as typeof WebSocket; + + setSocketFactory(sentinelFactory); + resetSocketFactory(); + + const socket = createSocket("wss://irc.example.com/webirc"); + + expect(sentinelFactory).not.toHaveBeenCalled(); + expect(socket.constructor.name).toBe("WebSocketWrapper"); + + if (originalWebSocket) { + globalThis.WebSocket = originalWebSocket; + } + }); + + test("resolveSocketProtocol preserves current supported routing", () => { + expect(resolveSocketProtocol("wss://irc.example.com/webirc")).toBe("wss"); + expect(resolveSocketProtocol("irc://irc.example.com:6667")).toBe("irc"); + expect(resolveSocketProtocol("ircs://irc.example.com:6697")).toBe("ircs"); + expect(() => resolveSocketProtocol("https://example.com/socket")).toThrow( + "Unsupported socket protocol", + ); }); test("close during connect suppresses a later onopen", async () => { From d1046286933aadf411cc4ecebcb3817b24394546 Mon Sep 17 00:00:00 2001 From: louzt <179385168+louzt@users.noreply.github.com> Date: Thu, 23 Apr 2026 21:20:32 -0600 Subject: [PATCH 5/5] refactor: centralize transport target resolution and extend tests --- src/lib/irc/IRCClient.ts | 48 ++++--------------------------------- src/lib/socket.ts | 44 ++++++++++++++++++++++++++++++++++ tests/lib/ircClient.test.ts | 24 +++++++++++++++++++ tests/lib/socket.test.ts | 34 ++++++++++++++++++++++++++ 4 files changed, 107 insertions(+), 43 deletions(-) diff --git a/src/lib/irc/IRCClient.ts b/src/lib/irc/IRCClient.ts index f0f5dc43..a1990c3e 100644 --- a/src/lib/irc/IRCClient.ts +++ b/src/lib/irc/IRCClient.ts @@ -11,9 +11,8 @@ import type { Server, User, } from "../../types"; -import { parseIrcUrl } from "../ircUrlParser"; import { parseMessageTags } from "../ircUtils"; -import { createSocket, type ISocket } from "../socket"; +import { createSocket, type ISocket, resolveSocketTarget } from "../socket"; import { IRC_DISPATCH } from "./handlers"; import type { IRCClientContext } from "./IRCClientContext"; @@ -514,47 +513,10 @@ export class IRCClient implements IRCClientContext { // Create a new connection promise and store it const connectionPromise = new Promise((resolve, reject) => { - let protocol: "wss" | "ircs" | "irc" = "wss"; - let actualHost = host; - let actualPort = port; - let actualPath = ""; - - if (host.startsWith("irc://") || host.startsWith("ircs://")) { - // Parse the IRC URL using centralized parser (Android-compatible) - const parsed = parseIrcUrl(host); - - protocol = parsed.scheme; - actualHost = parsed.host; - actualPort = parsed.port; - } else if (host.startsWith("wss://")) { - // Use URL constructor to preserve path/query (e.g. wss://host/websocket?token=...) - try { - const parsed = new URL(host); - actualHost = parsed.hostname; - actualPort = parsed.port ? Number.parseInt(parsed.port, 10) : port; - actualPath = - parsed.pathname !== "/" - ? parsed.pathname + parsed.search - : parsed.search; - } catch { - // malformed URL — leave actualHost/Port from the default - } - } else if (host.startsWith("ws://")) { - // Upgrade legacy ws:// to wss:// — unencrypted WebSockets are no longer supported - try { - const parsed = new URL(host); - actualHost = parsed.hostname; - actualPort = parsed.port ? Number.parseInt(parsed.port, 10) : port; - actualPath = - parsed.pathname !== "/" - ? parsed.pathname + parsed.search - : parsed.search; - } catch { - // malformed URL — leave actualHost/Port from the default - } - } - - const url = `${protocol}://${actualHost}:${actualPort}${actualPath}`; + const target = resolveSocketTarget(host, port); + const url = target.url; + const actualHost = target.host; + const actualPort = target.port; const socket = createSocket(url); // Create server object immediately and add to servers map diff --git a/src/lib/socket.ts b/src/lib/socket.ts index efb49116..89d1652d 100644 --- a/src/lib/socket.ts +++ b/src/lib/socket.ts @@ -2,9 +2,17 @@ import { invoke } from "@tauri-apps/api/core"; import { listen } from "@tauri-apps/api/event"; +import { parseIrcUrl } from "./ircUrlParser"; type SocketProtocol = "wss" | "irc" | "ircs"; +export interface ResolvedSocketTarget { + url: string; + host: string; + port: number; + protocol: SocketProtocol; +} + export interface ISocket { onopen: (() => void) | null; onmessage: ((event: { data: string }) => void) | null; @@ -235,6 +243,42 @@ export function resolveSocketProtocol(url: string): SocketProtocol { throw new Error("Unsupported socket protocol"); } +export function resolveSocketTarget( + rawHost: string, + port: number, +): ResolvedSocketTarget { + let protocol: SocketProtocol = "wss"; + let host = rawHost; + let resolvedPort = port; + let path = ""; + + if (rawHost.startsWith("irc://") || rawHost.startsWith("ircs://")) { + const parsed = parseIrcUrl(rawHost); + protocol = parsed.scheme; + host = parsed.host; + resolvedPort = parsed.port; + } else if (rawHost.startsWith("wss://") || rawHost.startsWith("ws://")) { + try { + const parsed = new URL(rawHost); + host = parsed.hostname; + resolvedPort = parsed.port ? Number.parseInt(parsed.port, 10) : port; + path = + parsed.pathname !== "/" + ? parsed.pathname + parsed.search + : parsed.search; + } catch { + // Keep the original host/port if the explicit websocket URL is malformed. + } + } + + return { + url: `${protocol}://${host}:${resolvedPort}${path}`, + host, + port: resolvedPort, + protocol, + }; +} + const defaultSocketFactory: SocketFactory = ({ url, protocol }) => { if (protocol === "wss") { return new WebSocketWrapper(url); diff --git a/tests/lib/ircClient.test.ts b/tests/lib/ircClient.test.ts index 1809a2a3..5348f300 100644 --- a/tests/lib/ircClient.test.ts +++ b/tests/lib/ircClient.test.ts @@ -102,6 +102,30 @@ describe("IRCClient", () => { expect(mockSocket.sentMessages).toContain("CAP LS 302"); }); + test("preserves secure websocket paths and query strings through the transport seam", async () => { + const mockSocket = new MockWebSocket( + "wss://irc.example.com:443/webirc?token=abc", + ); + MockWebSocketSpy.mockReturnValue(mockSocket); + + const connectionPromise = client.connect( + "WebIRC", + "wss://irc.example.com/webirc?token=abc", + 443, + "testuser", + ); + + mockSocket.simulateOpen(); + const server = await connectionPromise; + + expect(MockWebSocketSpy).toHaveBeenCalledWith( + "wss://irc.example.com:443/webirc?token=abc", + ); + expect(server.host).toBe("irc.example.com"); + expect(server.port).toBe(443); + expect(server.name).toBe("WebIRC"); + }); + test.skip("should handle connection errors", async () => { vi.useFakeTimers(); diff --git a/tests/lib/socket.test.ts b/tests/lib/socket.test.ts index 6f40dddb..da243520 100644 --- a/tests/lib/socket.test.ts +++ b/tests/lib/socket.test.ts @@ -17,6 +17,7 @@ import { createSocket, resetSocketFactory, resolveSocketProtocol, + resolveSocketTarget, setSocketFactory, TCPSocket, } from "../../src/lib/socket"; @@ -162,6 +163,39 @@ describe("TCPSocket", () => { ); }); + test("resolveSocketTarget preserves secure websocket paths and query strings", () => { + expect( + resolveSocketTarget("wss://irc.example.com/webirc?token=abc", 443), + ).toEqual({ + url: "wss://irc.example.com:443/webirc?token=abc", + host: "irc.example.com", + port: 443, + protocol: "wss", + }); + }); + + test("resolveSocketTarget upgrades ws urls to the secure websocket transport", () => { + expect(resolveSocketTarget("ws://irc.example.com/socket?x=1", 443)).toEqual( + { + url: "wss://irc.example.com:443/socket?x=1", + host: "irc.example.com", + port: 443, + protocol: "wss", + }, + ); + }); + + test("resolveSocketTarget keeps irc and ircs parsing in the transport seam", () => { + expect( + resolveSocketTarget("ircs://irc.libera.chat:6697/#chat", 443), + ).toEqual({ + url: "ircs://irc.libera.chat:6697", + host: "irc.libera.chat", + port: 6697, + protocol: "ircs", + }); + }); + test("close during connect suppresses a later onopen", async () => { const socket = new TCPSocket("irc://irc.example.com:6667"); const onopen = vi.fn();