580 lines
17 KiB
TypeScript
580 lines
17 KiB
TypeScript
import type { ServerWebSocket } from 'bun';
|
|
import { randomUUID } from 'crypto';
|
|
import type { ClientMessage, ServerMessage, Message, PiEvent } from './types';
|
|
import { sessionManager } from './session-manager';
|
|
import * as storage from './storage';
|
|
import * as piBridge from './pi-bridge';
|
|
import { sendClaudeCodeStreaming } from '@@/channels/send-claude-code';
|
|
import { join } from 'path';
|
|
import { getHomeDirForRole } from '../../../servers/data-path';
|
|
import { getUserSettings } from 'officerdb';
|
|
import { logger } from './logger';
|
|
|
|
// Default model when no user preference is set
|
|
const DEFAULT_MODEL = 'claude-code';
|
|
|
|
async function getUserDefaultModel(userId: number): Promise<string | null> {
|
|
try {
|
|
const settings = await getUserSettings(userId);
|
|
const chat = settings?.chat as Record<string, unknown> | undefined;
|
|
return (chat?.defaultModel as string) || null;
|
|
} catch (err) {
|
|
logger.error('Failed to read user settings for default model', { userId, error: String(err) });
|
|
}
|
|
return null;
|
|
}
|
|
|
|
type WSData = {
|
|
userId: number;
|
|
email: string;
|
|
username: string;
|
|
role: string;
|
|
sandboxed: boolean;
|
|
provider: string;
|
|
};
|
|
|
|
const IDLE_TIMEOUT_MS = 60 * 60 * 1000; // 1 hour
|
|
|
|
const resolveCwd = (email: string, role: string, cwd?: string) => {
|
|
const root = getHomeDirForRole(email, role);
|
|
if (!cwd || cwd === '~') return root;
|
|
if (cwd.startsWith('~/')) return join(root, cwd.slice(2));
|
|
if (cwd.startsWith('/')) return join(root, cwd.slice(1));
|
|
return join(root, cwd);
|
|
};
|
|
|
|
export const resolveBaseCwd = (email: string, role: string, cwd?: string) => {
|
|
return resolveCwd(email, role, cwd);
|
|
};
|
|
|
|
const wsToSessionMap = new WeakMap<any, string>();
|
|
|
|
function sendToClient(ws: ServerWebSocket<WSData> | null, msg: ServerMessage): void {
|
|
if (ws?.readyState === 1) {
|
|
ws.send(JSON.stringify(msg));
|
|
}
|
|
}
|
|
|
|
export async function open(ws: ServerWebSocket<WSData>): Promise<void> {
|
|
// logger.info('WebSocket connection opened', { email: ws.data.email });
|
|
}
|
|
|
|
export function message(ws: ServerWebSocket<WSData>, raw: string | Buffer): void {
|
|
const data = typeof raw === 'string' ? raw : raw.toString();
|
|
|
|
(async () => {
|
|
try {
|
|
const clientMsg = JSON.parse(data) as ClientMessage;
|
|
|
|
if (clientMsg.type === 'chat') {
|
|
await handleChat(ws, clientMsg);
|
|
} else if (clientMsg.type === 'resume') {
|
|
await handleResume(ws, clientMsg);
|
|
} else if (clientMsg.type === 'stop') {
|
|
await handleStop(ws);
|
|
}
|
|
} catch (err) {
|
|
logger.error('Error handling WebSocket message', { email: ws.data.email, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Failed to process message' });
|
|
}
|
|
})();
|
|
}
|
|
|
|
export function close(ws: ServerWebSocket<WSData>): void {
|
|
// logger.info('WebSocket connection closed', { email: ws.data.email });
|
|
|
|
const sessionId = wsToSessionMap.get(ws);
|
|
if (sessionId) {
|
|
sessionManager.detachWs(sessionId);
|
|
sessionManager.setIdleTimeout(sessionId, IDLE_TIMEOUT_MS);
|
|
}
|
|
}
|
|
|
|
function createEventHandler(sessionId: string, model: string, cwd: string, storageDir: string) {
|
|
return async (event: PiEvent): Promise<void> => {
|
|
const session = sessionManager.getSession(sessionId);
|
|
if (!session) return;
|
|
|
|
const ws = session.ws as ServerWebSocket<WSData> | null;
|
|
|
|
switch (event.type) {
|
|
case 'delta': {
|
|
sendToClient(ws, { type: 'assistant:delta', text: event.text });
|
|
session.streamBuffer += event.text;
|
|
break;
|
|
}
|
|
|
|
case 'text': {
|
|
// Flush streaming buffer as complete text
|
|
const text = event.text || session.streamBuffer;
|
|
if (text) {
|
|
sendToClient(ws, { type: 'assistant:text', text });
|
|
|
|
const assistantMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'assistant',
|
|
text,
|
|
model,
|
|
};
|
|
session.messages.push(assistantMsg);
|
|
session.meta.messageCount += 1;
|
|
session.streamBuffer = '';
|
|
}
|
|
break;
|
|
}
|
|
|
|
case 'tool:start': {
|
|
// Flush any pending streaming text first
|
|
if (session.streamBuffer) {
|
|
sendToClient(ws, { type: 'assistant:text', text: session.streamBuffer });
|
|
|
|
const assistantMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'assistant',
|
|
text: session.streamBuffer,
|
|
model,
|
|
};
|
|
session.messages.push(assistantMsg);
|
|
session.meta.messageCount += 1;
|
|
session.streamBuffer = '';
|
|
}
|
|
|
|
sendToClient(ws, {
|
|
type: 'tool:start',
|
|
toolCallId: event.toolCallId,
|
|
toolName: event.toolName,
|
|
toolInput: event.toolInput,
|
|
});
|
|
|
|
const toolMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'tool',
|
|
toolCallId: event.toolCallId,
|
|
toolName: event.toolName,
|
|
toolInput: event.toolInput,
|
|
};
|
|
session.messages.push(toolMsg);
|
|
session.meta.messageCount += 1;
|
|
break;
|
|
}
|
|
|
|
case 'tool:result': {
|
|
sendToClient(ws, {
|
|
type: 'tool:result',
|
|
toolCallId: event.toolCallId,
|
|
output: event.output,
|
|
isError: event.isError,
|
|
});
|
|
|
|
// Update existing tool message with output
|
|
for (let i = session.messages.length - 1; i >= 0; i--) {
|
|
const m = session.messages[i]!;
|
|
if (m.role === 'tool' && m.toolCallId === event.toolCallId) {
|
|
m.output = event.output;
|
|
m.isError = event.isError;
|
|
break;
|
|
}
|
|
}
|
|
break;
|
|
}
|
|
|
|
case 'result': {
|
|
// Flush any remaining streaming buffer
|
|
if (session.streamBuffer) {
|
|
sendToClient(ws, { type: 'assistant:text', text: session.streamBuffer });
|
|
|
|
const assistantMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'assistant',
|
|
text: session.streamBuffer,
|
|
model,
|
|
cost: event.cost,
|
|
};
|
|
session.messages.push(assistantMsg);
|
|
session.meta.messageCount += 1;
|
|
session.streamBuffer = '';
|
|
}
|
|
|
|
sendToClient(ws, { type: 'result', sessionId, cost: event.cost });
|
|
|
|
session.isGenerating = false;
|
|
session.meta.cost.inputTokens += event.cost.inputTokens;
|
|
session.meta.cost.outputTokens += event.cost.outputTokens;
|
|
session.meta.cost.totalUSD += event.cost.totalUSD;
|
|
session.meta.updatedAt = Date.now();
|
|
|
|
// Save session to disk
|
|
try {
|
|
await storage.saveSession(storageDir, sessionId, session.meta, session.messages);
|
|
logger.info('Session saved to disk', { sessionId, messageCount: session.messages.length });
|
|
} catch (err) {
|
|
logger.error('Failed to save session', { sessionId, error: String(err) });
|
|
}
|
|
break;
|
|
}
|
|
|
|
case 'error': {
|
|
sendToClient(ws, { type: 'error', message: event.message });
|
|
session.isGenerating = false;
|
|
break;
|
|
}
|
|
|
|
case 'stopped': {
|
|
sendToClient(ws, { type: 'stopped' });
|
|
session.isGenerating = false;
|
|
break;
|
|
}
|
|
}
|
|
};
|
|
}
|
|
|
|
async function handleChat(
|
|
ws: ServerWebSocket<WSData>,
|
|
msg: {
|
|
prompt: string;
|
|
displayText?: string;
|
|
sessionId?: string;
|
|
model?: string;
|
|
cwd?: string;
|
|
cwdRoot?: string;
|
|
sandboxed?: boolean;
|
|
groupSlug?: string;
|
|
attachmentIds?: string[];
|
|
thinking?: string;
|
|
context?: string;
|
|
contextId?: string;
|
|
},
|
|
): Promise<void> {
|
|
const { email, username, userId } = ws.data;
|
|
const sessionId = msg.sessionId || randomUUID();
|
|
|
|
// Use provided model, or fall back to user default, or use system default
|
|
let model = msg.model;
|
|
let modelSource = 'client-provided';
|
|
let userDefault = null;
|
|
if (!model) {
|
|
userDefault = await getUserDefaultModel(userId);
|
|
if (userDefault) {
|
|
model = userDefault;
|
|
modelSource = 'user-settings';
|
|
} else {
|
|
model = DEFAULT_MODEL;
|
|
modelSource = 'system-default';
|
|
}
|
|
}
|
|
|
|
logger.info('Model selected for chat', {
|
|
sessionId,
|
|
model,
|
|
modelSource,
|
|
clientModel: msg.model || null,
|
|
userDefault,
|
|
});
|
|
|
|
if (model === 'claude-code') {
|
|
return handleClaudeCodeChat(ws, sessionId, model, msg);
|
|
}
|
|
|
|
const homeDir = getHomeDirForRole(email, ws.data.role);
|
|
const cwd = resolveCwd(email, ws.data.role, msg.cwd);
|
|
const groupSlug = msg.groupSlug || null;
|
|
const session = sessionManager.getOrCreate(sessionId, email, cwd, model, groupSlug, msg.context, msg.contextId);
|
|
session.userId = userId;
|
|
sessionManager.attachWs(sessionId, ws);
|
|
wsToSessionMap.set(ws as any, sessionId);
|
|
|
|
sendToClient(ws, {
|
|
type: 'session:init',
|
|
sessionId,
|
|
model,
|
|
cwd,
|
|
context: session.meta.context,
|
|
contextId: session.meta.contextId,
|
|
});
|
|
if (!session.piProcess) {
|
|
try {
|
|
const onEvent = createEventHandler(sessionId, model, cwd, homeDir);
|
|
|
|
// If session has history, save to disk and pass --session for context replay
|
|
let spawnOptions: { sessionFile?: string; username?: string; role?: string } | undefined;
|
|
if (session.messages.length > 0) {
|
|
await storage.saveSession(homeDir, sessionId, session.meta, session.messages);
|
|
const hostPath = await storage.getSessionFilePath(homeDir, sessionId);
|
|
if (hostPath) {
|
|
spawnOptions = { sessionFile: hostPath, username, role: ws.data.role };
|
|
}
|
|
}
|
|
if (!spawnOptions) spawnOptions = { username, role: ws.data.role };
|
|
|
|
session.piProcess = await piBridge.spawnPi(cwd, model, userId, email, onEvent, spawnOptions);
|
|
|
|
// Null out piProcess when the process dies so next message triggers respawn
|
|
const proc = session.piProcess;
|
|
proc.exited.then(() => {
|
|
if (session.piProcess === proc) {
|
|
session.piProcess = null;
|
|
logger.info('Pi process exited, nulled reference', { sessionId });
|
|
}
|
|
});
|
|
|
|
logger.info('Spawned Pi process for session', {
|
|
sessionId,
|
|
model,
|
|
cwd,
|
|
hasSessionFile: !!spawnOptions?.sessionFile,
|
|
});
|
|
} catch (err) {
|
|
logger.error('Failed to spawn Pi process', { sessionId, model, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Failed to start Pi process' });
|
|
return;
|
|
}
|
|
}
|
|
|
|
// Add user message to session
|
|
const userMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'user',
|
|
text: msg.prompt,
|
|
};
|
|
session.messages.push(userMsg);
|
|
session.meta.messageCount += 1;
|
|
session.meta.updatedAt = Date.now();
|
|
|
|
if (!session.meta.title) {
|
|
session.meta.title = (msg.displayText ?? msg.prompt).slice(0, 100);
|
|
}
|
|
|
|
// Set thinking level if provided
|
|
console.log(`[pi] model: ${msg.model ?? 'default'}, thinking: ${msg.thinking ?? 'not set'}`);
|
|
if (msg.thinking) {
|
|
piBridge.setThinkingLevel(session.piProcess, msg.thinking);
|
|
}
|
|
|
|
// Send prompt to Pi
|
|
const requestId = randomUUID();
|
|
session.isGenerating = true;
|
|
piBridge.sendPrompt(session.piProcess, msg.prompt, requestId);
|
|
}
|
|
|
|
async function handleClaudeCodeChat(
|
|
ws: ServerWebSocket<WSData>,
|
|
sessionId: string,
|
|
model: string,
|
|
msg: {
|
|
prompt: string;
|
|
displayText?: string;
|
|
groupSlug?: string;
|
|
context?: string;
|
|
contextId?: string;
|
|
cwd?: string;
|
|
cwdRoot?: string;
|
|
sandboxed?: boolean;
|
|
},
|
|
): Promise<void> {
|
|
const { email, username, userId } = ws.data;
|
|
const homeDir = getHomeDirForRole(email, ws.data.role);
|
|
|
|
const cwd = resolveCwd(email, ws.data.role, msg.cwd);
|
|
|
|
const groupSlug = msg.groupSlug || null;
|
|
|
|
const session = sessionManager.getOrCreate(sessionId, email, cwd, model, groupSlug, msg.context, msg.contextId);
|
|
session.userId = userId;
|
|
sessionManager.attachWs(sessionId, ws);
|
|
wsToSessionMap.set(ws as any, sessionId);
|
|
|
|
sendToClient(ws, {
|
|
type: 'session:init',
|
|
sessionId,
|
|
model,
|
|
cwd,
|
|
context: session.meta.context,
|
|
contextId: session.meta.contextId,
|
|
});
|
|
|
|
// Add user message to session
|
|
const userMsg: Message = {
|
|
id: randomUUID(),
|
|
timestamp: Date.now(),
|
|
role: 'user',
|
|
text: msg.prompt,
|
|
};
|
|
session.messages.push(userMsg);
|
|
session.meta.messageCount += 1;
|
|
session.meta.updatedAt = Date.now();
|
|
|
|
if (!session.meta.title) {
|
|
session.meta.title = (msg.displayText ?? msg.prompt).slice(0, 100);
|
|
}
|
|
|
|
session.isGenerating = true;
|
|
|
|
const onEvent = createEventHandler(sessionId, model, cwd, homeDir);
|
|
|
|
try {
|
|
const handle = await sendClaudeCodeStreaming({
|
|
userId,
|
|
email,
|
|
username,
|
|
prompt: msg.prompt,
|
|
sessionKey: sessionId,
|
|
cwd,
|
|
onEvent,
|
|
});
|
|
|
|
// Store proc as piProcess so handleStop can kill it
|
|
session.piProcess = handle.proc;
|
|
|
|
// Null out when process exits so next message spawns a new one
|
|
handle.proc.exited.then(() => {
|
|
if (session.piProcess === handle.proc) {
|
|
session.piProcess = null;
|
|
}
|
|
});
|
|
} catch (err) {
|
|
logger.error('Failed to start Claude Code streaming', { sessionId, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Failed to start Claude Code' });
|
|
session.isGenerating = false;
|
|
}
|
|
}
|
|
|
|
async function handleResume(
|
|
ws: ServerWebSocket<WSData>,
|
|
msg: { sessionId: string; cwd?: string; cwdRoot?: string },
|
|
): Promise<void> {
|
|
const { email } = ws.data;
|
|
const { sessionId } = msg;
|
|
|
|
try {
|
|
let session = sessionManager.getSession(sessionId);
|
|
|
|
if (!session) {
|
|
const homeDir = getHomeDirForRole(email, ws.data.role);
|
|
|
|
try {
|
|
const { meta, messages } = await storage.loadSession(homeDir, sessionId);
|
|
|
|
session = sessionManager.getOrCreate(sessionId, email, meta.cwd, meta.model);
|
|
session.messages = messages;
|
|
session.meta = meta;
|
|
|
|
logger.info('Loaded session from disk', { sessionId, messageCount: messages.length });
|
|
} catch (err) {
|
|
logger.error('Failed to load session from disk', { sessionId, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Session not found', errorCode: 'SESSION_NOT_FOUND' });
|
|
return;
|
|
}
|
|
}
|
|
|
|
sessionManager.attachWs(sessionId, ws);
|
|
wsToSessionMap.set(ws as any, sessionId);
|
|
|
|
sendToClient(ws, {
|
|
type: 'session:init',
|
|
sessionId,
|
|
model: session.model,
|
|
cwd: session.cwd,
|
|
context: session.meta.context,
|
|
contextId: session.meta.contextId,
|
|
});
|
|
|
|
// Spawn fresh Pi process if needed
|
|
if (!session.piProcess) {
|
|
try {
|
|
const homeDir = getHomeDirForRole(email, ws.data.role);
|
|
const onEvent = createEventHandler(sessionId, session.model, session.cwd, homeDir);
|
|
|
|
// If session has history, save to disk and pass --session for context replay
|
|
let spawnOptions: { sessionFile?: string; username?: string; role?: string } | undefined;
|
|
if (session.messages.length > 0) {
|
|
await storage.saveSession(homeDir, sessionId, session.meta, session.messages);
|
|
const hostPath = await storage.getSessionFilePath(homeDir, sessionId);
|
|
if (hostPath) {
|
|
spawnOptions = { sessionFile: hostPath, username: ws.data.username, role: ws.data.role };
|
|
}
|
|
}
|
|
if (!spawnOptions) spawnOptions = { username: ws.data.username, role: ws.data.role };
|
|
|
|
session.piProcess = await piBridge.spawnPi(
|
|
session.cwd,
|
|
session.model,
|
|
session.userId!,
|
|
email,
|
|
onEvent,
|
|
spawnOptions,
|
|
);
|
|
|
|
// Null out piProcess when the process dies so next message triggers respawn
|
|
const proc = session.piProcess;
|
|
proc.exited.then(() => {
|
|
if (session.piProcess === proc) {
|
|
session.piProcess = null;
|
|
logger.info('Pi process exited, nulled reference', { sessionId });
|
|
}
|
|
});
|
|
|
|
logger.info('Spawned fresh Pi process for resumed session', {
|
|
sessionId,
|
|
model: session.model,
|
|
hasSessionFile: !!spawnOptions?.sessionFile,
|
|
});
|
|
} catch (err) {
|
|
logger.error('Failed to spawn Pi process for resume', { sessionId, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Failed to start Pi process' });
|
|
return;
|
|
}
|
|
}
|
|
|
|
sendToClient(ws, {
|
|
type: 'sync:messages',
|
|
sessionId,
|
|
messages: session.messages,
|
|
isGenerating: session.isGenerating,
|
|
streamingText: session.streamBuffer,
|
|
});
|
|
|
|
logger.info('Session resumed successfully', { sessionId, messageCount: session.messages.length });
|
|
} catch (err) {
|
|
logger.error('Unexpected error in handleResume', { sessionId, error: String(err) });
|
|
sendToClient(ws, { type: 'error', message: 'Failed to resume session' });
|
|
}
|
|
}
|
|
|
|
async function handleStop(ws: ServerWebSocket<WSData>): Promise<void> {
|
|
const sessionId = wsToSessionMap.get(ws);
|
|
|
|
if (sessionId) {
|
|
const session = sessionManager.getSession(sessionId);
|
|
|
|
if (session?.piProcess) {
|
|
try {
|
|
if (session.model === 'claude-code') {
|
|
// Claude Code: kill the process directly
|
|
session.piProcess.kill();
|
|
logger.info('Killed Claude Code process', { sessionId });
|
|
} else {
|
|
piBridge.abort(session.piProcess, randomUUID());
|
|
logger.info('Sent abort to Pi process', { sessionId });
|
|
}
|
|
session.isGenerating = false;
|
|
} catch (err) {
|
|
logger.error('Failed to stop process', { sessionId, error: String(err) });
|
|
}
|
|
}
|
|
}
|
|
|
|
sendToClient(ws, { type: 'stopped' });
|
|
}
|
|
|
|
export const piWebsocket = {
|
|
open,
|
|
message,
|
|
close,
|
|
drain() {},
|
|
};
|