Adding a model selector

This commit is contained in:
2026-07-09 13:08:49 -05:00
parent 6fa32c1020
commit a6029da45b
9 changed files with 611 additions and 35 deletions
+106 -24
View File
@@ -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<ModelSummary[]>([]);
const [selectedModel, setSelectedModel] = useState<ModelRef | undefined>();
const [modelLoading, setModelLoading] = useState(false);
const [modelError, setModelError] = useState<string | null>(null);
const assistantIdRef = useRef<string | null>(null);
const toolIdsRef = useRef(new Map<string, string>());
const pendingMessageRef = useRef<string | null>(null);
const listRef = useRef<FlatList<ChatItem>>(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<NativeScrollEvent>) => {
const { contentOffset, contentSize, layoutMeasurement } = scrollEvent.nativeEvent;
atBottomRef.current =
@@ -400,29 +469,39 @@ export default function ChatScreen() {
/>
<View style={styles.composer}>
{streaming ? (
<Pressable
style={styles.abortButton}
onPress={() => activeRuntimeId && abortPrompt(activeRuntimeId)}
>
<Text style={styles.abortText}>Stop</Text>
</Pressable>
) : null}
<TextInput
value={draft}
onChangeText={setDraft}
placeholder="Message Pi…"
placeholderTextColor={colors.textMuted}
style={styles.input}
multiline
<ModelSelector
models={models}
selected={selectedModel}
loading={modelLoading}
error={modelError}
disabled={streaming || modelLoading}
onSelect={(model) => void selectModel(model)}
/>
<Pressable
style={[styles.sendButton, !canSend && styles.sendButtonDisabled]}
onPress={send}
disabled={!canSend}
>
<Text style={styles.sendText}>Send</Text>
</Pressable>
<View style={styles.inputRow}>
{streaming ? (
<Pressable
style={styles.abortButton}
onPress={() => activeRuntimeId && abortPrompt(activeRuntimeId)}
>
<Text style={styles.abortText}>Stop</Text>
</Pressable>
) : null}
<TextInput
value={draft}
onChangeText={setDraft}
placeholder="Message Pi…"
placeholderTextColor={colors.textMuted}
style={styles.input}
multiline
/>
<Pressable
style={[styles.sendButton, !canSend && styles.sendButtonDisabled]}
onPress={() => void send()}
disabled={!canSend}
>
<Text style={styles.sendText}>Send</Text>
</Pressable>
</View>
</View>
</KeyboardAvoidingView>
);
@@ -454,6 +533,9 @@ const styles = StyleSheet.create({
paddingHorizontal: 12,
paddingVertical: 10,
gap: 8,
},
inputRow: {
gap: 8,
flexDirection: "row",
alignItems: "flex-end",
},
+200
View File
@@ -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 (
<>
<Pressable
style={({ pressed }) => [styles.trigger, pressed && styles.pressed, disabled && styles.disabled]}
onPress={() => setOpen(true)}
disabled={disabled}
>
<View style={styles.triggerText}>
<Text style={styles.modelId} numberOfLines={1}>
{labelFor(selected, models)}
</Text>
<Text style={styles.provider} numberOfLines={1}>
{selectedSummary?.providerDisplayName ?? selected?.provider ?? "available models"}
</Text>
</View>
{loading ? (
<ActivityIndicator size="small" color={colors.textSecondary} />
) : (
<Ionicons name="chevron-down" size={16} color={colors.textMuted} />
)}
</Pressable>
<Modal visible={open} animationType="slide" transparent onRequestClose={() => setOpen(false)}>
<View style={styles.overlay}>
<View style={styles.sheet}>
<View style={styles.header}>
<Text style={styles.title}>Choose model</Text>
<Pressable onPress={() => setOpen(false)} hitSlop={10}>
<Ionicons name="close" size={22} color={colors.textSecondary} />
</Pressable>
</View>
<TextInput
value={query}
onChangeText={setQuery}
placeholder="Search models…"
placeholderTextColor={colors.textMuted}
style={styles.search}
autoCorrect={false}
autoCapitalize="none"
/>
{error ? <Text style={styles.error}>{error}</Text> : null}
{models.length === 0 ? (
<Text style={styles.empty}>{loading ? "Loading models…" : "No available models"}</Text>
) : (
<FlatList
data={filteredModels}
keyExtractor={(model) => `${model.provider}:${model.id}`}
keyboardShouldPersistTaps="handled"
renderItem={({ item }) => {
const isSelected =
item.provider === selected?.provider && item.id === selected?.id;
return (
<Pressable
style={({ pressed }) => [
styles.option,
pressed && styles.pressed,
isSelected && styles.selected,
]}
onPress={() => {
onSelect({ provider: item.provider, id: item.id });
setOpen(false);
}}
>
<View style={styles.optionText}>
<Text style={styles.optionId}>{item.id}</Text>
<Text style={styles.optionMeta}>
{item.providerDisplayName ?? item.provider}
{item.name && item.name !== item.id ? ` · ${item.name}` : ""}
</Text>
</View>
{isSelected ? (
<Ionicons name="checkmark" size={18} color={colors.accent} />
) : null}
</Pressable>
);
}}
ListEmptyComponent={<Text style={styles.empty}>No matching models</Text>}
/>
)}
</View>
</View>
</Modal>
</>
);
}
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 },
});
+55 -4
View File
@@ -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<void>;
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<ModelListPayload>;
setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise<void>;
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<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>());
@@ -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,
+21 -1
View File
@@ -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<void> {
await AsyncStorage.setItem(HOST_KEY, JSON.stringify(config));
}
export async function loadLastModel(): Promise<ModelRef | null> {
const raw = await AsyncStorage.getItem(LAST_MODEL_KEY);
if (!raw) {
return null;
}
try {
const parsed = JSON.parse(raw) as Partial<ModelRef>;
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<void> {
await AsyncStorage.setItem(LAST_MODEL_KEY, JSON.stringify(model));
}
export async function loadSavedComputers(): Promise<SavedComputer[]> {
const raw = await AsyncStorage.getItem(COMPUTERS_KEY);
if (raw) {
+50
View File
@@ -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<ModelListPayload> {
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<LiveRuntimeInfo> {
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);
}
},
);
});
}
+22
View File
@@ -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 = {