Add thinking level management to runtime and socket handlers

This commit is contained in:
2026-07-09 13:14:14 -05:00
parent a6029da45b
commit a0cfc5ba4a
8 changed files with 321 additions and 20 deletions
+34 -1
View File
@@ -16,6 +16,7 @@ import type {
ModelListPayload,
ModelRef,
ModelSummary,
ThinkingLevel,
PersistedSessionInfo,
} from "./types.js";
@@ -53,6 +54,8 @@ export class RuntimeRegistry {
lastMessage: lastMessagePreview(session.messages),
isStreaming: session.isStreaming,
currentModel: toModelRef(session.model),
thinkingLevel: session.thinkingLevel as ThinkingLevel,
availableThinkingLevels: session.getAvailableThinkingLevels() as ThinkingLevel[],
};
});
}
@@ -76,7 +79,12 @@ export class RuntimeRegistry {
return this.runtimes.get(runtimeId);
}
async createRuntime(cwd: string, name?: string, modelRef?: ModelRef): Promise<LiveRuntimeInfo> {
async createRuntime(
cwd: string,
name?: string,
modelRef?: ModelRef,
thinkingLevel?: ThinkingLevel,
): Promise<LiveRuntimeInfo> {
const sessionManager = SessionManager.create(cwd);
const runtime = await createAgentSessionRuntime(createRuntimeFactory, {
cwd,
@@ -87,6 +95,9 @@ export class RuntimeRegistry {
if (modelRef) {
await this.applyModel(runtime, modelRef);
}
if (thinkingLevel) {
runtime.session.setThinkingLevel(thinkingLevel);
}
if (name) {
runtime.session.sessionManager.appendSessionInfo(name);
@@ -147,6 +158,12 @@ export class RuntimeRegistry {
.getAvailable()
.map((model) => toModelSummary(model, registry)),
currentModel: record ? toModelRef(record.runtime.session.model) : undefined,
currentThinkingLevel: record
? (record.runtime.session.thinkingLevel as ThinkingLevel)
: undefined,
availableThinkingLevels: record
? (record.runtime.session.getAvailableThinkingLevels() as ThinkingLevel[])
: undefined,
loadError: registry.getError(),
};
}
@@ -157,6 +174,18 @@ export class RuntimeRegistry {
return this.getLiveInfo(runtimeId, record);
}
setThinkingLevel(runtimeId: string, level: ThinkingLevel): LiveRuntimeInfo {
const record = this.requireRuntime(runtimeId);
record.runtime.session.setThinkingLevel(level);
return this.getLiveInfo(runtimeId, record);
}
cycleThinkingLevel(runtimeId: string): LiveRuntimeInfo {
const record = this.requireRuntime(runtimeId);
record.runtime.session.cycleThinkingLevel();
return this.getLiveInfo(runtimeId, record);
}
dispose(runtimeId: string): void {
const record = this.runtimes.get(runtimeId);
if (!record) {
@@ -189,6 +218,8 @@ export class RuntimeRegistry {
lastMessage: lastMessagePreview(session.messages),
isStreaming: session.isStreaming,
currentModel: toModelRef(session.model),
thinkingLevel: session.thinkingLevel as ThinkingLevel,
availableThinkingLevels: session.getAvailableThinkingLevels() as ThinkingLevel[],
};
}
@@ -203,6 +234,8 @@ export class RuntimeRegistry {
lastMessage: lastMessagePreview(session.messages),
isStreaming: session.isStreaming,
currentModel: toModelRef(session.model),
thinkingLevel: session.thinkingLevel as ThinkingLevel,
availableThinkingLevels: session.getAvailableThinkingLevels() as ThinkingLevel[],
};
}
+47 -3
View File
@@ -7,6 +7,7 @@ import type {
ModelListPayload,
ModelRef,
RuntimeCreatedPayload,
ThinkingLevel,
} from "./types.js";
type AttachState = {
@@ -64,9 +65,51 @@ export function registerSocketHandlers(io: Server, registry: RuntimeRegistry) {
},
);
socket.on("runtime:create", async (payload: { cwd: string; name?: string; model?: ModelRef }) => {
socket.on(
"runtime:thinking:set",
async (
payload: { runtimeId: string; level: ThinkingLevel },
ack?: (response: { runtime: RuntimeCreatedPayload } | { error: string }) => void,
) => {
try {
const runtime = registry.setThinkingLevel(payload.runtimeId, payload.level);
ack?.({ runtime });
} catch (error) {
ack?.({ error: errorMessage(error) });
}
},
);
socket.on(
"runtime:thinking:cycle",
async (
payload: { runtimeId: string },
ack?: (response: { runtime: RuntimeCreatedPayload } | { error: string }) => void,
) => {
try {
const runtime = registry.cycleThinkingLevel(payload.runtimeId);
ack?.({ runtime });
} catch (error) {
ack?.({ error: errorMessage(error) });
}
},
);
socket.on(
"runtime:create",
async (payload: {
cwd: string;
name?: string;
model?: ModelRef;
thinkingLevel?: ThinkingLevel;
}) => {
try {
const info = await registry.createRuntime(payload.cwd, payload.name, payload.model);
const info = await registry.createRuntime(
payload.cwd,
payload.name,
payload.model,
payload.thinkingLevel,
);
const created: RuntimeCreatedPayload = info;
socket.emit("runtime:created", created);
io.emit("sessions:list:result", {
@@ -76,7 +119,8 @@ export function registerSocketHandlers(io: Server, registry: RuntimeRegistry) {
} catch (error) {
emitError(socket, error);
}
});
},
);
socket.on("runtime:open", async (payload: { sessionFile: string }) => {
try {
+7
View File
@@ -5,6 +5,8 @@ export type ModelRef = {
id: string;
};
export type ThinkingLevel = "off" | "minimal" | "low" | "medium" | "high" | "xhigh";
export type ModelSummary = ModelRef & {
name?: string;
providerDisplayName?: string;
@@ -17,6 +19,8 @@ export type ModelSummary = ModelRef & {
export type ModelListPayload = {
models: ModelSummary[];
currentModel?: ModelRef;
currentThinkingLevel?: ThinkingLevel;
availableThinkingLevels?: ThinkingLevel[];
loadError?: string;
};
@@ -29,6 +33,8 @@ export type LiveRuntimeInfo = {
lastMessage?: string;
isStreaming: boolean;
currentModel?: ModelRef;
thinkingLevel?: ThinkingLevel;
availableThinkingLevels?: ThinkingLevel[];
};
export type PersistedSessionInfo = {
@@ -72,4 +78,5 @@ export type RuntimeCreatedPayload = {
sessionFile?: string;
lastMessage?: string;
currentModel?: ModelRef;
thinkingLevel?: ThinkingLevel;
};
+83 -11
View File
@@ -15,13 +15,22 @@ import {
import { AssistantMessage } from "@/components/AssistantMessage";
import { MessageRow, type Sender } from "@/components/MessageRow";
import { ModelSelector } from "@/components/ModelSelector";
import { ReasoningEffortButton } from "@/components/ReasoningEffortButton";
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, ModelRef, ModelSummary } from "@/lib/types";
import type {
ChatHistoryMessage,
ChatItem,
ModelRef,
ModelSummary,
ThinkingLevel,
} from "@/lib/types";
const ALL_THINKING_LEVELS: ThinkingLevel[] = ["off", "minimal", "low", "medium", "high", "xhigh"];
function extractText(content: unknown): string {
if (typeof content === "string") {
@@ -127,6 +136,7 @@ export default function ChatScreen() {
listModels,
sendPrompt,
setRuntimeModel,
setRuntimeThinkingLevel,
abortPrompt,
subscribeChat,
subscribeChatHistory,
@@ -143,6 +153,9 @@ export default function ChatScreen() {
const [selectedModel, setSelectedModel] = useState<ModelRef | undefined>();
const [modelLoading, setModelLoading] = useState(false);
const [modelError, setModelError] = useState<string | null>(null);
const [thinkingLevel, setThinkingLevel] = useState<ThinkingLevel>("off");
const [availableThinkingLevels, setAvailableThinkingLevels] =
useState<ThinkingLevel[]>(ALL_THINKING_LEVELS);
const assistantIdRef = useRef<string | null>(null);
const toolIdsRef = useRef(new Map<string, string>());
const pendingMessageRef = useRef<string | null>(null);
@@ -171,6 +184,17 @@ export default function ChatScreen() {
persistedModel ??
(payload.models[0] ? { provider: payload.models[0].provider, id: payload.models[0].id } : undefined),
);
const selected = payload.currentModel ?? persistedModel ?? payload.models[0];
const selectedSummary = selected
? payload.models.find(
(model) => model.provider === selected.provider && model.id === selected.id,
)
: undefined;
setAvailableThinkingLevels(
payload.availableThinkingLevels ??
(selectedSummary?.reasoning ? ALL_THINKING_LEVELS : ["off"]),
);
setThinkingLevel(payload.currentThinkingLevel ?? "off");
if (payload.loadError) {
setModelError(payload.loadError);
}
@@ -396,7 +420,7 @@ export default function ChatScreen() {
pendingMessageRef.current = message;
setCreating(true);
createRuntime(projectCwd, undefined, selectedModel);
createRuntime(projectCwd, undefined, selectedModel, thinkingLevel);
try {
const runtime = await waitForRuntime("created");
setActiveRuntimeId(runtime.runtimeId);
@@ -421,17 +445,47 @@ export default function ChatScreen() {
setSelectedModel(model);
setModelError(null);
if (!activeRuntimeId) {
const summary = models.find(
(entry) => entry.provider === model.provider && entry.id === model.id,
);
const nextLevels = summary?.reasoning ? ALL_THINKING_LEVELS : ["off" as ThinkingLevel];
setAvailableThinkingLevels(nextLevels);
if (!nextLevels.includes(thinkingLevel)) {
setThinkingLevel("off");
}
void saveLastModel(model);
return;
}
try {
await setRuntimeModel(activeRuntimeId, model);
const runtime = await setRuntimeModel(activeRuntimeId, model);
setThinkingLevel(runtime.thinkingLevel ?? "off");
setAvailableThinkingLevels(runtime.availableThinkingLevels ?? ["off"]);
} catch (error: unknown) {
setSelectedModel(previous);
setModelError(error instanceof Error ? error.message : "Failed to set model");
}
};
const cycleReasoningEffort = async () => {
if (streaming || availableThinkingLevels.length < 2) {
return;
}
const currentIndex = availableThinkingLevels.indexOf(thinkingLevel);
const nextLevel =
availableThinkingLevels[(currentIndex + 1) % availableThinkingLevels.length] ?? "off";
setThinkingLevel(nextLevel);
if (!activeRuntimeId) {
return;
}
try {
const runtime = await setRuntimeThinkingLevel(activeRuntimeId, nextLevel);
setThinkingLevel(runtime.thinkingLevel ?? nextLevel);
setAvailableThinkingLevels(runtime.availableThinkingLevels ?? availableThinkingLevels);
} catch {
setThinkingLevel(thinkingLevel);
}
};
const handleScroll = (scrollEvent: NativeSyntheticEvent<NativeScrollEvent>) => {
const { contentOffset, contentSize, layoutMeasurement } = scrollEvent.nativeEvent;
atBottomRef.current =
@@ -469,14 +523,23 @@ 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.selectorRow}>
<View style={styles.modelSelector}>
<ModelSelector
models={models}
selected={selectedModel}
loading={modelLoading}
error={modelError}
disabled={streaming || modelLoading}
onSelect={(model) => void selectModel(model)}
/>
</View>
<ReasoningEffortButton
level={thinkingLevel}
disabled={streaming || modelLoading || availableThinkingLevels.length < 2}
onPress={() => void cycleReasoningEffort()}
/>
</View>
<View style={styles.inputRow}>
{streaming ? (
<Pressable
@@ -534,6 +597,15 @@ const styles = StyleSheet.create({
paddingVertical: 10,
gap: 8,
},
selectorRow: {
flexDirection: "row",
alignItems: "stretch",
gap: 8,
},
modelSelector: {
flex: 1,
minWidth: 0,
},
inputRow: {
gap: 8,
flexDirection: "row",
@@ -0,0 +1,72 @@
import { Ionicons } from "@expo/vector-icons";
import { Pressable, StyleSheet, Text } from "react-native";
import { colors, monoFont } from "@/lib/theme";
import type { ThinkingLevel } from "@/lib/types";
type ReasoningEffortButtonProps = {
level: ThinkingLevel;
disabled?: boolean;
onPress: () => void;
};
const labels: Record<ThinkingLevel, string> = {
off: "Off",
minimal: "Minimal",
low: "Low",
medium: "Medium",
high: "High",
xhigh: "X-high",
};
export function ReasoningEffortButton({
level,
disabled,
onPress,
}: ReasoningEffortButtonProps) {
return (
<Pressable
accessibilityLabel={`Reasoning effort ${labels[level]}. Tap to change.`}
style={({ pressed }) => [
styles.button,
pressed && styles.pressed,
disabled && styles.disabled,
level !== "off" && styles.active,
]}
onPress={onPress}
disabled={disabled}
>
<Ionicons
name="sparkles-outline"
size={14}
color={level === "off" ? colors.textMuted : colors.accent}
/>
<Text style={styles.label}>{labels[level]}</Text>
</Pressable>
);
}
const styles = StyleSheet.create({
button: {
minHeight: 40,
flexDirection: "row",
alignItems: "center",
gap: 6,
borderRadius: 8,
borderWidth: 1,
borderColor: colors.border,
backgroundColor: colors.surfaceDeep,
paddingHorizontal: 10,
},
active: {
borderColor: "rgba(147, 197, 253, 0.45)",
backgroundColor: "rgba(147, 197, 253, 0.1)",
},
label: {
color: colors.text,
fontFamily: monoFont,
fontSize: 12,
fontWeight: "700",
},
pressed: { backgroundColor: colors.surface },
disabled: { opacity: 0.45 },
});
+47 -5
View File
@@ -20,7 +20,12 @@ import {
saveLastModel,
upsertSavedComputer,
} from "@/lib/api";
import { createSocket, listModels as requestModels, setRuntimeModel as requestSetRuntimeModel } from "@/lib/socket";
import {
createSocket,
listModels as requestModels,
setRuntimeModel as requestSetRuntimeModel,
setRuntimeThinkingLevel as requestSetRuntimeThinkingLevel,
} from "@/lib/socket";
import type {
ChatEvent,
ChatHistoryPayload,
@@ -30,6 +35,7 @@ import type {
PersistedSessionInfo,
RuntimeCreatedPayload,
SavedComputer,
ThinkingLevel,
} from "@/lib/types";
import type { Socket } from "socket.io-client";
@@ -55,10 +61,16 @@ type SocketContextValue = {
removeSavedComputerEntry: (id: string) => Promise<void>;
disconnect: () => void;
refreshSessions: () => void;
createRuntime: (cwd: string, name?: string, model?: ModelRef) => void;
createRuntime: (
cwd: string,
name?: string,
model?: ModelRef,
thinkingLevel?: ThinkingLevel,
) => void;
lastModel: ModelRef | null;
listModels: (payload: { runtimeId?: string; cwd?: string }) => Promise<ModelListPayload>;
setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise<void>;
setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise<LiveRuntimeInfo>;
setRuntimeThinkingLevel: (runtimeId: string, level: ThinkingLevel) => Promise<LiveRuntimeInfo>;
openRuntime: (sessionFile: string) => void;
attachRuntime: (runtimeId: string) => void;
sendPrompt: (runtimeId: string, message: string) => void;
@@ -249,8 +261,13 @@ export function SocketProvider({ children }: { children: ReactNode }) {
socketRef.current?.emit("sessions:list");
}, []);
const createRuntime = useCallback((cwd: string, name?: string, model?: ModelRef) => {
socketRef.current?.emit("runtime:create", { cwd, name, model });
const createRuntime = useCallback((
cwd: string,
name?: string,
model?: ModelRef,
thinkingLevel?: ThinkingLevel,
) => {
socketRef.current?.emit("runtime:create", { cwd, name, model, thinkingLevel });
}, []);
const listModels = useCallback(
@@ -282,6 +299,7 @@ export function SocketProvider({ children }: { children: ReactNode }) {
setLiveRuntimes((current) =>
current.map((entry) => (entry.runtimeId === runtime.runtimeId ? runtime : entry)),
);
return runtime;
} catch (requestError) {
const message = requestError instanceof Error ? requestError.message : "Failed to set model";
setError(message);
@@ -289,6 +307,28 @@ export function SocketProvider({ children }: { children: ReactNode }) {
}
}, []);
const setRuntimeThinkingLevel = useCallback(
async (runtimeId: string, level: ThinkingLevel) => {
const socket = socketRef.current;
if (!socket) {
throw new Error("Not connected");
}
try {
const runtime = await requestSetRuntimeThinkingLevel(socket, runtimeId, level);
setLiveRuntimes((current) =>
current.map((entry) => (entry.runtimeId === runtime.runtimeId ? runtime : entry)),
);
return runtime;
} catch (requestError) {
const message =
requestError instanceof Error ? requestError.message : "Failed to set reasoning effort";
setError(message);
throw requestError;
}
},
[],
);
const openRuntime = useCallback((sessionFile: string) => {
socketRef.current?.emit("runtime:open", { sessionFile });
}, []);
@@ -366,6 +406,7 @@ export function SocketProvider({ children }: { children: ReactNode }) {
createRuntime,
listModels,
setRuntimeModel,
setRuntimeThinkingLevel,
openRuntime,
attachRuntime,
sendPrompt,
@@ -397,6 +438,7 @@ export function SocketProvider({ children }: { children: ReactNode }) {
createRuntime,
listModels,
setRuntimeModel,
setRuntimeThinkingLevel,
openRuntime,
attachRuntime,
sendPrompt,
+24
View File
@@ -7,6 +7,7 @@ import type {
ModelRef,
PersistedSessionInfo,
RuntimeCreatedPayload,
ThinkingLevel,
} from "./types";
export type SocketCallbacks = {
@@ -90,3 +91,26 @@ export function setRuntimeModel(
);
});
}
export function setRuntimeThinkingLevel(
socket: Socket,
runtimeId: string,
level: ThinkingLevel,
timeoutMs = 10000,
): Promise<LiveRuntimeInfo> {
return new Promise((resolve, reject) => {
socket.timeout(timeoutMs).emit(
"runtime:thinking:set",
{ runtimeId, level },
(error: Error | null, response: { runtime: LiveRuntimeInfo } | { error: string }) => {
if (error) {
reject(new Error("Timed out setting reasoning effort"));
} else if ("error" in response) {
reject(new Error(response.error));
} else {
resolve(response.runtime);
}
},
);
});
}
+7
View File
@@ -3,6 +3,8 @@ export type ModelRef = {
id: string;
};
export type ThinkingLevel = "off" | "minimal" | "low" | "medium" | "high" | "xhigh";
export type ModelSummary = ModelRef & {
name?: string;
providerDisplayName?: string;
@@ -15,6 +17,8 @@ export type ModelSummary = ModelRef & {
export type ModelListPayload = {
models: ModelSummary[];
currentModel?: ModelRef;
currentThinkingLevel?: ThinkingLevel;
availableThinkingLevels?: ThinkingLevel[];
loadError?: string;
};
@@ -27,6 +31,8 @@ export type LiveRuntimeInfo = {
lastMessage?: string;
isStreaming: boolean;
currentModel?: ModelRef;
thinkingLevel?: ThinkingLevel;
availableThinkingLevels?: ThinkingLevel[];
};
export type PersistedSessionInfo = {
@@ -58,6 +64,7 @@ export type RuntimeCreatedPayload = {
sessionFile?: string;
lastMessage?: string;
currentModel?: ModelRef;
thinkingLevel?: ThinkingLevel;
};
export type ChatEvent = {