Files
pi-mobile/apps/mobile/contexts/SocketContext.tsx
T
2026-07-09 13:08:49 -05:00

420 lines
12 KiB
TypeScript

import {
createContext,
useCallback,
useContext,
useEffect,
useMemo,
useRef,
useState,
type ReactNode,
} from "react";
import {
buildBaseUrl,
checkHealth,
loadSavedComputers,
loadSavedProjects,
loadLastModel,
removeSavedComputer,
saveHostConfig,
saveSavedProjects,
saveLastModel,
upsertSavedComputer,
} from "@/lib/api";
import { createSocket, listModels as requestModels, setRuntimeModel as requestSetRuntimeModel } from "@/lib/socket";
import type {
ChatEvent,
ChatHistoryPayload,
LiveRuntimeInfo,
ModelListPayload,
ModelRef,
PersistedSessionInfo,
RuntimeCreatedPayload,
SavedComputer,
} from "@/lib/types";
import type { Socket } from "socket.io-client";
type SocketContextValue = {
host: string;
port: string;
computerName: string;
baseUrl: string | null;
connected: boolean;
connecting: boolean;
error: string | null;
savedComputers: SavedComputer[];
persistedSessions: PersistedSessionInfo[];
liveRuntimes: LiveRuntimeInfo[];
savedProjects: string[];
addProject: (cwd: string) => Promise<void>;
removeProject: (cwd: string) => Promise<void>;
setHost: (host: string) => void;
setPort: (port: string) => void;
setComputerName: (name: string) => void;
connect: () => Promise<boolean>;
connectToSavedComputer: (computer: SavedComputer) => Promise<boolean>;
removeSavedComputerEntry: (id: string) => Promise<void>;
disconnect: () => void;
refreshSessions: () => void;
createRuntime: (cwd: string, name?: string, model?: ModelRef) => void;
lastModel: ModelRef | null;
listModels: (payload: { runtimeId?: string; cwd?: string }) => Promise<ModelListPayload>;
setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise<void>;
openRuntime: (sessionFile: string) => void;
attachRuntime: (runtimeId: string) => void;
sendPrompt: (runtimeId: string, message: string) => void;
abortPrompt: (runtimeId: string) => void;
subscribeChat: (listener: (event: ChatEvent) => void) => () => void;
subscribeChatHistory: (listener: (payload: ChatHistoryPayload) => void) => () => void;
waitForRuntime: (
kind: "created" | "opened",
timeoutMs?: number,
) => Promise<RuntimeCreatedPayload>;
};
const SocketContext = createContext<SocketContextValue | null>(null);
export function SocketProvider({ children }: { children: ReactNode }) {
const [host, setHost] = useState("");
const [port, setPort] = useState("8787");
const [computerName, setComputerName] = useState("");
const [baseUrl, setBaseUrl] = useState<string | null>(null);
const [connected, setConnected] = useState(false);
const [connecting, setConnecting] = useState(false);
const [error, setError] = useState<string | null>(null);
const [savedComputers, setSavedComputers] = useState<SavedComputer[]>([]);
const [persistedSessions, setPersistedSessions] = useState<PersistedSessionInfo[]>([]);
const [liveRuntimes, setLiveRuntimes] = useState<LiveRuntimeInfo[]>([]);
const [savedProjects, setSavedProjects] = useState<string[]>([]);
const [lastModel, setLastModel] = useState<ModelRef | null>(null);
const socketRef = useRef<Socket | null>(null);
const chatListenersRef = useRef(new Set<(event: ChatEvent) => void>());
const chatHistoryListenersRef = useRef(new Set<(payload: ChatHistoryPayload) => void>());
const runtimeWaitersRef = useRef<{
created: Array<(payload: RuntimeCreatedPayload) => void>;
opened: Array<(payload: RuntimeCreatedPayload) => void>;
}>({ created: [], opened: [] });
useEffect(() => {
void loadLastModel().then(setLastModel);
void loadSavedComputers().then((computers) => {
setSavedComputers(computers);
const recent = computers[0];
if (!recent) {
return;
}
setHost(recent.host);
setPort(recent.port);
setComputerName(recent.name);
void loadSavedProjects(recent.host, recent.port).then(setSavedProjects);
});
}, []);
const removeSavedComputerEntry = useCallback(
async (id: string) => {
const next = await removeSavedComputer(savedComputers, id);
setSavedComputers(next);
},
[savedComputers],
);
const addProject = useCallback(
async (cwd: string) => {
const trimmed = cwd.trim();
if (!trimmed) {
return;
}
setSavedProjects((current) => {
if (current.some((entry) => entry === trimmed)) {
return current;
}
const next = [...current, trimmed];
void saveSavedProjects(host, port, next);
return next;
});
},
[host, port],
);
const removeProject = useCallback(
async (cwd: string) => {
setSavedProjects((current) => {
const next = current.filter((entry) => entry !== cwd);
void saveSavedProjects(host, port, next);
return next;
});
},
[host, port],
);
const disconnect = useCallback(() => {
socketRef.current?.disconnect();
socketRef.current = null;
setConnected(false);
setBaseUrl(null);
}, []);
const connectWithHost = useCallback(
async (nextHost: string, nextPort: string, nextName?: string) => {
setConnecting(true);
setError(null);
try {
const nextBaseUrl = buildBaseUrl(nextHost, nextPort);
const healthy = await checkHealth(nextBaseUrl);
if (!healthy) {
throw new Error("Could not reach Pi Mobile API");
}
disconnect();
const socket = createSocket(nextBaseUrl, {
onConnect: () => {
setConnected(true);
setConnecting(false);
socket.emit("sessions:list");
},
onDisconnect: () => {
setConnected(false);
},
onError: (message) => setError(message),
onSessionsList: (payload) => {
setPersistedSessions(payload.persisted);
setLiveRuntimes(payload.live);
},
onRuntimeCreated: (payload) => {
for (const resolve of runtimeWaitersRef.current.created.splice(0)) {
resolve(payload);
}
},
onRuntimeOpened: (payload) => {
for (const resolve of runtimeWaitersRef.current.opened.splice(0)) {
resolve(payload);
}
},
onChatEvent: (payload) => {
for (const listener of chatListenersRef.current) {
listener(payload);
}
},
onChatHistory: (payload) => {
for (const listener of chatHistoryListenersRef.current) {
listener(payload);
}
},
});
socketRef.current = socket;
setBaseUrl(nextBaseUrl);
setHost(nextHost);
setPort(nextPort);
if (nextName !== undefined) {
setComputerName(nextName);
}
await saveHostConfig({ host: nextHost, port: nextPort });
const nextComputers = await upsertSavedComputer(
savedComputers,
nextHost,
nextPort,
nextName ?? computerName,
);
setSavedComputers(nextComputers);
const projects = await loadSavedProjects(nextHost, nextPort);
setSavedProjects(projects);
return true;
} catch (connectError) {
const message =
connectError instanceof Error ? connectError.message : "Failed to connect";
setError(message);
setConnecting(false);
return false;
}
},
[computerName, disconnect, savedComputers],
);
const connect = useCallback(async () => {
return connectWithHost(host, port, computerName);
}, [connectWithHost, host, port, computerName]);
const connectToSavedComputer = useCallback(
async (computer: SavedComputer) => {
return connectWithHost(computer.host, computer.port, computer.name);
},
[connectWithHost],
);
const refreshSessions = useCallback(() => {
socketRef.current?.emit("sessions:list");
}, []);
const createRuntime = useCallback((cwd: string, name?: string, model?: ModelRef) => {
socketRef.current?.emit("runtime:create", { cwd, name, model });
}, []);
const listModels = useCallback(
async (payload: { runtimeId?: string; cwd?: string }) => {
const socket = socketRef.current;
if (!socket) {
throw new Error("Not connected");
}
try {
return await requestModels(socket, payload);
} catch (requestError) {
const message = requestError instanceof Error ? requestError.message : "Failed to load models";
setError(message);
throw requestError;
}
},
[],
);
const setRuntimeModel = useCallback(async (runtimeId: string, model: ModelRef) => {
const socket = socketRef.current;
if (!socket) {
throw new Error("Not connected");
}
setLastModel(model);
try {
const runtime = await requestSetRuntimeModel(socket, runtimeId, model);
await saveLastModel(model);
setLiveRuntimes((current) =>
current.map((entry) => (entry.runtimeId === runtime.runtimeId ? runtime : entry)),
);
} catch (requestError) {
const message = requestError instanceof Error ? requestError.message : "Failed to set model";
setError(message);
throw requestError;
}
}, []);
const openRuntime = useCallback((sessionFile: string) => {
socketRef.current?.emit("runtime:open", { sessionFile });
}, []);
const attachRuntime = useCallback((runtimeId: string) => {
socketRef.current?.emit("runtime:attach", { runtimeId });
}, []);
const sendPrompt = useCallback((runtimeId: string, message: string) => {
socketRef.current?.emit("chat:prompt", { runtimeId, message });
}, []);
const abortPrompt = useCallback((runtimeId: string) => {
socketRef.current?.emit("chat:abort", { runtimeId });
}, []);
const subscribeChat = useCallback((listener: (event: ChatEvent) => void) => {
chatListenersRef.current.add(listener);
return () => {
chatListenersRef.current.delete(listener);
};
}, []);
const subscribeChatHistory = useCallback((listener: (payload: ChatHistoryPayload) => void) => {
chatHistoryListenersRef.current.add(listener);
return () => {
chatHistoryListenersRef.current.delete(listener);
};
}, []);
const waitForRuntime = useCallback(
(kind: "created" | "opened", timeoutMs = 15000) =>
new Promise<RuntimeCreatedPayload>((resolve, reject) => {
const timer = setTimeout(() => {
runtimeWaitersRef.current[kind] = runtimeWaitersRef.current[kind].filter(
(entry) => entry !== resolve,
);
reject(new Error("Timed out waiting for runtime"));
}, timeoutMs);
const wrappedResolve = (payload: RuntimeCreatedPayload) => {
clearTimeout(timer);
resolve(payload);
};
runtimeWaitersRef.current[kind].push(wrappedResolve);
}),
[],
);
const value = useMemo(
() => ({
host,
port,
computerName,
baseUrl,
connected,
connecting,
error,
savedComputers,
persistedSessions,
liveRuntimes,
savedProjects,
lastModel,
addProject,
removeProject,
setHost,
setPort,
setComputerName,
connect,
connectToSavedComputer,
removeSavedComputerEntry,
disconnect,
refreshSessions,
createRuntime,
listModels,
setRuntimeModel,
openRuntime,
attachRuntime,
sendPrompt,
abortPrompt,
subscribeChat,
subscribeChatHistory,
waitForRuntime,
}),
[
host,
port,
computerName,
baseUrl,
connected,
connecting,
error,
savedComputers,
persistedSessions,
liveRuntimes,
savedProjects,
lastModel,
addProject,
removeProject,
connect,
connectToSavedComputer,
removeSavedComputerEntry,
disconnect,
refreshSessions,
createRuntime,
listModels,
setRuntimeModel,
openRuntime,
attachRuntime,
sendPrompt,
abortPrompt,
subscribeChat,
subscribeChatHistory,
waitForRuntime,
],
);
return <SocketContext.Provider value={value}>{children}</SocketContext.Provider>;
}
export function useSocket() {
const context = useContext(SocketContext);
if (!context) {
throw new Error("useSocket must be used within SocketProvider");
}
return context;
}