import { createContext, type ReactNode, useCallback, useContext, useEffect, useMemo, useRef, useState, } from "react"; import { invoke, isTauri } from "@tauri-apps/api/core"; import { listen, type UnlistenFn } from "@tauri-apps/api/event"; import { MTPClient } from "mtp"; import { type z } from "zod"; import { ConnectionState } from "mtp"; import createAsyncQueue from "@tensamin/shared/asyncQueue"; import { toast as sonnerToast } from "@methanium/ui"; import { type Calls, type Communities, type Contacts, mtp as schemas, type MTP as Schemas, } from "@tensamin/shared/data"; import { log } from "@tensamin/shared/log"; import { ProtocolError } from "@tensamin/shared/errors"; import { useStorage } from "@tensamin/storage/context"; import { RECONNECT_RESET, RECONNECT_TRIES, RETRY_INTERVAL } from "./values"; function base64ToUint8Array(b64: string) { const bin = atob(b64); const out = new Uint8Array(bin.length); for (let i = 0; i < bin.length; i++) { out[i] = bin.charCodeAt(i); } return out; } export type ProtocolMessage< T extends keyof Schemas & string = keyof Schemas & string, > = { id?: number; type: T | string; data: z.infer; }; export type BoundSendFn = ( type: T, data?: z.infer, options?: { id?: number }, ) => Promise>; export type PushHandler = (message: ProtocolMessage) => void | Promise; const PUSH_TYPES = [ "MessageLive", "MessageEditLive", "MessageReactionLive", "MessageDeleteLive", "MessageState", "CallInvite", "GetStates", "ClientChanged", "ErrorNoIota", ] as const; export function isPushType(type: string): boolean { return (PUSH_TYPES as readonly string[]).includes(type); } function removeMissingContacts( contacts: Contacts, message: ProtocolMessage, ): Contacts { if (message.type !== "GetStates") return contacts; const data = message.data as { MissingUserIds?: unknown }; if (!Array.isArray(data.MissingUserIds)) return contacts; const missing = new Set( data.MissingUserIds.filter( (userId): userId is number => typeof userId === "number", ), ); return contacts.filter((contact) => !missing.has(contact.UserId)); } export type MTPExchange = { type: keyof Schemas & string; data: unknown; response: ProtocolMessage; }; export type MTPInterceptor = (exchange: MTPExchange) => void | Promise; type ContextType = { send: BoundSendFn; subscribe: ( type: T, handler: (message: ProtocolMessage) => void, ) => () => void; subscribePush: (handler: PushHandler) => () => void; addInterceptor: (interceptor: MTPInterceptor) => () => void; readyState: number; identified: boolean; freshContacts: Contacts; freshCommunities: Communities; freshCalls: Calls; contextReady: boolean; loadingDescription: string; }; const MTPContext = createContext(undefined); function getProtocolErrorDetails(error: unknown) { if (typeof error !== "object" || error === null || !("type" in error)) { return null; } const protocolError = error as { id?: unknown; type?: unknown; data?: unknown; }; return { id: protocolError.id, type: protocolError.type, data: protocolError.data, }; } // Zod schema validation export function validateResponse( type: T, message: { id?: number; type: string; data: unknown }, ): ProtocolMessage { if (message.type.startsWith("Error")) { return message as ProtocolMessage; } const schema = schemas[message.type as keyof Schemas & string]?.response ?? schemas[type]?.response; if (!schema) { return message as ProtocolMessage; } const parsed = schema.safeParse(message.data); if (!parsed.success) { throw new Error( `Response validation failed for ${type}: ${parsed.error.message}`, ); } return { id: message.id, type: message.type, data: parsed.data, } as ProtocolMessage; } function useMessageHandlers() { const interceptorsRef = useRef(new Set()); const pushHandlersRef = useRef(new Set()); const lastInitialStateRef = useRef(null); const subscribePush = useCallback((handler: PushHandler) => { pushHandlersRef.current.add(handler); const initialState = lastInitialStateRef.current; if (initialState?.type === "GetStates") { void Promise.resolve(handler(initialState)).catch(() => undefined); } return () => pushHandlersRef.current.delete(handler); }, []); const addInterceptor = useCallback((interceptor: MTPInterceptor) => { interceptorsRef.current.add(interceptor); return () => interceptorsRef.current.delete(interceptor); }, []); return { addInterceptor, interceptorsRef, lastInitialStateRef, pushHandlersRef, subscribePush, }; } function BrowserProvider(props: { children: ReactNode; blockConnection?: boolean; }) { const { load } = useStorage(); const [readyState, setReadyState] = useState( ConnectionState.Disconnected, ); const [identified, setIdentified] = useState(false); const [identifying, setIdentifying] = useState(false); const [freshCommunities, setFreshCommunities] = useState([]); const [freshContacts, setFreshContacts] = useState([]); const [freshCalls, setFreshCalls] = useState([]); const clientRef = useRef> | null>( null, ); const { addInterceptor, interceptorsRef, lastInitialStateRef, pushHandlersRef, subscribePush, } = useMessageHandlers(); const connected = readyState === ConnectionState.Connected; // MTP url const [mtpUrl, setMtpUrl] = useState(null); useEffect(() => { load("omega_url").then(setMtpUrl); }, [load]); // Validation override functions const send: BoundSendFn = useMemo( () => async (type, data, options) => { const client = clientRef.current; if (!client) { throw new Error("mtp is not connected"); } const message = await client.request( type, (data ?? {}) as Record, options, ); const response = validateResponse(type, message); setFreshContacts((contacts) => removeMissingContacts(contacts, response)); if (response.type.startsWith("Error")) { const errorData = response.data as Record; throw new ProtocolError({ type: response.type, requestId: response.id, errorType: typeof errorData.ErrorType === "string" ? errorData.ErrorType : undefined, }); } return response; }, [], ); const subscribe = useCallback((type, handler) => { const client = clientRef.current; if (!client) { return () => {}; } return client.subscribe(type, (message) => { handler(validateResponse(type, message)); }); }, []); // Reconnect stuff const resolveConnectionRef = useRef(() => {}); useEffect(() => { if (!mtpUrl) return; let attempts = 0; let reconnectTimer: ReturnType | null = null; let reconnectResetTimer: ReturnType | null = null; let reconnectScheduled = false; let disposed = false; let connectionGeneration = 0; const clearReconnectTimer = () => { if (!reconnectTimer) return; clearTimeout(reconnectTimer); reconnectTimer = null; reconnectScheduled = false; }; const clearReconnectResetTimer = () => { if (!reconnectResetTimer) return; clearTimeout(reconnectResetTimer); reconnectResetTimer = null; }; const scheduleReconnect = (error: unknown) => { if (disposed || reconnectScheduled) return; if (attempts >= RECONNECT_TRIES) { log(0, "mtp", "red", "Reconnection attempts exhausted", error); sonnerToast.error("Connection failed", { id: "mtp-connection-toast", description: error instanceof Error ? error.message.split(":")[0] : "Connection lost", icon: null, duration: Infinity, closeButton: true, promise: null, } as unknown as Parameters[1]); return; } attempts += 1; sonnerToast.loading( `Reconnecting to server... (attempt ${attempts} of ${RECONNECT_TRIES})`, { id: "mtp-connection-toast" }, ); reconnectScheduled = true; reconnectTimer = setTimeout(() => { reconnectScheduled = false; reconnectTimer = null; void connect(); }, RETRY_INTERVAL); }; async function connect() { if (disposed || props.blockConnection) return; const generation = ++connectionGeneration; let client: Awaited> | null = null; let failed = false; const cleanup = () => { client?.disconnect(); if (clientRef.current === client) { clientRef.current = null; } clearReconnectResetTimer(); if (generation === connectionGeneration) { setReadyState(ConnectionState.Disconnected); setIdentified(false); setIdentifying(false); } }; try { setIdentified(false); setIdentifying(false); const [userId, keyring] = await Promise.all([ load("user_id"), load("mtp_keyring"), ]); if (!userId || !keyring) { throw new Error("Missing login credentials"); } const forcedOmikronUrl = await load("forced_omikron_url"); const forcedOmikronPublicKey = await load("forced_omikron_public_key"); let url = null; let omikronPublicKey = null; if (forcedOmikronUrl && forcedOmikronPublicKey) { url = forcedOmikronUrl; omikronPublicKey = forcedOmikronPublicKey; } else { log(2, "mtp", "purple", "Fetching Omikron data."); const data = await fetch(`${mtpUrl}api/get/omikron/${userId}`); if (data.status === 404) { sonnerToast.error("We couldn't reach your Iota", { description: "Check your network connection and try restarting your Iota", icon: null, duration: Infinity, closeButton: true, }); resolveConnectionRef.current?.(); cleanup(); return; } const omikronData = (await data.json()) as { id: number; ip_address: string; port: number; public_key: string; status: string; }; if ( !omikronData.ip_address || !omikronData.port || !omikronData.public_key ) throw new Error("Invalid Omikron data"); url = `https://${omikronData.ip_address}:${omikronData.port}`; omikronPublicKey = omikronData.public_key; } //codec.decode(new Uint8Array(await res.arrayBuffer())), if (!url || !omikronPublicKey) throw new Error("Missing Omikron URL or Public Key"); log(2, "mtp", "green", "Connecting to: " + url); client = await MTPClient.create({ url, credentials: { clientId: userId, keyring: base64ToUint8Array(keyring), }, hostPublicKey: omikronPublicKey, descriptor: "client", pings: true, logger: (event) => { if (event.type === "state") { if (generation !== connectionGeneration) return; const state = client?.state ?? ConnectionState.Disconnected; setReadyState(state); if ( state === ConnectionState.Disconnected && clientRef.current === client && !failed ) { failed = true; clientRef.current = null; setIdentified(false); setIdentifying(false); scheduleReconnect(new Error("MTP connection lost")); } } if (event.type !== "Pong" && event.type !== "Ping") { log( 2, "mtp", event.type === "state" ? "purple" : event.direction === "recv" ? "cyan" : event.direction === "send" ? "gray" : "blue", event.type === "state" ? event.data : event.direction === "recv" ? "< " + event.type : event.direction === "send" ? "> " + event.type : event.type, event, ); } }, }); if (disposed || generation !== connectionGeneration) { client.disconnect(); return; } const activeClient = client; clientRef.current = activeClient; for (const type of PUSH_TYPES) { activeClient.subscribe(type, (message) => { let validated: ProtocolMessage; try { validated = validateResponse(type, message); } catch (error) { log(1, "mtp", "red", "Failed to validate push message", error, { type, data: message.data, }); return; } setFreshContacts((contacts) => removeMissingContacts(contacts, validated), ); for (const handler of [...pushHandlersRef.current]) { void Promise.resolve() .then(() => handler(validated)) .catch((error) => { log(1, "mtp", "red", "Push handler failed", error, { type }); }); } if (validated.type === "GetStates") { lastInitialStateRef.current = validated; } }); } setReadyState(activeClient.state); clearReconnectTimer(); // Schedule reconnect reset clearReconnectResetTimer(); reconnectResetTimer = setTimeout(() => { attempts = 0; reconnectResetTimer = null; }, RECONNECT_RESET * 1_000); setReadyState(activeClient.state); setIdentifying(true); const stateSync = new Promise>( (resolve, reject) => { let unsubscribeStateSync = () => {}; let unsubscribeNoIota = () => {}; const cleanupStateSync = () => { clearTimeout(timeout); unsubscribeStateSync(); unsubscribeNoIota(); }; const timeout = setTimeout(() => { cleanupStateSync(); reject(new Error("Initial state synchronization timed out")); }, 120_000); unsubscribeStateSync = activeClient.subscribe( "ClientStateSync", (message) => { cleanupStateSync(); try { resolve(validateResponse("ClientStateSync", message)); } catch (error) { reject(error); } }, ); unsubscribeNoIota = activeClient.subscribe("ErrorNoIota", () => { cleanupStateSync(); reject(new Error("No Iota is currently connected")); }); }, ); const [, finalResponse] = await Promise.all([ activeClient.auth(), stateSync, ]); if (finalResponse.type.startsWith("Error")) { throw new Error( `State synchronization failed: ${finalResponse.type}`, ); } const acknowledgement = await activeClient.request("ClientStateAck", { SessionId: finalResponse.data.SessionId, VersionNumber: finalResponse.data.VersionNumber, }); if (acknowledgement.type.startsWith("Error")) { throw new Error( `State acknowledgement failed: ${acknowledgement.type}`, ); } if (disposed || clientRef.current !== activeClient) return; setFreshContacts(finalResponse.data.Contacts); setFreshCommunities(finalResponse.data.Communities); setFreshCalls(finalResponse.data.Calls); setIdentifying(false); setIdentified(true); resolveConnectionRef.current?.(); } catch (connectError) { if (disposed || generation !== connectionGeneration) { client?.disconnect(); return; } failed = true; cleanup(); const connectErrorMessage = connectError instanceof Error ? connectError.message : String(connectError ?? "Unknown error"); log( 0, "mtp", "red", `Connection/authentication attempt failed: ${connectErrorMessage}`, getProtocolErrorDetails(connectError) ?? connectError, ); scheduleReconnect(connectError); } } void connect(); return () => { disposed = true; clearReconnectTimer(); clearReconnectResetTimer(); clientRef.current?.disconnect(); clientRef.current = null; setReadyState(ConnectionState.Disconnected); setIdentified(false); setIdentifying(false); sonnerToast.dismiss("mtp-connection-toast"); }; }, [ lastInitialStateRef, mtpUrl, props.blockConnection, load, pushHandlersRef, ]); // No Iota check useEffect(() => { if (!connected) return; return subscribe("ErrorNoIota", () => { setIdentified(false); setIdentifying(false); sonnerToast.error("We couldn't reach your Iota", { description: "Check your network connection and try restarting your Iota", icon: null, duration: Infinity, closeButton: true, }); resolveConnectionRef.current?.(); }); }, [connected, subscribe]); // Async queue const loadingDescription = useMemo(() => { if (!mtpUrl) return "Loading connection details"; if (readyState === ConnectionState.Connecting || !connected) { return "Establishing transport channel"; } if (identifying || !identified) return "Waiting for authenticated session"; return "Loading..."; }, [connected, identified, identifying, readyState, mtpUrl]); const contextReady = connected && identified && mtpUrl !== null; const mtpRef = useMemo( () => createAsyncQueue<{ send: typeof send; subscribe: typeof subscribe; subscribePush: typeof subscribePush; }>(), [], ); useEffect(() => { if (connected && identified && mtpUrl) { mtpRef.set({ send, subscribe, subscribePush, }); } }, [connected, identified, mtpUrl, send, subscribe, subscribePush, mtpRef]); const sendQueued: BoundSendFn = useMemo( () => async (type, data, options) => { const mtp = await mtpRef.get(); const response = await mtp.send(type, data, options); for (const interceptor of interceptorsRef.current) { void Promise.resolve( interceptor({ type, data, response: response as ProtocolMessage }), ).catch((error) => { log(1, "mtp", "yellow", "MTP interceptor failed", error, { type }); }); } return response; }, [interceptorsRef, mtpRef], ); return ( {props.children} ); } type NativeSnapshot = { generation: number; readyState: number; identified: boolean; state?: unknown; error?: string; }; function TauriProvider(props: { children: ReactNode; blockConnection?: boolean; }) { const [snapshot, setSnapshot] = useState({ generation: 0, readyState: ConnectionState.Disconnected, identified: false, }); const [freshContacts, setFreshContacts] = useState([]); const [freshCommunities, setFreshCommunities] = useState([]); const [freshCalls, setFreshCalls] = useState([]); const generationRef = useRef(0); const { addInterceptor, interceptorsRef, lastInitialStateRef, pushHandlersRef, subscribePush, } = useMessageHandlers(); const subscriptionsRef = useRef( new Map void>>(), ); const applySnapshot = useCallback((next: NativeSnapshot) => { if (next.generation < generationRef.current) return; generationRef.current = next.generation; if (next.error) { log(0, "android", "orange", "MTP connection failed", next.error); } setSnapshot(next); if (!next.identified || next.state === undefined) return; const parsed = schemas.ClientStateSync.response.safeParse(next.state); if (!parsed.success) { log(0, "mtp", "red", "Invalid native MTP state", parsed.error); return; } setFreshContacts(parsed.data.Contacts); setFreshCommunities(parsed.data.Communities); setFreshCalls(parsed.data.Calls); }, []); const dispatchMessage = useCallback( (raw: unknown) => { if (!raw || typeof raw !== "object" || !("type" in raw)) return; const message = raw as { id?: number; type: string; data: unknown }; let validated: ProtocolMessage; try { validated = validateResponse( message.type as keyof Schemas & string, message, ); } catch (error) { log(1, "mtp", "red", "Failed to validate native MTP message", error); return; } for (const handler of subscriptionsRef.current.get(validated.type) ?? []) { handler(validated); } if (!isPushType(validated.type)) return; setFreshContacts((contacts) => removeMissingContacts(contacts, validated), ); for (const handler of [...pushHandlersRef.current]) { void Promise.resolve(handler(validated)).catch((error) => { log(1, "mtp", "red", "Native MTP push handler failed", error, { type: validated.type, }); }); } if (validated.type === "GetStates") { lastInitialStateRef.current = validated; } }, [lastInitialStateRef, pushHandlersRef], ); useEffect(() => { if (props.blockConnection) return; let disposed = false; let unlisten: UnlistenFn | undefined; void (async () => { try { const nextUnlisten = await listen< | { kind: "state"; snapshot: NativeSnapshot } | { kind: "message"; generation: number; message: unknown } | { kind: "log"; level: number; message: string; details?: unknown; } >("mtp://event", ({ payload }) => { if (disposed) return; if (payload.kind === "state") { applySnapshot(payload.snapshot); return; } if (payload.kind === "message") { if (payload.generation === generationRef.current) { dispatchMessage(payload.message); } return; } log( payload.level, "android", "orange", payload.message, payload.details, ); }); if (disposed) nextUnlisten(); else unlisten = nextUnlisten; } catch (error) { log(0, "mtp", "red", "Failed to subscribe to native MTP events", error); } try { const current = await invoke("mtp_status"); if (!disposed) applySnapshot(current); } catch (error) { log(0, "mtp", "red", "Failed to load native MTP status", error); } })(); return () => { disposed = true; unlisten?.(); }; }, [applySnapshot, dispatchMessage, props.blockConnection]); useEffect(() => { if (props.blockConnection) return; const updateVisibility = () => { void invoke("mtp_set_ui_visible", { visible: document.visibilityState === "visible" && document.hasFocus(), }); }; updateVisibility(); document.addEventListener("visibilitychange", updateVisibility); window.addEventListener("focus", updateVisibility); window.addEventListener("blur", updateVisibility); return () => { document.removeEventListener("visibilitychange", updateVisibility); window.removeEventListener("focus", updateVisibility); window.removeEventListener("blur", updateVisibility); void invoke("mtp_set_ui_visible", { visible: false }); }; }, [props.blockConnection]); const send = useCallback( async (type, data, options) => { const response = await invoke("mtp_request", { typeName: type, data: data ?? {}, id: options?.id, }); const validated = validateResponse(type, response); setFreshContacts((contacts) => removeMissingContacts(contacts, validated), ); if (validated.type.startsWith("Error")) { const errorData = validated.data as Record; throw new ProtocolError({ type: validated.type, requestId: validated.id, errorType: typeof errorData.ErrorType === "string" ? errorData.ErrorType : undefined, }); } for (const interceptor of interceptorsRef.current) { void Promise.resolve( interceptor({ type, data, response: validated as ProtocolMessage }), ).catch((error) => { log(1, "mtp", "yellow", "MTP interceptor failed", error, { type }); }); } return validated; }, [interceptorsRef], ); const subscribe = useCallback((type, handler) => { const handlers = subscriptionsRef.current.get(type) ?? new Set<(message: ProtocolMessage) => void>(); handlers.add(handler as (message: ProtocolMessage) => void); subscriptionsRef.current.set(type, handlers); return () => { handlers.delete(handler as (message: ProtocolMessage) => void); if (handlers.size === 0) subscriptionsRef.current.delete(type); }; }, []); const connected = snapshot.readyState === ConnectionState.Connected; const contextReady = connected && snapshot.identified; return ( {props.children} ); } export function Provider(props: { children: ReactNode; blockConnection?: boolean; }) { const [wasmReady, setWasmReady] = useState(false); const [wasmError, setWasmError] = useState(); useEffect(() => { let active = true; void MTPClient.init().then( () => { if (active) setWasmReady(true); }, (error: unknown) => { if (active) setWasmError(() => error); }, ); return () => { active = false; }; }, []); if (wasmError) throw wasmError; if (!wasmReady) return null; return isTauri() ? ( ) : ( ); } export function useMTP(): ContextType { const context = useContext(MTPContext); if (!context) { throw new Error("useMTP must be used within an MTPProvider"); } return context; }