Adding a model selector
This commit is contained in:
@@ -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<LiveRuntimeInfo> {
|
||||
async createRuntime(cwd: string, name?: string, modelRef?: ModelRef): Promise<LiveRuntimeInfo> {
|
||||
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<ModelListPayload> {
|
||||
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<LiveRuntimeInfo> {
|
||||
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<void> {
|
||||
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,
|
||||
};
|
||||
}
|
||||
|
||||
@@ -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 info = await registry.createRuntime(payload.cwd, payload.name);
|
||||
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, 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";
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
@@ -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,6 +469,15 @@ export default function ChatScreen() {
|
||||
/>
|
||||
|
||||
<View style={styles.composer}>
|
||||
<ModelSelector
|
||||
models={models}
|
||||
selected={selectedModel}
|
||||
loading={modelLoading}
|
||||
error={modelError}
|
||||
disabled={streaming || modelLoading}
|
||||
onSelect={(model) => void selectModel(model)}
|
||||
/>
|
||||
<View style={styles.inputRow}>
|
||||
{streaming ? (
|
||||
<Pressable
|
||||
style={styles.abortButton}
|
||||
@@ -418,12 +496,13 @@ export default function ChatScreen() {
|
||||
/>
|
||||
<Pressable
|
||||
style={[styles.sendButton, !canSend && styles.sendButtonDisabled]}
|
||||
onPress={send}
|
||||
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",
|
||||
},
|
||||
|
||||
@@ -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 },
|
||||
});
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
},
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user