From a0cfc5ba4ab5969af09785286353cbc01e22d01c Mon Sep 17 00:00:00 2001 From: Owen Qwen Date: Thu, 9 Jul 2026 13:14:14 -0500 Subject: [PATCH] Add thinking level management to runtime and socket handlers --- apps/api/src/runtime-registry.ts | 35 ++++++- apps/api/src/socket-handlers.ts | 50 +++++++++- apps/api/src/types.ts | 7 ++ apps/mobile/app/chat/[runtimeId].tsx | 94 ++++++++++++++++--- .../components/ReasoningEffortButton.tsx | 72 ++++++++++++++ apps/mobile/contexts/SocketContext.tsx | 52 +++++++++- apps/mobile/lib/socket.ts | 24 +++++ apps/mobile/lib/types.ts | 7 ++ 8 files changed, 321 insertions(+), 20 deletions(-) create mode 100644 apps/mobile/components/ReasoningEffortButton.tsx diff --git a/apps/api/src/runtime-registry.ts b/apps/api/src/runtime-registry.ts index 1db80bf..d3e31a3 100644 --- a/apps/api/src/runtime-registry.ts +++ b/apps/api/src/runtime-registry.ts @@ -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 { + async createRuntime( + cwd: string, + name?: string, + modelRef?: ModelRef, + thinkingLevel?: ThinkingLevel, + ): Promise { 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[], }; } diff --git a/apps/api/src/socket-handlers.ts b/apps/api/src/socket-handlers.ts index c075246..308567b 100644 --- a/apps/api/src/socket-handlers.ts +++ b/apps/api/src/socket-handlers.ts @@ -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 { diff --git a/apps/api/src/types.ts b/apps/api/src/types.ts index e7ef156..d1987fd 100644 --- a/apps/api/src/types.ts +++ b/apps/api/src/types.ts @@ -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; }; diff --git a/apps/mobile/app/chat/[runtimeId].tsx b/apps/mobile/app/chat/[runtimeId].tsx index c5e38b5..849be46 100644 --- a/apps/mobile/app/chat/[runtimeId].tsx +++ b/apps/mobile/app/chat/[runtimeId].tsx @@ -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(); const [modelLoading, setModelLoading] = useState(false); const [modelError, setModelError] = useState(null); + const [thinkingLevel, setThinkingLevel] = useState("off"); + const [availableThinkingLevels, setAvailableThinkingLevels] = + useState(ALL_THINKING_LEVELS); const assistantIdRef = useRef(null); const toolIdsRef = useRef(new Map()); const pendingMessageRef = useRef(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) => { const { contentOffset, contentSize, layoutMeasurement } = scrollEvent.nativeEvent; atBottomRef.current = @@ -469,14 +523,23 @@ export default function ChatScreen() { /> - void selectModel(model)} - /> + + + void selectModel(model)} + /> + + void cycleReasoningEffort()} + /> + {streaming ? ( void; +}; + +const labels: Record = { + off: "Off", + minimal: "Minimal", + low: "Low", + medium: "Medium", + high: "High", + xhigh: "X-high", +}; + +export function ReasoningEffortButton({ + level, + disabled, + onPress, +}: ReasoningEffortButtonProps) { + return ( + [ + styles.button, + pressed && styles.pressed, + disabled && styles.disabled, + level !== "off" && styles.active, + ]} + onPress={onPress} + disabled={disabled} + > + + {labels[level]} + + ); +} + +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 }, +}); diff --git a/apps/mobile/contexts/SocketContext.tsx b/apps/mobile/contexts/SocketContext.tsx index 8bf2b0f..cbe94ae 100644 --- a/apps/mobile/contexts/SocketContext.tsx +++ b/apps/mobile/contexts/SocketContext.tsx @@ -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; 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; - setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise; + setRuntimeModel: (runtimeId: string, model: ModelRef) => Promise; + setRuntimeThinkingLevel: (runtimeId: string, level: ThinkingLevel) => Promise; 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, diff --git a/apps/mobile/lib/socket.ts b/apps/mobile/lib/socket.ts index e4b4c88..60338ba 100644 --- a/apps/mobile/lib/socket.ts +++ b/apps/mobile/lib/socket.ts @@ -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 { + 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); + } + }, + ); + }); +} diff --git a/apps/mobile/lib/types.ts b/apps/mobile/lib/types.ts index 28026be..a5cef46 100644 --- a/apps/mobile/lib/types.ts +++ b/apps/mobile/lib/types.ts @@ -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 = {