From a6029da45bda7d008359ddf4ebf7cc5b6ccdab3d Mon Sep 17 00:00:00 2001 From: Owen Qwen Date: Thu, 9 Jul 2026 13:08:49 -0500 Subject: [PATCH] Adding a model selector --- apps/api/src/runtime-registry.ts | 89 +++++++++- apps/api/src/socket-handlers.ts | 52 +++++- apps/api/src/types.ts | 22 +++ apps/mobile/app/chat/[runtimeId].tsx | 130 ++++++++++++--- apps/mobile/components/ModelSelector.tsx | 200 +++++++++++++++++++++++ apps/mobile/contexts/SocketContext.tsx | 59 ++++++- apps/mobile/lib/api.ts | 22 ++- apps/mobile/lib/socket.ts | 50 ++++++ apps/mobile/lib/types.ts | 22 +++ 9 files changed, 611 insertions(+), 35 deletions(-) create mode 100644 apps/mobile/components/ModelSelector.tsx diff --git a/apps/api/src/runtime-registry.ts b/apps/api/src/runtime-registry.ts index 80565e1..1db80bf 100644 --- a/apps/api/src/runtime-registry.ts +++ b/apps/api/src/runtime-registry.ts @@ -3,6 +3,7 @@ import { type AgentSessionEvent, type AgentSessionRuntime, createAgentSessionRuntime, + createAgentSessionServices, getAgentDir, initTheme, SessionManager, @@ -10,7 +11,13 @@ import { import { lastMessagePreview } from "./message-history.js"; import { createRuntimeFactory } from "./pi-factory.js"; import { lastMessageFromSessionFile } from "./session-preview.js"; -import type { LiveRuntimeInfo, PersistedSessionInfo } from "./types.js"; +import type { + LiveRuntimeInfo, + ModelListPayload, + ModelRef, + ModelSummary, + PersistedSessionInfo, +} from "./types.js"; let themeReady = false; @@ -45,6 +52,7 @@ export class RuntimeRegistry { sessionFile: session.sessionFile, lastMessage: lastMessagePreview(session.messages), isStreaming: session.isStreaming, + currentModel: toModelRef(session.model), }; }); } @@ -68,7 +76,7 @@ export class RuntimeRegistry { return this.runtimes.get(runtimeId); } - async createRuntime(cwd: string, name?: string): Promise { + async createRuntime(cwd: string, name?: string, modelRef?: ModelRef): Promise { const sessionManager = SessionManager.create(cwd); const runtime = await createAgentSessionRuntime(createRuntimeFactory, { cwd, @@ -76,6 +84,10 @@ export class RuntimeRegistry { sessionManager, }); + if (modelRef) { + await this.applyModel(runtime, modelRef); + } + if (name) { runtime.session.sessionManager.appendSessionInfo(name); } @@ -121,6 +133,30 @@ export class RuntimeRegistry { await record.runtime.session.abort(); } + async listModels(runtimeId?: string, cwd?: string): Promise { + const record = runtimeId ? this.runtimes.get(runtimeId) : undefined; + if (runtimeId && !record) { + throw new Error(`Runtime not found: ${runtimeId}`); + } + + const registry = record + ? record.runtime.session.modelRegistry + : (await createAgentSessionServices({ cwd: cwd || process.cwd() })).modelRegistry; + return { + models: registry + .getAvailable() + .map((model) => toModelSummary(model, registry)), + currentModel: record ? toModelRef(record.runtime.session.model) : undefined, + loadError: registry.getError(), + }; + } + + async setModel(runtimeId: string, modelRef: ModelRef): Promise { + const record = this.requireRuntime(runtimeId); + await this.applyModel(record.runtime, modelRef); + return this.getLiveInfo(runtimeId, record); + } + dispose(runtimeId: string): void { const record = this.runtimes.get(runtimeId); if (!record) { @@ -152,9 +188,33 @@ export class RuntimeRegistry { sessionFile: session.sessionFile, lastMessage: lastMessagePreview(session.messages), isStreaming: session.isStreaming, + currentModel: toModelRef(session.model), }; } + private getLiveInfo(runtimeId: string, record: RuntimeRecord): LiveRuntimeInfo { + const session = record.runtime.session; + return { + runtimeId, + sessionId: session.sessionId, + cwd: record.cwd, + sessionName: session.sessionManager.getSessionName(), + sessionFile: session.sessionFile, + lastMessage: lastMessagePreview(session.messages), + isStreaming: session.isStreaming, + currentModel: toModelRef(session.model), + }; + } + + private async applyModel(runtime: AgentSessionRuntime, modelRef: ModelRef): Promise { + const registry = runtime.session.modelRegistry; + const model = registry.find(modelRef.provider, modelRef.id); + if (!model || !registry.getAvailable().some((entry) => entry.provider === modelRef.provider && entry.id === modelRef.id)) { + throw new Error(`Model unavailable: ${modelRef.provider}/${modelRef.id}`); + } + await runtime.session.setModel(model); + } + private bindSession(record: RuntimeRecord): () => void { record.unsubscribe?.(); @@ -182,3 +242,28 @@ export class RuntimeRegistry { return record; } } + +function toModelRef(model?: { provider: string; id: string }): ModelRef | undefined { + return model ? { provider: model.provider, id: model.id } : undefined; +} + +function toModelSummary(model: { + provider: string; + id: string; + name?: string; + reasoning: boolean; + input: readonly ("text" | "image")[]; + contextWindow: number; + maxTokens: number; +}, registry?: { getProviderDisplayName(provider: string): string }): ModelSummary { + return { + provider: model.provider, + id: model.id, + name: model.name, + providerDisplayName: registry?.getProviderDisplayName(model.provider), + reasoning: model.reasoning, + input: [...model.input], + contextWindow: model.contextWindow, + maxTokens: model.maxTokens, + }; +} diff --git a/apps/api/src/socket-handlers.ts b/apps/api/src/socket-handlers.ts index 0b31856..c075246 100644 --- a/apps/api/src/socket-handlers.ts +++ b/apps/api/src/socket-handlers.ts @@ -1,7 +1,13 @@ import type { Server, Socket } from "socket.io"; import { serializeMessages } from "./message-history.js"; import type { RuntimeRegistry } from "./runtime-registry.js"; -import type { ChatEvent, ChatHistoryPayload, RuntimeCreatedPayload } from "./types.js"; +import type { + ChatEvent, + ChatHistoryPayload, + ModelListPayload, + ModelRef, + RuntimeCreatedPayload, +} from "./types.js"; type AttachState = { runtimeId?: string; @@ -24,9 +30,43 @@ export function registerSocketHandlers(io: Server, registry: RuntimeRegistry) { } }); - socket.on("runtime:create", async (payload: { cwd: string; name?: string }) => { + socket.on( + "models:list", + async ( + payload: { runtimeId?: string; cwd?: string }, + ack?: (response: ModelListPayload | { error: string }) => void, + ) => { + try { + const result = await registry.listModels(payload.runtimeId, payload.cwd); + ack?.(result); + } catch (error) { + ack?.({ error: errorMessage(error) }); + } + }, + ); + + socket.on( + "runtime:model:set", + async ( + payload: { runtimeId: string; model: ModelRef }, + ack?: (response: { runtime: RuntimeCreatedPayload } | { error: string }) => void, + ) => { + try { + const runtime = await registry.setModel(payload.runtimeId, payload.model); + ack?.({ runtime }); + io.emit("sessions:list:result", { + persisted: await registry.listPersisted(), + live: registry.listLive(), + }); + } catch (error) { + ack?.({ error: errorMessage(error) }); + } + }, + ); + + socket.on("runtime:create", async (payload: { cwd: string; name?: string; model?: ModelRef }) => { try { - const info = await registry.createRuntime(payload.cwd, payload.name); + const info = await registry.createRuntime(payload.cwd, payload.name, payload.model); const created: RuntimeCreatedPayload = info; socket.emit("runtime:created", created); io.emit("sessions:list:result", { @@ -97,6 +137,10 @@ export function registerSocketHandlers(io: Server, registry: RuntimeRegistry) { } function emitError(socket: Socket, error: unknown) { - const message = error instanceof Error ? error.message : "Unknown error"; + const message = errorMessage(error); socket.emit("error", { message }); } + +function errorMessage(error: unknown): string { + return error instanceof Error ? error.message : "Unknown error"; +} diff --git a/apps/api/src/types.ts b/apps/api/src/types.ts index fff4bfd..e7ef156 100644 --- a/apps/api/src/types.ts +++ b/apps/api/src/types.ts @@ -1,5 +1,25 @@ import type { AgentSessionEvent } from "@earendil-works/pi-coding-agent"; +export type ModelRef = { + provider: string; + id: string; +}; + +export type ModelSummary = ModelRef & { + name?: string; + providerDisplayName?: string; + reasoning: boolean; + input: ("text" | "image")[]; + contextWindow: number; + maxTokens: number; +}; + +export type ModelListPayload = { + models: ModelSummary[]; + currentModel?: ModelRef; + loadError?: string; +}; + export type LiveRuntimeInfo = { runtimeId: string; sessionId: string; @@ -8,6 +28,7 @@ export type LiveRuntimeInfo = { sessionFile?: string; lastMessage?: string; isStreaming: boolean; + currentModel?: ModelRef; }; export type PersistedSessionInfo = { @@ -50,4 +71,5 @@ export type RuntimeCreatedPayload = { sessionName?: string; sessionFile?: string; lastMessage?: string; + currentModel?: ModelRef; }; diff --git a/apps/mobile/app/chat/[runtimeId].tsx b/apps/mobile/app/chat/[runtimeId].tsx index 9e7c318..c5e38b5 100644 --- a/apps/mobile/app/chat/[runtimeId].tsx +++ b/apps/mobile/app/chat/[runtimeId].tsx @@ -14,12 +14,14 @@ import { } from "react-native"; import { AssistantMessage } from "@/components/AssistantMessage"; import { MessageRow, type Sender } from "@/components/MessageRow"; +import { ModelSelector } from "@/components/ModelSelector"; import { ToolCallCard } from "@/components/ToolCallCard"; import { UserMessage } from "@/components/UserMessage"; import { useSocket } from "@/contexts/SocketContext"; +import { saveLastModel } from "@/lib/api"; import { extractResultText, toolArgsPreview, truncateDetail } from "@/lib/tool-format"; import { colors, monoFont } from "@/lib/theme"; -import type { ChatHistoryMessage, ChatItem } from "@/lib/types"; +import type { ChatHistoryMessage, ChatItem, ModelRef, ModelSummary } from "@/lib/types"; function extractText(content: unknown): string { if (typeof content === "string") { @@ -121,7 +123,10 @@ export default function ChatScreen() { const { attachRuntime, createRuntime, + lastModel, + listModels, sendPrompt, + setRuntimeModel, abortPrompt, subscribeChat, subscribeChatHistory, @@ -134,12 +139,57 @@ export default function ChatScreen() { const [draft, setDraft] = useState(""); const [streaming, setStreaming] = useState(false); const [creating, setCreating] = useState(false); + const [models, setModels] = useState([]); + const [selectedModel, setSelectedModel] = useState(); + const [modelLoading, setModelLoading] = useState(false); + const [modelError, setModelError] = useState(null); const assistantIdRef = useRef(null); const toolIdsRef = useRef(new Map()); const pendingMessageRef = useRef(null); const listRef = useRef>(null); const atBottomRef = useRef(true); + useEffect(() => { + let cancelled = false; + setModelLoading(true); + setModelError(null); + void listModels(isDraftChat ? { cwd: projectCwd } : { runtimeId: activeRuntimeId ?? undefined }) + .then((payload) => { + if (cancelled) { + return; + } + setModels(payload.models); + const persistedModel = lastModel + ? payload.models.some( + (model) => model.provider === lastModel.provider && model.id === lastModel.id, + ) + ? lastModel + : undefined + : undefined; + setSelectedModel( + payload.currentModel ?? + persistedModel ?? + (payload.models[0] ? { provider: payload.models[0].provider, id: payload.models[0].id } : undefined), + ); + if (payload.loadError) { + setModelError(payload.loadError); + } + }) + .catch((error: unknown) => { + if (!cancelled) { + setModelError(error instanceof Error ? error.message : "Failed to load models"); + } + }) + .finally(() => { + if (!cancelled) { + setModelLoading(false); + } + }); + return () => { + cancelled = true; + }; + }, [activeRuntimeId, isDraftChat, lastModel, listModels, projectCwd]); + useEffect(() => { if (!activeRuntimeId) { return; @@ -346,7 +396,7 @@ export default function ChatScreen() { pendingMessageRef.current = message; setCreating(true); - createRuntime(projectCwd); + createRuntime(projectCwd, undefined, selectedModel); try { const runtime = await waitForRuntime("created"); setActiveRuntimeId(runtime.runtimeId); @@ -363,6 +413,25 @@ export default function ChatScreen() { sendPrompt(activeRuntimeId, message); }; + const selectModel = async (model: ModelRef) => { + if (streaming) { + return; + } + const previous = selectedModel; + setSelectedModel(model); + setModelError(null); + if (!activeRuntimeId) { + void saveLastModel(model); + return; + } + try { + await setRuntimeModel(activeRuntimeId, model); + } catch (error: unknown) { + setSelectedModel(previous); + setModelError(error instanceof Error ? error.message : "Failed to set model"); + } + }; + const handleScroll = (scrollEvent: NativeSyntheticEvent) => { const { contentOffset, contentSize, layoutMeasurement } = scrollEvent.nativeEvent; atBottomRef.current = @@ -400,29 +469,39 @@ export default function ChatScreen() { /> - {streaming ? ( - activeRuntimeId && abortPrompt(activeRuntimeId)} - > - Stop - - ) : null} - void selectModel(model)} /> - - Send - + + {streaming ? ( + activeRuntimeId && abortPrompt(activeRuntimeId)} + > + Stop + + ) : null} + + void send()} + disabled={!canSend} + > + Send + + ); @@ -454,6 +533,9 @@ const styles = StyleSheet.create({ paddingHorizontal: 12, paddingVertical: 10, gap: 8, + }, + inputRow: { + gap: 8, flexDirection: "row", alignItems: "flex-end", }, diff --git a/apps/mobile/components/ModelSelector.tsx b/apps/mobile/components/ModelSelector.tsx new file mode 100644 index 0000000..ee26ce3 --- /dev/null +++ b/apps/mobile/components/ModelSelector.tsx @@ -0,0 +1,200 @@ +import { Ionicons } from "@expo/vector-icons"; +import { useMemo, useState } from "react"; +import { + ActivityIndicator, + FlatList, + Modal, + Pressable, + StyleSheet, + Text, + TextInput, + View, +} from "react-native"; +import { colors, monoFont } from "@/lib/theme"; +import type { ModelRef, ModelSummary } from "@/lib/types"; + +type ModelSelectorProps = { + models: ModelSummary[]; + selected?: ModelRef; + loading?: boolean; + error?: string | null; + disabled?: boolean; + onSelect: (model: ModelRef) => void; +}; + +function labelFor(model?: ModelRef, models: ModelSummary[] = []) { + if (!model) { + return "Select model"; + } + return models.find((entry) => entry.provider === model.provider && entry.id === model.id)?.id ?? model.id; +} + +export function ModelSelector({ + models, + selected, + loading, + error, + disabled, + onSelect, +}: ModelSelectorProps) { + const [open, setOpen] = useState(false); + const [query, setQuery] = useState(""); + const filteredModels = useMemo(() => { + const needle = query.trim().toLowerCase(); + if (!needle) { + return models; + } + return models.filter((model) => + `${model.id} ${model.name ?? ""} ${model.providerDisplayName ?? model.provider}` + .toLowerCase() + .includes(needle), + ); + }, [models, query]); + + const selectedSummary = selected + ? models.find((model) => model.provider === selected.provider && model.id === selected.id) + : undefined; + + return ( + <> + [styles.trigger, pressed && styles.pressed, disabled && styles.disabled]} + onPress={() => setOpen(true)} + disabled={disabled} + > + + + {labelFor(selected, models)} + + + {selectedSummary?.providerDisplayName ?? selected?.provider ?? "available models"} + + + {loading ? ( + + ) : ( + + )} + + + setOpen(false)}> + + + + Choose model + setOpen(false)} hitSlop={10}> + + + + + {error ? {error} : null} + {models.length === 0 ? ( + {loading ? "Loading models…" : "No available models"} + ) : ( + `${model.provider}:${model.id}`} + keyboardShouldPersistTaps="handled" + renderItem={({ item }) => { + const isSelected = + item.provider === selected?.provider && item.id === selected?.id; + return ( + [ + styles.option, + pressed && styles.pressed, + isSelected && styles.selected, + ]} + onPress={() => { + onSelect({ provider: item.provider, id: item.id }); + setOpen(false); + }} + > + + {item.id} + + {item.providerDisplayName ?? item.provider} + {item.name && item.name !== item.id ? ` · ${item.name}` : ""} + + + {isSelected ? ( + + ) : null} + + ); + }} + ListEmptyComponent={No matching models} + /> + )} + + + + + ); +} + +const styles = StyleSheet.create({ + trigger: { + minHeight: 40, + flexDirection: "row", + alignItems: "center", + justifyContent: "space-between", + backgroundColor: colors.surfaceDeep, + borderRadius: 8, + borderWidth: 1, + borderColor: colors.border, + paddingHorizontal: 11, + paddingVertical: 7, + }, + triggerText: { flex: 1, gap: 1 }, + modelId: { color: colors.text, fontFamily: monoFont, fontSize: 13 }, + provider: { color: colors.textMuted, fontFamily: monoFont, fontSize: 11 }, + pressed: { backgroundColor: colors.surface }, + disabled: { opacity: 0.45 }, + overlay: { flex: 1, justifyContent: "flex-end", backgroundColor: "rgba(0, 0, 0, 0.65)" }, + sheet: { + maxHeight: "78%", + backgroundColor: colors.surface, + borderTopLeftRadius: 16, + borderTopRightRadius: 16, + borderWidth: 1, + borderColor: colors.border, + padding: 16, + gap: 12, + }, + header: { flexDirection: "row", alignItems: "center", justifyContent: "space-between" }, + title: { color: colors.text, fontSize: 17, fontWeight: "700" }, + search: { + color: colors.text, + backgroundColor: colors.surfaceDeep, + borderWidth: 1, + borderColor: colors.border, + borderRadius: 8, + paddingHorizontal: 11, + paddingVertical: 9, + fontFamily: monoFont, + fontSize: 13, + }, + option: { + flexDirection: "row", + alignItems: "center", + justifyContent: "space-between", + paddingVertical: 12, + paddingHorizontal: 10, + borderRadius: 8, + }, + selected: { backgroundColor: "rgba(147, 197, 253, 0.12)" }, + optionText: { flex: 1, gap: 2 }, + optionId: { color: colors.text, fontFamily: monoFont, fontSize: 13 }, + optionMeta: { color: colors.textMuted, fontFamily: monoFont, fontSize: 11 }, + error: { color: colors.errorMuted, fontSize: 12 }, + empty: { color: colors.textMuted, fontFamily: monoFont, fontSize: 13, paddingVertical: 20 }, +}); diff --git a/apps/mobile/contexts/SocketContext.tsx b/apps/mobile/contexts/SocketContext.tsx index 3be78f1..8bf2b0f 100644 --- a/apps/mobile/contexts/SocketContext.tsx +++ b/apps/mobile/contexts/SocketContext.tsx @@ -13,16 +13,20 @@ import { checkHealth, loadSavedComputers, loadSavedProjects, + loadLastModel, removeSavedComputer, saveHostConfig, saveSavedProjects, + saveLastModel, upsertSavedComputer, } from "@/lib/api"; -import { createSocket } from "@/lib/socket"; +import { createSocket, listModels as requestModels, setRuntimeModel as requestSetRuntimeModel } from "@/lib/socket"; import type { ChatEvent, ChatHistoryPayload, LiveRuntimeInfo, + ModelListPayload, + ModelRef, PersistedSessionInfo, RuntimeCreatedPayload, SavedComputer, @@ -51,7 +55,10 @@ type SocketContextValue = { removeSavedComputerEntry: (id: string) => Promise; disconnect: () => void; refreshSessions: () => void; - createRuntime: (cwd: string, name?: string) => void; + createRuntime: (cwd: string, name?: string, model?: ModelRef) => void; + lastModel: ModelRef | null; + listModels: (payload: { runtimeId?: string; cwd?: string }) => Promise; + setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise; openRuntime: (sessionFile: string) => void; attachRuntime: (runtimeId: string) => void; sendPrompt: (runtimeId: string, message: string) => void; @@ -78,6 +85,7 @@ export function SocketProvider({ children }: { children: ReactNode }) { const [persistedSessions, setPersistedSessions] = useState([]); const [liveRuntimes, setLiveRuntimes] = useState([]); const [savedProjects, setSavedProjects] = useState([]); + const [lastModel, setLastModel] = useState(null); const socketRef = useRef(null); const chatListenersRef = useRef(new Set<(event: ChatEvent) => void>()); @@ -88,6 +96,7 @@ export function SocketProvider({ children }: { children: ReactNode }) { }>({ created: [], opened: [] }); useEffect(() => { + void loadLastModel().then(setLastModel); void loadSavedComputers().then((computers) => { setSavedComputers(computers); const recent = computers[0]; @@ -240,8 +249,44 @@ export function SocketProvider({ children }: { children: ReactNode }) { socketRef.current?.emit("sessions:list"); }, []); - const createRuntime = useCallback((cwd: string, name?: string) => { - socketRef.current?.emit("runtime:create", { cwd, name }); + 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) => { @@ -307,6 +352,7 @@ export function SocketProvider({ children }: { children: ReactNode }) { persistedSessions, liveRuntimes, savedProjects, + lastModel, addProject, removeProject, setHost, @@ -318,6 +364,8 @@ export function SocketProvider({ children }: { children: ReactNode }) { disconnect, refreshSessions, createRuntime, + listModels, + setRuntimeModel, openRuntime, attachRuntime, sendPrompt, @@ -338,6 +386,7 @@ export function SocketProvider({ children }: { children: ReactNode }) { persistedSessions, liveRuntimes, savedProjects, + lastModel, addProject, removeProject, connect, @@ -346,6 +395,8 @@ export function SocketProvider({ children }: { children: ReactNode }) { disconnect, refreshSessions, createRuntime, + listModels, + setRuntimeModel, openRuntime, attachRuntime, sendPrompt, diff --git a/apps/mobile/lib/api.ts b/apps/mobile/lib/api.ts index d18fa66..0503930 100644 --- a/apps/mobile/lib/api.ts +++ b/apps/mobile/lib/api.ts @@ -1,9 +1,10 @@ import AsyncStorage from "@react-native-async-storage/async-storage"; -import type { HostConfig, SavedComputer } from "./types"; +import type { HostConfig, ModelRef, SavedComputer } from "./types"; const HOST_KEY = "pi-mobile:last-host"; const COMPUTERS_KEY = "pi-mobile:computers"; const PROJECTS_PREFIX = "pi-mobile:projects:"; +const LAST_MODEL_KEY = "pi-mobile:last-model"; function generateId(): string { return `${Date.now()}-${Math.random().toString(36).slice(2, 9)}`; @@ -47,6 +48,25 @@ export async function saveHostConfig(config: HostConfig): Promise { await AsyncStorage.setItem(HOST_KEY, JSON.stringify(config)); } +export async function loadLastModel(): Promise { + const raw = await AsyncStorage.getItem(LAST_MODEL_KEY); + if (!raw) { + return null; + } + try { + const parsed = JSON.parse(raw) as Partial; + return typeof parsed.provider === "string" && typeof parsed.id === "string" + ? { provider: parsed.provider, id: parsed.id } + : null; + } catch { + return null; + } +} + +export async function saveLastModel(model: ModelRef): Promise { + await AsyncStorage.setItem(LAST_MODEL_KEY, JSON.stringify(model)); +} + export async function loadSavedComputers(): Promise { const raw = await AsyncStorage.getItem(COMPUTERS_KEY); if (raw) { diff --git a/apps/mobile/lib/socket.ts b/apps/mobile/lib/socket.ts index 3f5f511..e4b4c88 100644 --- a/apps/mobile/lib/socket.ts +++ b/apps/mobile/lib/socket.ts @@ -3,6 +3,8 @@ import type { ChatEvent, ChatHistoryPayload, LiveRuntimeInfo, + ModelListPayload, + ModelRef, PersistedSessionInfo, RuntimeCreatedPayload, } from "./types"; @@ -40,3 +42,51 @@ export function createSocket(baseUrl: string, callbacks: SocketCallbacks): Socke return socket; } + +export function listModels( + socket: Socket, + payload: { runtimeId?: string; cwd?: string }, + timeoutMs = 10000, +): Promise { + return new Promise((resolve, reject) => { + socket.timeout(timeoutMs).emit( + "models:list", + payload, + (error: Error | null, response: ModelListPayload | { error: string }) => { + if (error) { + reject(new Error("Timed out loading models")); + } else if ("error" in response) { + reject(new Error(response.error)); + } else { + resolve(response); + } + }, + ); + }); +} + +export function setRuntimeModel( + socket: Socket, + runtimeId: string, + model: ModelRef, + timeoutMs = 10000, +): Promise { + return new Promise((resolve, reject) => { + socket.timeout(timeoutMs).emit( + "runtime:model:set", + { runtimeId, model }, + ( + error: Error | null, + response: { runtime: LiveRuntimeInfo } | { error: string }, + ) => { + if (error) { + reject(new Error("Timed out setting model")); + } else if ("error" in response) { + reject(new Error(response.error)); + } else { + resolve(response.runtime); + } + }, + ); + }); +} diff --git a/apps/mobile/lib/types.ts b/apps/mobile/lib/types.ts index fd9d2da..28026be 100644 --- a/apps/mobile/lib/types.ts +++ b/apps/mobile/lib/types.ts @@ -1,3 +1,23 @@ +export type ModelRef = { + provider: string; + id: string; +}; + +export type ModelSummary = ModelRef & { + name?: string; + providerDisplayName?: string; + reasoning: boolean; + input: ("text" | "image")[]; + contextWindow: number; + maxTokens: number; +}; + +export type ModelListPayload = { + models: ModelSummary[]; + currentModel?: ModelRef; + loadError?: string; +}; + export type LiveRuntimeInfo = { runtimeId: string; sessionId: string; @@ -6,6 +26,7 @@ export type LiveRuntimeInfo = { sessionFile?: string; lastMessage?: string; isStreaming: boolean; + currentModel?: ModelRef; }; export type PersistedSessionInfo = { @@ -36,6 +57,7 @@ export type RuntimeCreatedPayload = { sessionName?: string; sessionFile?: string; lastMessage?: string; + currentModel?: ModelRef; }; export type ChatEvent = {