diff --git a/app/(app)/_layout.tsx b/app/(app)/_layout.tsx index ac7d6bc..46b17aa 100644 --- a/app/(app)/_layout.tsx +++ b/app/(app)/_layout.tsx @@ -188,9 +188,12 @@ export default function AppLayout() { }); }, [refreshActiveServerSession]); - const onApiAuthError = useCallback(async () => { - // Token expired on an API call (prompt, steer, etc.) — refresh and retry - return refreshActiveServerSession(); + const onApiAuthError = useCallback(async (): Promise => { + const ok = await refreshActiveServerSession(); + if (!ok) return null; + const state = useAuthStore.getState(); + const sid = state.activeServerId; + return sid ? state.tokens[sid]?.accessToken ?? null : null; }, [refreshActiveServerSession]); const piClientConfig = useMemo( diff --git a/app/(app)/chat/[sessionId].tsx b/app/(app)/chat/[sessionId].tsx index 665867c..ffcc3d2 100644 --- a/app/(app)/chat/[sessionId].tsx +++ b/app/(app)/chat/[sessionId].tsx @@ -20,9 +20,11 @@ import { ExtensionUiDialog } from '@/features/agent/components/extension-ui-dial import { DiffPanelProvider } from '@/features/agent/components/diff-panel/context'; import { MobileDiffSheetProvider } from '@/features/agent/components/message-list/mobile-diff-sheet'; import { useAgentSession, useChatSessions, useConnection } from '@pi-ui/client'; +import type { ImageContent } from '@pi-ui/client'; import { useChatStore } from '@/features/chat/store'; import { useWorkspaceStore } from '@/features/workspace/store'; import type { PendingExtensionUiRequest as LegacyPendingUiRequest } from '@/features/agent/extension-ui'; +import type { Attachment } from '@/features/workspace/components/prompt-input/constants'; export default function ChatSessionScreen() { const { sessionId } = useLocalSearchParams<{ sessionId: string }>(); @@ -51,10 +53,24 @@ export default function ChatSessionScreen() { }, [sessionId, selectSession, registerSessionWorkspace]); const handleSend = useCallback( - async (text: string, _attachments: unknown[], options?: { queueBehavior?: 'steer' | 'followUp' }) => { + async (text: string, attachments: Attachment[], options?: { queueBehavior?: 'steer' | 'followUp' }) => { if (!sessionId || inputBlocked) return; setAlertMessage(null); + let images: ImageContent[] | undefined; + const imageAttachments = attachments.filter((a) => a.type === "image" && a.preview); + if (imageAttachments.length > 0) { + images = imageAttachments.map((a) => { + const dataUrl = a.preview!; + const commaIdx = dataUrl.indexOf(","); + const meta = dataUrl.slice(0, commaIdx); + const base64 = dataUrl.slice(commaIdx + 1); + const mimeMatch = meta.match(/data:([^;]+)/); + const mimeType = mimeMatch?.[1] ?? "image/png"; + return { type: "image" as const, data: base64, mimeType }; + }); + } + const isFirst = !session.messages.length; const behavior = options?.queueBehavior ?? (session.isStreaming ? 'steer' : undefined); const sendFn = behavior === 'steer' @@ -64,7 +80,7 @@ export default function ChatSessionScreen() { : session.prompt; try { - await sendFn(text); + await sendFn(text, images ? { images } : undefined); if (isFirst) setTimeout(() => invalidateChatSessions(), 2000); } catch (error) { setAlertMessage(error instanceof Error ? error.message : 'Failed to send prompt'); diff --git a/app/(app)/servers.tsx b/app/(app)/servers.tsx index e7a7b0b..775bea2 100644 --- a/app/(app)/servers.tsx +++ b/app/(app)/servers.tsx @@ -162,6 +162,8 @@ function ServerFormFields({ ); } +type ServerFormData = Omit & { username: string; password: string }; + function ServerFormDesktopModal({ visible, onClose, @@ -173,7 +175,7 @@ function ServerFormDesktopModal({ }: { visible: boolean; onClose: () => void; - onSave: (data: Omit) => void; + onSave: (data: ServerFormData) => void; initial?: Server; isDark: boolean; loading?: boolean; @@ -189,8 +191,8 @@ function ServerFormDesktopModal({ if (visible) { setName(initial?.name ?? ""); setAddress(initial?.address ?? ""); - setUsername(initial?.username ?? ""); - setPassword(initial?.password ?? ""); + setUsername(""); + setPassword(""); setShowPassword(false); } }, [visible, initial]); @@ -301,7 +303,7 @@ function ServerFormSheet({ }: { visible: boolean; onClose: () => void; - onSave: (data: Omit) => void; + onSave: (data: ServerFormData) => void; initial?: Server; isDark: boolean; loading?: boolean; @@ -336,8 +338,8 @@ function ServerFormSheet({ if (visible) { setName(initial?.name ?? ""); setAddress(initial?.address ?? ""); - setUsername(initial?.username ?? ""); - setPassword(initial?.password ?? ""); + setUsername(""); + setPassword(""); setShowPassword(false); translateY.value = withTiming(0, TIMING_CONFIG); overlayOpacity.value = withTiming(1, TIMING_CONFIG); @@ -539,7 +541,7 @@ function ServerFormSheet({ function ServerFormModal(props: { visible: boolean; onClose: () => void; - onSave: (data: Omit) => void; + onSave: (data: ServerFormData) => void; initial?: Server; isDark: boolean; loading?: boolean; @@ -636,24 +638,16 @@ export default function ServersScreen() { let connected = false; - // Has stored token — activate and verify session if (auth.hasToken(server.id)) { connected = await auth.activateServer(server); } - // Try login with stored credentials - if (!connected && server.username) { - const result = await auth.loginToServer(server); - connected = result.success; - } - if (connected) { router.replace('/'); } else { - // No valid token and no credentials — prompt to edit/re-enter credentials setConnecting(null); setEditingServer(server); - setLoginError("Not connected. Enter credentials to connect."); + setLoginError("Session expired. Enter credentials to reconnect."); setFormVisible(true); } } catch (e) { @@ -664,22 +658,22 @@ export default function ServersScreen() { }; const handleSave = useCallback( - async (data: Omit) => { + async (data: Omit & { username: string; password: string }) => { setLoginLoading(true); setLoginError(null); + const { username, password, ...serverData } = data; let server: Server; if (editingServer) { - await updateServer(editingServer.id, data); - server = { ...editingServer, ...data }; + await updateServer(editingServer.id, serverData); + server = { ...editingServer, ...serverData }; } else { - await addServer(data); - // get the newly added server (last in list after addServer) + await addServer(serverData); const servers = useServersStore.getState().servers; server = servers[servers.length - 1]; } - const result = await loginToServer(server); + const result = await loginToServer(server, { username, password }); setLoginLoading(false); if (result.success) { diff --git a/app/(app)/workspace/[workspaceId]/s/[sessionId].tsx b/app/(app)/workspace/[workspaceId]/s/[sessionId].tsx index 83649d2..d26b5de 100644 --- a/app/(app)/workspace/[workspaceId]/s/[sessionId].tsx +++ b/app/(app)/workspace/[workspaceId]/s/[sessionId].tsx @@ -23,9 +23,11 @@ import { DiffPanelProvider } from "@/features/agent/components/diff-panel/contex import { DiffSidebar } from "@/features/agent/components/diff-panel"; import { MobileDiffSheetProvider } from "@/features/agent/components/message-list/mobile-diff-sheet"; import { useAgentSession, useConnection, useWorkspaceSessions as useSessions } from "@pi-ui/client"; +import type { ImageContent } from "@pi-ui/client"; import { requestBrowserNotificationPermission } from "@/features/agent/browser-notifications"; import type { PendingExtensionUiRequest as LegacyPendingUiRequest } from "@/features/agent/extension-ui"; import type { ChatMessage } from "@/features/agent/types"; +import type { Attachment } from "@/features/workspace/components/prompt-input/constants"; export default function SessionScreen() { const { workspaceId, sessionId } = useLocalSearchParams<{ @@ -76,13 +78,27 @@ export default function SessionScreen() { const handleSend = useCallback( async ( text: string, - _attachments: unknown[], + attachments: Attachment[], options?: { queueBehavior?: "steer" | "followUp" }, ) => { if (!sessionId || inputBlockedByConnection) return; setAlertMessage(null); requestBrowserNotificationPermission(); + let images: ImageContent[] | undefined; + const imageAttachments = attachments.filter((a) => a.type === "image" && a.preview); + if (imageAttachments.length > 0) { + images = imageAttachments.map((a) => { + const dataUrl = a.preview!; + const commaIdx = dataUrl.indexOf(","); + const meta = dataUrl.slice(0, commaIdx); + const base64 = dataUrl.slice(commaIdx + 1); + const mimeMatch = meta.match(/data:([^;]+)/); + const mimeType = mimeMatch?.[1] ?? "image/png"; + return { type: "image" as const, data: base64, mimeType }; + }); + } + const behavior = options?.queueBehavior ?? (agentSession.isStreaming ? "steer" : undefined); const sendFn = behavior === "steer" ? agentSession.steer @@ -91,7 +107,7 @@ export default function SessionScreen() { : agentSession.prompt; try { - await sendFn(text); + await sendFn(text, images ? { images } : undefined); } catch (error) { setAlertMessage( error instanceof Error ? error.message : "Failed to send prompt", diff --git a/app/connect.tsx b/app/connect.tsx index e8b37b4..a9d8ee5 100644 --- a/app/connect.tsx +++ b/app/connect.tsx @@ -101,8 +101,6 @@ export default function DirectConnectScreen() { id: serverId, name: existingServer?.name || connectParams.hostname || "Pico Server", address: baseUrl, - username: existingServer?.username ?? "", - password: existingServer?.password ?? "", }); await fetchWorkspaces(); diff --git a/assets/images/android-icon-background.png b/assets/images/android-icon-background.png index 5ffefc5..209d1a4 100644 Binary files a/assets/images/android-icon-background.png and b/assets/images/android-icon-background.png differ diff --git a/assets/images/android-icon-foreground.png b/assets/images/android-icon-foreground.png index 95d76df..96c9c1c 100644 Binary files a/assets/images/android-icon-foreground.png and b/assets/images/android-icon-foreground.png differ diff --git a/assets/images/android-icon-monochrome.png b/assets/images/android-icon-monochrome.png index 77484eb..120d2b8 100644 Binary files a/assets/images/android-icon-monochrome.png and b/assets/images/android-icon-monochrome.png differ diff --git a/assets/images/favicon.png b/assets/images/favicon.png index 1266b1a..949b3ff 100644 Binary files a/assets/images/favicon.png and b/assets/images/favicon.png differ diff --git a/assets/images/icon.png b/assets/images/icon.png index b30c43c..2712811 100644 Binary files a/assets/images/icon.png and b/assets/images/icon.png differ diff --git a/assets/images/splash-icon.png b/assets/images/splash-icon.png index 0f46b03..2712811 100644 Binary files a/assets/images/splash-icon.png and b/assets/images/splash-icon.png differ diff --git a/backend/src/models/agent.rs b/backend/src/models/agent.rs index 73710fe..7042f6a 100644 --- a/backend/src/models/agent.rs +++ b/backend/src/models/agent.rs @@ -167,16 +167,30 @@ pub struct WsStreamQuery { } #[derive(Debug, Deserialize, ToSchema)] -pub struct SessionStreamQuery { - pub last_message_id: Option, - pub before: Option, - pub limit: Option, -} +pub struct SessionStreamQuery {} #[derive(Debug, Deserialize)] pub struct WsSessionStreamQuery { - pub last_message_id: Option, + pub access_token: Option, +} + +#[derive(Debug, Deserialize, ToSchema)] +pub struct SetActiveSessionRequest { + pub connection_id: String, + pub session_id: Option, + pub from_event_id: Option, + pub from_delta_event_id: Option, +} + +#[derive(Debug, Deserialize, ToSchema)] +pub struct SessionHistoryQuery { pub before: Option, pub limit: Option, - pub access_token: Option, +} + +#[derive(Debug, Serialize, ToSchema)] +pub struct SessionHistoryResponse { + pub messages: Vec, + pub has_more: bool, + pub oldest_entry_id: Option, } diff --git a/backend/src/routes/agent.rs b/backend/src/routes/agent.rs index 22a59f8..0ddb795 100644 --- a/backend/src/routes/agent.rs +++ b/backend/src/routes/agent.rs @@ -487,11 +487,15 @@ async fn handle_ws_stream( return; } + let (connection_id, conn, mut inject_rx) = state.sse_registry.register().await; + let hello = serde_json::json!({ "type": "server_hello", "instance_id": *state.instance_id, + "connection_id": connection_id, }); if !send_ws_batch(&mut socket, vec![hello]).await { + state.sse_registry.unregister(&connection_id).await; return; } @@ -504,6 +508,7 @@ async fn handle_ws_stream( replay_payloads.push(stream_event_value(&event)); if replay_payloads.len() >= WS_MAX_BATCH_EVENTS { if !send_ws_batch(&mut socket, replay_payloads).await { + state.sse_registry.unregister(&connection_id).await; return; } replay_payloads = Vec::with_capacity(WS_MAX_BATCH_EVENTS); @@ -512,6 +517,7 @@ async fn handle_ws_stream( if !replay_payloads.is_empty() && !send_ws_batch(&mut socket, replay_payloads).await { + state.sse_registry.unregister(&connection_id).await; return; } @@ -528,6 +534,7 @@ async fn handle_ws_stream( }; let payload = stream_event_value(&ports_event); if !send_ws_batch(&mut socket, vec![payload]).await { + state.sse_registry.unregister(&connection_id).await; return; } } @@ -536,6 +543,7 @@ async fn handle_ws_stream( let mut keepalive = tokio::time::interval(Duration::from_secs(WS_KEEPALIVE_SECS)); keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + let mut inject_hwm: u64 = 0; loop { tokio::select! { @@ -558,12 +566,35 @@ async fn handle_ws_stream( Some(Err(_)) => break, } } + injected = inject_rx.recv() => { + match injected { + Some(value) => { + if let Some(id) = value.get("id").and_then(|v| v.as_u64()) { + if id > inject_hwm { + inject_hwm = id; + } + } + if !send_ws_batch(&mut socket, vec![value]).await { + break; + } + } + None => break, + } + } result = rx.recv() => { let mut payloads = match result { Ok(event) => { - if !is_global_event(&event.event_type) { + if inject_hwm > 0 && event.id <= inject_hwm { + continue; + } + + if is_session_only_event(&event.event_type) + && is_skippable_for_inactive(&event.event_type) + && !conn.is_session_receiving_deltas(&event.session_id).await + { continue; } + vec![strip_live_event(stream_event_value(&event))] } Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { @@ -572,6 +603,8 @@ async fn handle_ws_stream( Err(tokio::sync::broadcast::error::RecvError::Closed) => break, }; + // Note: drain_ws_pending_payloads doesn't do per-connection filtering + // but that's acceptable for batching — the main filter is above let closed = drain_ws_pending_payloads(&mut rx, &mut payloads); if !send_ws_batch(&mut socket, payloads).await { @@ -584,6 +617,8 @@ async fn handle_ws_stream( } } } + + state.sse_registry.unregister(&connection_id).await; } // --- Session Management --- @@ -1069,8 +1104,122 @@ pub async fn preview_proxy_path( .await } +// --- Session History (REST) --- + +#[utoipa::path( + get, + path = "/api/sessions/{session_id}/history", + params( + ("session_id" = String, Path, description = "Session ID"), + ("before" = Option, Query, description = "Load messages before this entry ID"), + ("limit" = Option, Query, description = "Max messages to return (default 20)"), + ), + responses( + (status = 200, description = "Paginated session messages", body = SessionHistoryResponse), + (status = 401, description = "Unauthorized"), + (status = 404, description = "Session not found"), + ), + security(("bearer_auth" = [])), + tag = "agent" +)] +pub async fn session_history( + State(state): State, + headers: HeaderMap, + Path(session_id): Path, + Query(params): Query, +) -> (StatusCode, Json>) { + if let Err((code, msg)) = require_auth(&state, &headers).await { + return (code, Json(ApiResponse::err(msg))); + } + + let base = state.config.sessions_base_path(); + let sid = session_id.clone(); + let limit = params.limit.unwrap_or(20); + let before = params.before.clone(); + + let result = tokio::task::spawn_blocking(move || { + session::get_session_messages_paginated(&base, &sid, limit, before.as_deref()) + }) + .await + .unwrap(); + + match result { + Some(paginated) => ( + StatusCode::OK, + Json(ApiResponse::ok(SessionHistoryResponse { + messages: paginated.messages, + has_more: paginated.has_more, + oldest_entry_id: paginated.oldest_entry_id, + })), + ), + None => ( + StatusCode::NOT_FOUND, + Json(ApiResponse::err("Session not found")), + ), + } +} + +// --- Active Session --- + +#[utoipa::path( + post, + path = "/api/stream-active-session", + request_body = SetActiveSessionRequest, + responses( + (status = 200, description = "Active session updated"), + (status = 401, description = "Unauthorized"), + (status = 404, description = "Connection not found"), + ), + security(("bearer_auth" = [])), + tag = "agent" +)] +pub async fn set_active_session( + State(state): State, + headers: HeaderMap, + Json(req): Json, +) -> (StatusCode, Json>) { + if let Err((code, msg)) = require_auth(&state, &headers).await { + return (code, Json(ApiResponse::err(msg))); + } + + let conn = match state.sse_registry.get(&req.connection_id).await { + Some(c) => c, + None => { + return ( + StatusCode::NOT_FOUND, + Json(ApiResponse::err("Connection not found")), + ); + } + }; + + conn.set_active(req.session_id.clone()).await; + + if let Some(ref session_id) = req.session_id { + let buffered = state + .agent + .get_buffered_session_events(session_id, req.from_event_id, req.from_delta_event_id) + .await; + let collapsed = collapse_buffered_events(buffered); + for event in collapsed { + let val = stream_event_value(&event); + if conn.inject_tx.send(val).is_err() { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(ApiResponse::err("Connection closed")), + ); + } + } + } + + (StatusCode::OK, Json(ApiResponse::ok(json!({ "ok": true })))) +} + // --- SSE Stream --- +fn is_skippable_for_inactive(event_type: &str) -> bool { + event_type == "message_update" || event_type == "tool_execution_update" +} + #[utoipa::path( get, path = "/api/stream", @@ -1095,16 +1244,21 @@ pub async fn stream( let mut rx = state.agent.subscribe(); let instance_id = state.instance_id.clone(); let port_scanner = state.port_scanner.clone(); + let (connection_id, conn, mut inject_rx) = state.sse_registry.register().await; + let registry = state.sse_registry.clone(); + let conn_id_clone = connection_id.clone(); let stream = async_stream::stream! { let hello = serde_json::json!({ "type": "server_hello", "instance_id": *instance_id, + "connection_id": connection_id, }); yield Ok::<_, Infallible>( Event::default().data(serde_json::to_string(&hello).unwrap_or_default()), ); + // Replay buffered global events (no active session set yet, so skip session-only) for event in replay_events { if !is_global_event(&event.event_type) { continue; @@ -1115,7 +1269,7 @@ pub async fn stream( ); } - // Send current port state so the client knows about open ports immediately + // Send current port state { let ports_data = port_scanner.get_current_ports_event().await; let ports_event = StreamEvent { @@ -1130,26 +1284,61 @@ pub async fn stream( yield Ok::<_, Infallible>(Event::default().data(data)); } + let mut inject_hwm: u64 = 0; + loop { - match rx.recv().await { - Ok(event) => { - if !is_global_event(&event.event_type) { - continue; + tokio::select! { + result = rx.recv() => { + match result { + Ok(event) => { + if inject_hwm > 0 && event.id <= inject_hwm { + continue; + } + + if is_session_only_event(&event.event_type) + && is_skippable_for_inactive(&event.event_type) + && !conn.is_session_receiving_deltas(&event.session_id).await + { + continue; + } + + let val = strip_live_event(stream_event_value(&event)); + let data = serde_json::to_string(&val).unwrap_or_default(); + yield Ok::<_, Infallible>( + Event::default().id(event.id.to_string()).data(data), + ); + } + Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { + yield Ok::<_, Infallible>( + Event::default().data(stream_lagged_json(n.into())), + ); + } + Err(_) => break, } - let val = strip_live_event(stream_event_value(&event)); - let data = serde_json::to_string(&val).unwrap_or_default(); - yield Ok::<_, Infallible>( - Event::default().id(event.id.to_string()).data(data), - ); } - Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { - yield Ok::<_, Infallible>( - Event::default().data(stream_lagged_json(n.into())), - ); + injected = inject_rx.recv() => { + match injected { + Some(value) => { + let data = serde_json::to_string(&value).unwrap_or_default(); + let id = value.get("id").and_then(|v| v.as_u64()); + if let Some(id) = id { + if id > inject_hwm { + inject_hwm = id; + } + } + let mut event = Event::default().data(data); + if let Some(id) = id { + event = event.id(id.to_string()); + } + yield Ok::<_, Infallible>(event); + } + None => break, + } } - Err(_) => break, } } + + registry.unregister(&conn_id_clone).await; }; Sse::new(stream) @@ -1210,31 +1399,6 @@ fn collapse_buffered_events(events: Vec) -> Vec { result } -fn build_history_replay_events( - messages: Vec, - session_id: &str, - workspace_id: &str, - has_more: bool, - oldest_entry_id: Option, -) -> Vec { - let now = chrono::Utc::now().timestamp_millis(); - let data = json!({ - "type": "history_messages", - "messages": messages, - "has_more": has_more, - "oldest_entry_id": oldest_entry_id, - }); - - vec![StreamEvent { - id: 0, - session_id: session_id.to_string(), - workspace_id: workspace_id.to_string(), - event_type: "history_messages".to_string(), - data, - timestamp: now, - }] -} - #[utoipa::path( get, path = "/api/stream/{session_id}", @@ -1254,56 +1418,19 @@ pub async fn session_stream( State(state): State, headers: HeaderMap, Path(session_id): Path, - Query(params): Query, + Query(_params): Query, ) -> impl IntoResponse { if let Err((code, msg)) = require_auth(&state, &headers).await { return (code, Json(ApiResponse::::err(msg))).into_response(); } - let session_info = state.agent.get_session_info(&session_id).await; - let workspace_id = session_info - .as_ref() - .map(|s| s.workspace_id.clone()) - .unwrap_or_default(); - let mut rx = state.agent.subscribe(); - let skip_history = params.last_message_id.as_deref() == Some("SKIP_HISTORY"); - let msg_limit = params.limit.unwrap_or(20); - let before_cursor = params.before.clone(); - - let (history_messages, has_more, oldest_entry_id) = if skip_history { - (vec![], false, None) - } else { - let base = state.config.sessions_base_path(); - let sid = session_id.clone(); - let last_msg_id = params.last_message_id.clone(); - let before = before_cursor.clone(); - let result = tokio::task::spawn_blocking(move || { - if last_msg_id.is_some() { - let msgs = session::get_session_messages_after(&base, &sid, last_msg_id.as_deref()) - .unwrap_or_default(); - session::PaginatedMessages { messages: msgs, has_more: false, oldest_entry_id: None } - } else { - session::get_session_messages_paginated(&base, &sid, msg_limit, before.as_deref()) - .unwrap_or(session::PaginatedMessages { messages: vec![], has_more: false, oldest_entry_id: None }) - } - }) - .await - .unwrap(); - (result.messages, result.has_more, result.oldest_entry_id) - }; - - let history_events = build_history_replay_events( - history_messages, - &session_id, - &workspace_id, - has_more, - oldest_entry_id, - ); - let buffered_events = collapse_buffered_events( - state.agent.get_buffered_session_events(&session_id).await, + state + .agent + .get_buffered_session_events(&session_id, None, None) + .await, ); let high_water_mark = buffered_events.last().map(|e| e.id).unwrap_or(0); @@ -1318,18 +1445,6 @@ pub async fn session_stream( Event::default().data(serde_json::to_string(&hello).unwrap_or_default()), ); - for event in history_events { - let data = stream_event_json(&event); - yield Ok::<_, Infallible>( - Event::default().event("history").data(data), - ); - } - - let history_done = serde_json::json!({"type": "history_done"}); - yield Ok::<_, Infallible>( - Event::default().event("history_done").data(serde_json::to_string(&history_done).unwrap_or_default()), - ); - for event in buffered_events { let data = stream_event_json(&event); yield Ok::<_, Infallible>( @@ -1375,9 +1490,6 @@ async fn handle_ws_session_stream( state: AppState, session_id: String, access_token: Option, - last_message_id: Option, - before_cursor: Option, - msg_limit: u32, ) { let token = match access_token { Some(token) => token, @@ -1402,12 +1514,6 @@ async fn handle_ws_session_stream( return; } - let session_info = state.agent.get_session_info(&session_id).await; - let workspace_id = session_info - .as_ref() - .map(|s| s.workspace_id.clone()) - .unwrap_or_default(); - let hello = serde_json::json!({ "type": "session_stream_hello", "session_id": session_id, @@ -1418,59 +1524,11 @@ async fn handle_ws_session_stream( let mut rx = state.agent.subscribe(); - let skip_history = last_message_id.as_deref() == Some("SKIP_HISTORY"); - - let (history_messages, has_more, oldest_entry_id) = if skip_history { - (vec![], false, None) - } else { - let base = state.config.sessions_base_path(); - let sid = session_id.clone(); - let last_msg = last_message_id.clone(); - let before = before_cursor.clone(); - let result = tokio::task::spawn_blocking(move || { - if last_msg.is_some() { - let msgs = session::get_session_messages_after(&base, &sid, last_msg.as_deref()) - .unwrap_or_default(); - session::PaginatedMessages { messages: msgs, has_more: false, oldest_entry_id: None } - } else { - session::get_session_messages_paginated(&base, &sid, msg_limit, before.as_deref()) - .unwrap_or(session::PaginatedMessages { messages: vec![], has_more: false, oldest_entry_id: None }) - } - }) - .await - .unwrap(); - (result.messages, result.has_more, result.oldest_entry_id) - }; - - let history_events = build_history_replay_events( - history_messages, - &session_id, - &workspace_id, - has_more, - oldest_entry_id, - ); - - let mut history_payloads = Vec::with_capacity(WS_MAX_BATCH_EVENTS); - for event in history_events { - history_payloads.push(stream_event_value(&event)); - if history_payloads.len() >= WS_MAX_BATCH_EVENTS { - if !send_ws_batch(&mut socket, history_payloads).await { - return; - } - history_payloads = Vec::with_capacity(WS_MAX_BATCH_EVENTS); - } - } - if !history_payloads.is_empty() && !send_ws_batch(&mut socket, history_payloads).await { - return; - } - - let history_done = serde_json::json!({"type": "history_done"}); - if !send_ws_batch(&mut socket, vec![history_done]).await { - return; - } - let buffered_events = collapse_buffered_events( - state.agent.get_buffered_session_events(&session_id).await, + state + .agent + .get_buffered_session_events(&session_id, None, None) + .await, ); let high_water_mark = buffered_events.last().map(|e| e.id).unwrap_or(0); @@ -1559,7 +1617,7 @@ pub async fn ws_session_stream( .or(params.access_token); ws.protocols(["pi-stream-v1"]) .on_upgrade(move |socket| { - handle_ws_session_stream(socket, state, session_id, access_token, params.last_message_id, params.before, params.limit.unwrap_or(20)) + handle_ws_session_stream(socket, state, session_id, access_token) }) } diff --git a/backend/src/server/mod.rs b/backend/src/server/mod.rs index 9317adc..94c8173 100644 --- a/backend/src/server/mod.rs +++ b/backend/src/server/mod.rs @@ -112,6 +112,7 @@ pub async fn serve(cli: Cli, force_qr: bool) -> anyhow::Result<()> { desktop, http_client: reqwest::Client::new(), instance_id, + sse_registry: crate::services::sse_registry::SseConnectionRegistry::new(), }; let app = build_app(state); diff --git a/backend/src/server/openapi.rs b/backend/src/server/openapi.rs index 3f3e415..769c7f6 100644 --- a/backend/src/server/openapi.rs +++ b/backend/src/server/openapi.rs @@ -94,6 +94,8 @@ use crate::services; routes::agent::set_session_name, routes::agent::get_commands, routes::agent::extension_ui_response, + routes::agent::session_history, + routes::agent::set_active_session, routes::chat::create_session, routes::chat::list_sessions, routes::chat::delete_session, @@ -180,6 +182,9 @@ use crate::services; models::agent::AgentSetSessionNameRequest, models::agent::AgentNewSessionRequest, models::agent::AgentExtensionUiResponseRequest, + models::agent::SessionHistoryQuery, + models::agent::SessionHistoryResponse, + models::agent::SetActiveSessionRequest, services::agent::AgentSessionInfo, services::agent::ActiveSessionSummary, services::agent::StreamEvent, diff --git a/backend/src/server/router.rs b/backend/src/server/router.rs index 6db2c83..793cffd 100644 --- a/backend/src/server/router.rs +++ b/backend/src/server/router.rs @@ -146,7 +146,12 @@ fn agent_routes() -> Router { "/agent/sessions/{session_id}/preview/{hostname}/{port}/{*path}", any(routes::agent::preview_proxy_path), ) + .route( + "/sessions/{session_id}/history", + get(routes::agent::session_history), + ) .route("/stream", get(routes::agent::stream)) + .route("/stream-active-session", post(routes::agent::set_active_session)) .route("/ws/stream", get(routes::agent::ws_stream)) .route("/stream/{session_id}", get(routes::agent::session_stream)) .route("/ws/stream/{session_id}", get(routes::agent::ws_session_stream)) diff --git a/backend/src/server/state.rs b/backend/src/server/state.rs index 5a9424b..9bcb29e 100644 --- a/backend/src/server/state.rs +++ b/backend/src/server/state.rs @@ -8,6 +8,7 @@ use crate::services::agent::AgentManager; use crate::services::desktop::DesktopManager; use crate::services::pairing::PairingManager; use crate::services::port_scanner::PortScanner; +use crate::services::sse_registry::SseConnectionRegistry; use crate::services::task::TaskManager; #[derive(Clone, Debug, Serialize, Deserialize)] @@ -29,4 +30,5 @@ pub struct AppState { pub desktop: DesktopManager, pub http_client: reqwest::Client, pub instance_id: Arc, + pub sse_registry: SseConnectionRegistry, } diff --git a/backend/src/services/agent.rs b/backend/src/services/agent.rs index 64d41c3..1a6f125 100644 --- a/backend/src/services/agent.rs +++ b/backend/src/services/agent.rs @@ -13,7 +13,7 @@ use utoipa::ToSchema; use super::provider::{ AgentCapability, AgentCommand, AgentProcessHandle, AgentProvider, AgentSessionConfig, AgentStreamEvent, CommandResponse, ExtensionUiRequestKind, SessionSnapshot, - StreamingBehavior, + StreamingBehavior, TurnStats, }; const MAX_BUFFER_SIZE: usize = 10_000; @@ -481,15 +481,19 @@ impl AgentManager { } } - pub async fn get_buffered_session_events(&self, session_id: &str) -> Vec { + pub async fn get_buffered_session_events( + &self, + session_id: &str, + from_event_id: Option, + from_delta_event_id: Option, + ) -> Vec { let mut buffer = self.event_buffer.lock().await; - get_buffered_events_for_session(buffer.make_contiguous(), session_id) - } - - pub async fn get_session_info(&self, session_id: &str) -> Option { - let resolved_id = self.resolve_session_id(session_id).await; - let sessions = self.sessions.read().await; - sessions.get(&resolved_id).map(build_session_info) + get_buffered_events_for_session( + buffer.make_contiguous(), + session_id, + from_event_id, + from_delta_event_id, + ) } pub fn start_idle_cleanup_task(&self) { @@ -647,7 +651,7 @@ impl AgentManager { tokio::spawn(async move { let mut is_streaming = false; - while let Some(event) = event_rx.recv().await { + while let Some(mut event) = event_rx.recv().await { let is_exit = matches!(event, AgentStreamEvent::SessionProcessExited); let current_session_id = session_id_ref.lock().unwrap().clone(); @@ -660,6 +664,14 @@ impl AgentManager { _ => is_streaming, }; + if matches!(&event, AgentStreamEvent::AgentEnd { .. }) { + let buf = event_buffer.lock().await; + let stats = compute_turn_stats_from_buffer(&buf, ¤t_session_id); + if let AgentStreamEvent::AgentEnd { ref mut turn_stats, .. } = event { + *turn_stats = stats; + } + } + update_pending_extension_ui(&sessions, ¤t_session_id, &event).await; let (event_type, data) = stream_event_to_json(&event); @@ -1219,6 +1231,130 @@ fn stream_event_to_json(event: &AgentStreamEvent) -> (String, Value) { event.to_json() } +fn compute_turn_stats_from_buffer(buffer: &std::collections::VecDeque, session_id: &str) -> Option { + use std::collections::HashSet; + + let mut files_edited = HashSet::new(); + let mut files_created = HashSet::new(); + let mut lines_added: u32 = 0; + let mut lines_removed: u32 = 0; + + let mut agent_start_ts: Option = None; + let turn_events: Vec<&StreamEvent> = { + let mut events = Vec::new(); + for evt in buffer.iter().rev() { + if evt.session_id != session_id { + continue; + } + if evt.event_type == "agent_start" { + agent_start_ts = Some(evt.timestamp); + break; + } + events.push(evt); + } + events.reverse(); + events + }; + let now_ms = chrono::Utc::now().timestamp_millis(); + let duration_ms = agent_start_ts.map(|ts| now_ms - ts).unwrap_or(0); + + let mut tool_call_paths: HashMap = HashMap::new(); + let mut tool_call_content: HashMap = HashMap::new(); + for evt in &turn_events { + if evt.event_type == "tool_execution_start" { + let data = &evt.data; + let call_id = data.get("toolCallId").and_then(|v| v.as_str()).unwrap_or(""); + if call_id.is_empty() { continue; } + let path = extract_path_from_args(data); + if !path.is_empty() { + tool_call_paths.insert(call_id.to_string(), path); + } + if let Some(content) = extract_content_from_args(data) { + tool_call_content.insert(call_id.to_string(), content); + } + } + } + + for evt in &turn_events { + if evt.event_type != "tool_execution_end" { + continue; + } + let data = &evt.data; + let tool_name = data.get("toolName").and_then(|v| v.as_str()).unwrap_or(""); + let is_error = data.get("isError").and_then(|v| v.as_bool()).unwrap_or(false); + if is_error { + continue; + } + let call_id = data.get("toolCallId").and_then(|v| v.as_str()).unwrap_or(""); + + match tool_name { + "edit" => { + if let Some(path) = tool_call_paths.get(call_id) { + files_edited.insert(path.clone()); + } + if let Some(diff) = data + .get("result") + .and_then(|r| r.get("details")) + .and_then(|d| d.get("diff")) + .and_then(|v| v.as_str()) + { + for line in diff.lines() { + if line.starts_with('+') && !line.starts_with("++") { + lines_added += 1; + } + if line.starts_with('-') && !line.starts_with("--") { + lines_removed += 1; + } + } + } + } + "write" => { + if let Some(path) = tool_call_paths.get(call_id) { + files_created.insert(path.clone()); + } + if let Some(content) = tool_call_content.get(call_id) { + lines_added += content.lines().count() as u32; + } + } + _ => {} + } + } + + Some(TurnStats { + files_edited: files_edited.len() as u32, + files_created: files_created.len() as u32, + lines_added, + lines_removed, + duration_ms, + }) +} + +fn extract_path_from_args(data: &Value) -> String { + if let Some(args) = data.get("args") { + if let Some(path) = args.get("path").and_then(|v| v.as_str()) { + return path.to_string(); + } + } + let args_str = data.get("args").and_then(|v| v.as_str()).unwrap_or(""); + if let Ok(parsed) = serde_json::from_str::(args_str) { + if let Some(path) = parsed.get("path").and_then(|v| v.as_str()) { + return path.to_string(); + } + } + String::new() +} + +fn extract_content_from_args(data: &Value) -> Option { + if let Some(args) = data.get("args") { + if let Some(content) = args.get("content").and_then(|v| v.as_str()) { + return Some(content.to_string()); + } + } + let args_str = data.get("args").and_then(|v| v.as_str())?; + let parsed: Value = serde_json::from_str(args_str).ok()?; + parsed.get("content").and_then(|v| v.as_str()).map(|s| s.to_string()) +} + const SESSION_ONLY_EVENT_TYPES: &[&str] = &[ "message_start", "message_update", @@ -1241,9 +1377,15 @@ pub fn is_session_only_event(event_type: &str) -> bool { SESSION_ONLY_EVENT_TYPES.contains(&event_type) } +fn is_replay_delta_event(event_type: &str) -> bool { + event_type == "message_update" || event_type == "tool_execution_update" +} + pub fn get_buffered_events_for_session( events: &[StreamEvent], session_id: &str, + from_event_id: Option, + from_delta_event_id: Option, ) -> Vec { let session_events: Vec<&StreamEvent> = events .iter() @@ -1255,15 +1397,23 @@ pub fn get_buffered_events_for_session( .rposition(|e| { e.event_type == "turn_end" || e.event_type == "agent_end" - || e.event_type == "message_end" }); - match last_boundary { - Some(idx) => session_events[idx + 1..] - .iter() - .cloned() - .cloned() - .collect(), - None => session_events.into_iter().cloned().collect(), - } + let relevant = match last_boundary { + Some(idx) => &session_events[idx + 1..], + None => session_events.as_slice(), + }; + + relevant + .iter() + .filter(|event| { + let event = **event; + let after_general = from_event_id.is_none_or(|from| event.id > from); + let after_delta = is_replay_delta_event(&event.event_type) + && from_delta_event_id.is_some_and(|from| event.id > from); + after_general || after_delta + }) + .cloned() + .cloned() + .collect() } diff --git a/backend/src/services/mod.rs b/backend/src/services/mod.rs index 77517b9..f32f151 100644 --- a/backend/src/services/mod.rs +++ b/backend/src/services/mod.rs @@ -8,4 +8,5 @@ pub mod port_scanner; pub mod provider; pub mod runtime; pub mod session; +pub mod sse_registry; pub mod task; diff --git a/backend/src/services/provider/traits.rs b/backend/src/services/provider/traits.rs index 04403b9..a161956 100644 --- a/backend/src/services/provider/traits.rs +++ b/backend/src/services/provider/traits.rs @@ -257,6 +257,16 @@ pub struct ToolContent { pub details: Option, } +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct TurnStats { + pub files_edited: u32, + pub files_created: u32, + pub lines_added: u32, + pub lines_removed: u32, + pub duration_ms: i64, +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct ModelInfo { @@ -534,10 +544,12 @@ pub enum AgentStreamEvent { #[serde(rename = "agent_start")] AgentStart, - #[serde(rename = "agent_end")] + #[serde(rename = "agent_end", rename_all = "camelCase")] AgentEnd { #[serde(default)] messages: Value, + #[serde(default, skip_serializing_if = "Option::is_none")] + turn_stats: Option, }, #[serde(rename = "turn_start")] @@ -549,6 +561,8 @@ pub enum AgentStreamEvent { message: Option, #[serde(default, skip_serializing_if = "Option::is_none")] tool_results: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + turn_stats: Option, }, #[serde(rename = "message_start")] diff --git a/backend/src/services/session.rs b/backend/src/services/session.rs index 4c6b0b9..aa8b555 100644 --- a/backend/src/services/session.rs +++ b/backend/src/services/session.rs @@ -1,5 +1,5 @@ use serde_json::Value; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::path::{Path, PathBuf}; use crate::models::{PaginatedSessions, SessionDetail, SessionEntry, SessionHeader, SessionListItem, SessionTreeNode}; @@ -200,6 +200,15 @@ pub fn get_session_messages_paginated( all_messages.push((entry_id, msg)); } + // Stamp turn stats on ALL messages before paginating so + // turns that span page boundaries still get stats. + let mut all_values: Vec = all_messages.iter().map(|(_, msg)| msg.clone()).collect(); + stamp_turn_stats_on_messages(&mut all_values); + // Write stamped values back. + for (i, val) in all_values.into_iter().enumerate() { + all_messages[i].1 = val; + } + let end_index = if let Some(before_id) = before_entry_id { all_messages.iter().position(|(eid, _)| eid.as_deref() == Some(before_id)) .unwrap_or(all_messages.len()) @@ -228,61 +237,155 @@ pub fn get_session_messages_paginated( }) } -pub fn get_session_messages_after( - base_path: &Path, - session_id: &str, - last_message_id: Option<&str>, -) -> Option> { - let file_path = find_session_file_anywhere(base_path, session_id)?; - let content = std::fs::read_to_string(&file_path).ok()?; +fn is_diff_add_line(line: &str) -> bool { + line.starts_with('+') && !line.starts_with("++") +} - let mut messages = Vec::new(); - let mut found_marker = last_message_id.is_none(); +fn is_diff_rm_line(line: &str) -> bool { + line.starts_with('-') && !line.starts_with("--") +} - for line in content.lines().skip(1) { - if line.trim().is_empty() { - continue; - } +fn stamp_turn_stats_on_messages(messages: &mut [Value]) { - let Ok(val) = serde_json::from_str::(line) else { - continue; - }; + let mut turn_start_idx: Option = None; - if val.get("type").and_then(|v| v.as_str()) != Some("message") { + let mut i = 0; + while i < messages.len() { + let role = messages[i].get("role").and_then(|v| v.as_str()).unwrap_or(""); + + if role == "user" { + turn_start_idx = Some(i); + i += 1; continue; } - let Some(message) = val.get("message") else { + let is_final_assistant = role == "assistant" + && messages[i].get("stopReason").and_then(|v| v.as_str()) == Some("stop"); + + if !is_final_assistant || turn_start_idx.is_none() { + i += 1; continue; - }; + } - let entry_id = val.get("id").and_then(|v| v.as_str()); + let start = turn_start_idx.unwrap(); + let mut files_edited = HashSet::new(); + let mut files_created = HashSet::new(); + let mut lines_added: u32 = 0; + let mut lines_removed: u32 = 0; - if !found_marker { - let msg_id = message.get("id").and_then(|v| v.as_str()) - .or_else(|| message.get("messageId").and_then(|v| v.as_str())) - .or_else(|| message.get("entryId").and_then(|v| v.as_str())) - .or(entry_id); + for j in start..=i { + let msg = &messages[j]; + let Some(content) = msg.get("content").and_then(|v| v.as_array()) else { + continue; + }; + for block in content { + if block.get("type").and_then(|v| v.as_str()) != Some("toolCall") { + continue; + } + let tool_name = block.get("name").and_then(|v| v.as_str()).unwrap_or(""); + let args = block.get("arguments"); + let path_str = extract_path_from_block(block); + if path_str.is_empty() { + continue; + } - if msg_id == last_message_id { - found_marker = true; + match tool_name { + "edit" => { + files_edited.insert(path_str); + // Look for diff in toolResult (matched by toolCallId) + let tool_call_id = block.get("id").and_then(|v| v.as_str()).unwrap_or(""); + for k in start..=i { + let tr = &messages[k]; + if tr.get("role").and_then(|v| v.as_str()) != Some("toolResult") { + continue; + } + if tr.get("toolCallId").and_then(|v| v.as_str()) != Some(tool_call_id) { + continue; + } + if let Some(diff) = tr.get("details").and_then(|d| d.get("diff")).and_then(|v| v.as_str()) { + for line in diff.lines() { + if is_diff_add_line(line) { + lines_added += 1; + } + if is_diff_rm_line(line) { + lines_removed += 1; + } + } + } + break; + } + } + "write" => { + files_created.insert(path_str); + let content_str = if let Some(a) = args { + if a.is_object() { + a.get("content").and_then(|v| v.as_str()).unwrap_or("").to_string() + } else if let Some(s) = a.as_str() { + serde_json::from_str::(s) + .ok() + .and_then(|v| v.get("content")?.as_str().map(|s| s.to_string())) + .unwrap_or_default() + } else { + String::new() + } + } else { + String::new() + }; + if !content_str.is_empty() { + lines_added += content_str.lines().count() as u32; + } + } + _ => {} + } } - continue; } - let mut msg = message.clone(); - if let Some(eid) = entry_id { - if let Some(obj) = msg.as_object_mut() { - if !obj.contains_key("id") && !obj.contains_key("entryId") { - obj.insert("entryId".to_string(), Value::String(eid.to_string())); + let start_ts = messages[start].get("timestamp").and_then(|v| v.as_f64()).or_else(|| { + messages[start].get("timestamp").and_then(|v| v.as_str()).and_then(|s| s.parse::().ok()) + }); + let end_ts = messages[i].get("timestamp").and_then(|v| v.as_f64()).or_else(|| { + messages[i].get("timestamp").and_then(|v| v.as_str()).and_then(|s| s.parse::().ok()) + }); + + if let Some(obj) = messages[i].as_object_mut() { + if !files_edited.is_empty() || !files_created.is_empty() { + obj.insert("turnFileStats".to_string(), serde_json::json!({ + "filesEdited": files_edited.len(), + "filesCreated": files_created.len(), + "linesAdded": lines_added, + "linesRemoved": lines_removed, + })); + } + + if let (Some(s), Some(e)) = (start_ts, end_ts) { + let dur = (e - s) as i64; + if dur > 0 { + obj.insert("turnDurationMs".to_string(), serde_json::json!(dur)); } } } - messages.push(msg); + turn_start_idx = None; + i += 1; } +} - Some(messages) +fn extract_path_from_block(block: &Value) -> String { + let args = block.get("arguments"); + if let Some(a) = args { + if a.is_object() { + if let Some(p) = a.get("path").and_then(|v| v.as_str()) { + return p.to_string(); + } + } else if let Some(s) = a.as_str() { + if let Ok(parsed) = serde_json::from_str::(s) { + if let Some(p) = parsed.get("path").and_then(|v| v.as_str()) { + return p.to_string(); + } + } + } + } + String::new() } pub fn get_session_tree(base_path: &Path, cwd: &str, session_id: &str) -> Option> { diff --git a/backend/src/services/sse_registry.rs b/backend/src/services/sse_registry.rs new file mode 100644 index 0000000..a37ffdc --- /dev/null +++ b/backend/src/services/sse_registry.rs @@ -0,0 +1,87 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use serde_json::Value; +use tokio::sync::{mpsc, RwLock}; +use tokio::time::Instant; +use uuid::Uuid; + +const PREVIOUS_SESSION_GRACE_SECS: u64 = 10; + +struct SessionState { + active: Option, + previous: Option, + previous_deadline: Option, +} + +pub struct SseConnection { + session_state: RwLock, + pub inject_tx: mpsc::UnboundedSender, +} + +impl SseConnection { + pub async fn is_session_receiving_deltas(&self, session_id: &str) -> bool { + let state = self.session_state.read().await; + if state.active.as_deref() == Some(session_id) { + return true; + } + if state.previous.as_deref() == Some(session_id) { + if let Some(deadline) = state.previous_deadline { + return Instant::now() < deadline; + } + } + false + } + + pub async fn set_active(&self, session_id: Option) { + let mut state = self.session_state.write().await; + let old_active = state.active.take(); + + if let Some(ref old) = old_active { + if session_id.as_deref() != Some(old.as_str()) { + state.previous = Some(old.clone()); + state.previous_deadline = + Some(Instant::now() + Duration::from_secs(PREVIOUS_SESSION_GRACE_SECS)); + } + } + + state.active = session_id; + } +} + +#[derive(Clone)] +pub struct SseConnectionRegistry { + connections: Arc>>>, +} + +impl SseConnectionRegistry { + pub fn new() -> Self { + Self { + connections: Arc::new(RwLock::new(HashMap::new())), + } + } + + pub async fn register(&self) -> (String, Arc, mpsc::UnboundedReceiver) { + let id = Uuid::new_v4().to_string(); + let (inject_tx, inject_rx) = mpsc::unbounded_channel(); + let conn = Arc::new(SseConnection { + session_state: RwLock::new(SessionState { + active: None, + previous: None, + previous_deadline: None, + }), + inject_tx, + }); + self.connections.write().await.insert(id.clone(), conn.clone()); + (id, conn, inject_rx) + } + + pub async fn unregister(&self, connection_id: &str) { + self.connections.write().await.remove(connection_id); + } + + pub async fn get(&self, connection_id: &str) -> Option> { + self.connections.read().await.get(connection_id).cloned() + } +} diff --git a/components/ui/animated-list-item.tsx b/components/ui/animated-list-item.tsx new file mode 100644 index 0000000..4f7a34e --- /dev/null +++ b/components/ui/animated-list-item.tsx @@ -0,0 +1,16 @@ +import type { ReactNode } from "react"; +import Animated, { FadeIn, FadeOut, LinearTransition } from "react-native-reanimated"; + +const ITEM_LAYOUT = LinearTransition.springify().damping(18).stiffness(180).mass(0.7); + +export function AnimatedListItem({ children }: { children: ReactNode }) { + return ( + + {children} + + ); +} diff --git a/features/agent/components/message-list/assistant-message.tsx b/features/agent/components/message-list/assistant-message.tsx index 1ef8187..08cbe38 100644 --- a/features/agent/components/message-list/assistant-message.tsx +++ b/features/agent/components/message-list/assistant-message.tsx @@ -50,7 +50,8 @@ export const AssistantMessage = memo(function AssistantMessage({ const isThinkingOnly = hasThinking && !hasText && !hasToolCalls && isStreaming; const isMidTurn = message.stopReason === "toolUse"; const turnCompleted = !sessionStreaming; - const showToolbar = !isStreaming && !isMidTurn && (!!message.text || !!message.errorMessage); + const isFinalResponse = message.stopReason === "stop"; + const showToolbar = !isStreaming && isFinalResponse && (!!message.text || !!message.errorMessage); const [hovered, setHovered] = useState(false); const isWeb = Platform.OS === "web"; @@ -257,7 +258,7 @@ const styles = StyleSheet.create({ container: { paddingHorizontal: 16, paddingVertical: 4, - gap: 8, + gap: 12, }, textBlock: {}, errorBlock: { diff --git a/features/agent/components/message-list/index.tsx b/features/agent/components/message-list/index.tsx index 4935ecf..240b10b 100644 --- a/features/agent/components/message-list/index.tsx +++ b/features/agent/components/message-list/index.tsx @@ -15,7 +15,7 @@ import { ArrowDown } from "lucide-react-native"; import { useAgentSession } from "@pi-ui/client"; import { Colors, Fonts } from "@/constants/theme"; import { useColorScheme } from "@/hooks/use-color-scheme"; -import type { ChatMessage, ToolCallInfo } from "../../types"; +import type { ChatMessage, ToolCallInfo, TurnFileStats } from "../../types"; import { UserMessage } from "./user-message"; import { AssistantMessage } from "./assistant-message"; import { SystemMessage } from "./system-message"; @@ -24,13 +24,6 @@ interface MessageListProps { sessionId: string; } -interface TurnFileStats { - filesEdited: number; - filesCreated: number; - linesAdded: number; - linesRemoved: number; -} - interface VisibleMessageItem { key: string; message: ChatMessage; @@ -41,8 +34,6 @@ interface VisibleMessageItem { function mergeConsecutiveToolCalls( messages: ChatMessage[], - turnDurations: Map, - turnFileStatsMap: Map, ): VisibleMessageItem[] { const visible: VisibleMessageItem[] = []; let anchor: VisibleMessageItem | null = null; @@ -54,12 +45,10 @@ function mergeConsecutiveToolCalls( (!!msg.errorMessage && msg.errorMessage.length > 0) || (!!msg.thinking && msg.thinking.length > 0); const toolCalls = msg.toolCalls?.length ? msg.toolCalls : undefined; - const turnDurationMs = turnDurations.get(msg.id); - const turnFileStats = turnFileStatsMap.get(msg.id); if (msg.role === "user" || msg.role === "system") { anchor = null; - visible.push({ key: msg.id, message: msg, turnDurationMs }); + visible.push({ key: msg.id, message: msg }); continue; } @@ -68,8 +57,8 @@ function mergeConsecutiveToolCalls( key: msg.id, message: msg, toolCalls, - turnDurationMs, - turnFileStats, + turnDurationMs: msg.turnDurationMs, + turnFileStats: msg.turnFileStats, }; anchor = msg.isStreaming ? null : item; visible.push(item); @@ -80,8 +69,8 @@ function mergeConsecutiveToolCalls( anchor.toolCalls = anchor.toolCalls?.length ? [...anchor.toolCalls, ...toolCalls] : [...toolCalls]; - anchor.turnDurationMs = anchor.turnDurationMs ?? turnDurationMs; - anchor.turnFileStats = anchor.turnFileStats ?? turnFileStats; + anchor.turnDurationMs = anchor.turnDurationMs ?? msg.turnDurationMs; + anchor.turnFileStats = anchor.turnFileStats ?? msg.turnFileStats; } } @@ -109,64 +98,9 @@ export const MessageList = memo(function MessageList({ const prevMessageCountRef = useRef(messages.length); - const { turnDurations, turnFileStatsMap } = useMemo(() => { - const durations = new Map(); - const fileStats = new Map(); - let turnStartTs: number | null = null; - let turnStartIdx: number | null = null; - for (let i = 0; i < messages.length; i++) { - const msg = messages[i]!; - if (msg.role === "user") { - turnStartTs = msg.timestamp; - turnStartIdx = i; - } else if (msg.role === "assistant" && msg.stopReason === "stop" && turnStartTs) { - durations.set(msg.id, msg.timestamp - turnStartTs); - if (turnStartIdx !== null) { - const edited = new Set(); - const created = new Set(); - let linesAdded = 0; - let linesRemoved = 0; - for (let j = turnStartIdx; j <= i; j++) { - const tc = messages[j]!.toolCalls; - if (!tc) continue; - for (const call of tc) { - if (call.status !== "complete") continue; - let path = ""; - try { path = (JSON.parse(call.arguments || "{}") as Record).path as string || ""; } catch {} - if (!path) continue; - if (call.name === "edit") { - edited.add(path); - const diff = call.diff?.trim() || ""; - if (diff) { - for (const line of diff.split("\n")) { - if (/^\+(?!\+)/.test(line)) linesAdded++; - if (/^-(?!-)/.test(line)) linesRemoved++; - } - } - } - if (call.name === "write") { - created.add(path); - try { - const content = (JSON.parse(call.arguments || "{}") as Record).content as string || ""; - if (content) linesAdded += content.split("\n").length; - } catch {} - } - } - } - if (edited.size > 0 || created.size > 0) { - fileStats.set(msg.id, { filesEdited: edited.size, filesCreated: created.size, linesAdded, linesRemoved }); - } - } - turnStartTs = null; - turnStartIdx = null; - } - } - return { turnDurations: durations, turnFileStatsMap: fileStats }; - }, [messages]); - const visibleItems = useMemo( - () => mergeConsecutiveToolCalls(messages, turnDurations, turnFileStatsMap), - [messages, turnDurations, turnFileStatsMap], + () => mergeConsecutiveToolCalls(messages), + [messages], ); const reversed = useMemo(() => [...visibleItems].reverse(), [visibleItems]); @@ -204,11 +138,15 @@ export const MessageList = memo(function MessageList({ listRef.current?.scrollToOffset({ offset: 0, animated: true }); }, []); + const sessionRef = useRef(session); + sessionRef.current = session; + const handleLoadMore = useCallback(() => { - if (session.hasMoreMessages && !session.isLoadingOlderMessages) { - session.loadOlderMessages(); + const s = sessionRef.current; + if (s.hasMoreMessages && !s.isLoadingOlderMessages) { + s.loadOlderMessages(); } - }, [session]); + }, []); const renderItem = useCallback( ({ item }: ListRenderItemInfo) => ( @@ -449,9 +387,9 @@ const styles = StyleSheet.create({ height: 16, }, summaryBlock: { - width: 8, - height: 16, - borderRadius: 2, + width: 5, + height: 12, + borderRadius: 1, }, summaryText: { fontSize: 12, diff --git a/features/agent/components/message-list/tool-call/bash-tool-call.tsx b/features/agent/components/message-list/tool-call/bash-tool-call.tsx index 23ccc03..4774366 100644 --- a/features/agent/components/message-list/tool-call/bash-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/bash-tool-call.tsx @@ -4,6 +4,7 @@ import { Colors, Fonts } from "@/constants/theme"; import type { ToolCallInfo } from "../../../types"; import { isToolActive, parseToolArguments, truncateOutput } from "../utils"; import { AnimatedCollapse } from "../animated-collapse"; +import { ToolResultImages } from "./tool-result-images"; interface BashToolCallProps { tc: ToolCallInfo; @@ -82,6 +83,9 @@ export const BashToolCall = memo(function BashToolCall({ + {tc.resultImages && tc.resultImages.length > 0 && ( + + )} ); }); @@ -91,7 +95,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, ranLabel: { fontSize: 12, @@ -105,7 +109,7 @@ const styles = StyleSheet.create({ terminal: { borderRadius: 6, padding: 10, - marginTop: 4, + marginTop: 8, marginLeft: 12, }, promptLine: { diff --git a/features/agent/components/message-list/tool-call/download-tool-call.tsx b/features/agent/components/message-list/tool-call/download-tool-call.tsx index 94b335e..d4d8a18 100644 --- a/features/agent/components/message-list/tool-call/download-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/download-tool-call.tsx @@ -39,7 +39,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, label: { fontSize: 12, diff --git a/features/agent/components/message-list/tool-call/edit-tool-call.tsx b/features/agent/components/message-list/tool-call/edit-tool-call.tsx index 2411085..57fdbc7 100644 --- a/features/agent/components/message-list/tool-call/edit-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/edit-tool-call.tsx @@ -215,7 +215,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, titleRow: { flexDirection: "row", @@ -262,7 +262,7 @@ const styles = StyleSheet.create({ fontFamily: Fonts.sansMedium, }, diffWrap: { - marginTop: 4, + marginTop: 8, marginLeft: 12, borderRadius: 6, overflow: "hidden", diff --git a/features/agent/components/message-list/tool-call/generic-tool-call.tsx b/features/agent/components/message-list/tool-call/generic-tool-call.tsx index 87f922d..ff1b5a0 100644 --- a/features/agent/components/message-list/tool-call/generic-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/generic-tool-call.tsx @@ -1,33 +1,54 @@ -import { memo, useCallback, useState } from "react"; +import { memo, useCallback, useEffect, useRef, useState } from "react"; import { Pressable, StyleSheet, Text, View } from "react-native"; import { Colors, Fonts } from "@/constants/theme"; import type { ToolCallInfo } from "../../../types"; -import { toolDisplayName } from "../utils"; +import { toolDisplayName, isToolActive } from "../utils"; import { AnimatedCollapse } from "../animated-collapse"; +import { ToolResultImages } from "./tool-result-images"; interface GenericToolCallProps { tc: ToolCallInfo; isDark: boolean; + turnCompleted?: boolean; } export const GenericToolCall = memo(function GenericToolCall({ tc, isDark, + turnCompleted = false, }: GenericToolCallProps) { const colors = isDark ? Colors.dark : Colors.light; - const [expanded, setExpanded] = useState(false); + const active = isToolActive(tc); + const [expanded, setExpanded] = useState(() => active || (!turnCompleted && tc.status === "complete")); const toggle = useCallback(() => setExpanded((p) => !p), []); + useEffect(() => { + if (active) setExpanded(true); + }, [active]); + + const prevTurnCompleted = useRef(turnCompleted); + useEffect(() => { + const justCompleted = turnCompleted && !prevTurnCompleted.current; + prevTurnCompleted.current = turnCompleted; + if (justCompleted && !active) { + const timer = setTimeout(() => setExpanded(false), 400); + return () => clearTimeout(timer); + } + }, [turnCompleted, active]); + + const hasImages = !!(tc.resultImages && tc.resultImages.length > 0); const hasResult = !!tc.result || !!tc.partialResult; + const hasContent = hasResult || hasImages; const resultText = tc.result || tc.partialResult || ""; return ( - + {toolDisplayName(tc.name)} + {hasImages && } ; case "read": - return ; + return ; case "write": return ; case "edit": @@ -95,9 +95,9 @@ function SingleToolCall({ tc, isDark, turnCompleted }: { tc: ToolCallInfo; isDar case "download": return ; case "subagent": - return ; + return ; default: - return ; + return ; } } @@ -233,13 +233,12 @@ const GroupedToolCalls = memo(function GroupedToolCalls({ const styles = StyleSheet.create({ container: { - gap: 8, + gap: 12, }, groupHeader: { flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, }, labelRow: { flexDirection: "row", diff --git a/features/agent/components/message-list/tool-call/read-tool-call.tsx b/features/agent/components/message-list/tool-call/read-tool-call.tsx index c044112..b6716cd 100644 --- a/features/agent/components/message-list/tool-call/read-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/read-tool-call.tsx @@ -1,30 +1,49 @@ -import { memo, useCallback, useState } from "react"; +import { memo, useCallback, useEffect, useRef, useState } from "react"; import { Pressable, StyleSheet, Text, View } from "react-native"; import { Colors, Fonts } from "@/constants/theme"; import type { ToolCallInfo } from "../../../types"; -import { basename, parseToolArguments } from "../utils"; +import { basename, isToolActive, parseToolArguments } from "../utils"; import { CodePreview } from "../code-preview"; import { AnimatedCollapse } from "../animated-collapse"; +import { ToolResultImages } from "./tool-result-images"; interface ReadToolCallProps { tc: ToolCallInfo; isDark: boolean; + turnCompleted?: boolean; } export const ReadToolCall = memo(function ReadToolCall({ tc, isDark, + turnCompleted = false, }: ReadToolCallProps) { const colors = isDark ? Colors.dark : Colors.light; - const [expanded, setExpanded] = useState(false); + const active = isToolActive(tc); + const [expanded, setExpanded] = useState(() => active); const toggle = useCallback(() => setExpanded((p) => !p), []); + useEffect(() => { + if (active) setExpanded(true); + }, [active]); + + const prevTurnCompleted = useRef(turnCompleted); + useEffect(() => { + const justCompleted = turnCompleted && !prevTurnCompleted.current; + prevTurnCompleted.current = turnCompleted; + if (justCompleted && !active) { + const timer = setTimeout(() => setExpanded(false), 400); + return () => clearTimeout(timer); + } + }, [turnCompleted, active]); + const parsed = parseToolArguments(tc.arguments); const filePath = (parsed.path as string) || ""; const fileName = basename(filePath); const offset = (parsed.offset as number) || 1; const content = tc.result || ""; - const hasContent = !!content; + const hasImages = !!(tc.resultImages && tc.resultImages.length > 0); + const hasContent = !!content || hasImages; return ( @@ -33,7 +52,8 @@ export const ReadToolCall = memo(function ReadToolCall({ Read {fileName || filePath || "file"} - + {hasImages && } + @@ -47,7 +67,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, fileName: { fontSize: 12, @@ -56,7 +76,7 @@ const styles = StyleSheet.create({ flex: 1, }, previewWrap: { - marginTop: 4, + marginTop: 8, marginLeft: 12, }, }); diff --git a/features/agent/components/message-list/tool-call/subagent-tool-call.tsx b/features/agent/components/message-list/tool-call/subagent-tool-call.tsx index 46d91b7..dc85b3f 100644 --- a/features/agent/components/message-list/tool-call/subagent-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/subagent-tool-call.tsx @@ -10,6 +10,7 @@ import { AnimatedCollapse } from "../animated-collapse"; interface SubagentToolCallProps { tc: ToolCallInfo; isDark: boolean; + turnCompleted?: boolean; } const DETAIL_MAX_HEIGHT = 340; @@ -17,16 +18,27 @@ const DETAIL_MAX_HEIGHT = 340; export const SubagentToolCall = memo(function SubagentToolCall({ tc, isDark, + turnCompleted = false, }: SubagentToolCallProps) { const colors = isDark ? Colors.dark : Colors.light; const active = isToolActive(tc); const hasResult = !!tc.result; - const [expanded, setExpanded] = useState(active || hasResult); + const [expanded, setExpanded] = useState(() => active || (!turnCompleted && hasResult)); const scrollRef = useRef(null); useEffect(() => { - if (active || hasResult) setExpanded(true); - }, [active, hasResult]); + if (active) setExpanded(true); + }, [active]); + + const prevTurnCompleted = useRef(turnCompleted); + useEffect(() => { + const justCompleted = turnCompleted && !prevTurnCompleted.current; + prevTurnCompleted.current = turnCompleted; + if (justCompleted && !active) { + const timer = setTimeout(() => setExpanded(false), 400); + return () => clearTimeout(timer); + } + }, [turnCompleted, active]); const transcript = tc.result || tc.partialResult || ""; const parsed = parseToolArguments(tc.arguments); @@ -144,7 +156,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "flex-start", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, headerText: { flex: 1, @@ -174,7 +186,7 @@ const styles = StyleSheet.create({ detailBox: { borderRadius: 8, padding: 10, - marginTop: 4, + marginTop: 8, marginLeft: 12, }, section: { diff --git a/features/agent/components/message-list/tool-call/tool-result-images.tsx b/features/agent/components/message-list/tool-call/tool-result-images.tsx new file mode 100644 index 0000000..177decd --- /dev/null +++ b/features/agent/components/message-list/tool-call/tool-result-images.tsx @@ -0,0 +1,99 @@ +import { memo, useCallback, useState } from "react"; +import { Image, Pressable, StyleSheet, View, Modal, Platform } from "react-native"; +import type { ToolResultImage } from "../../../types"; + +interface ToolResultImagesProps { + images: ToolResultImage[]; + isDark: boolean; +} + +export const ToolResultImages = memo(function ToolResultImages({ + images, + isDark, +}: ToolResultImagesProps) { + const [previewUri, setPreviewUri] = useState(null); + + const openPreview = useCallback((uri: string) => setPreviewUri(uri), []); + const closePreview = useCallback(() => setPreviewUri(null), []); + + if (!images.length) return null; + + return ( + <> + + {images.map((img, i) => { + const uri = img.data.startsWith("data:") + ? img.data + : `data:${img.mimeType};base64,${img.data}`; + return ( + openPreview(uri)} + style={[ + styles.thumbWrap, + { backgroundColor: isDark ? "#1a1a1a" : "#f0f0f0" }, + ]} + > + + + ); + })} + + {previewUri && ( + + + + + + + + )} + + ); +}); + +const styles = StyleSheet.create({ + container: { + flexDirection: "row", + flexWrap: "wrap", + gap: 8, + marginTop: 6, + marginLeft: 12, + }, + thumbWrap: { + borderRadius: 8, + overflow: "hidden", + maxWidth: 400, + maxHeight: 300, + }, + thumb: { + width: 320, + height: 200, + ...(Platform.OS === "web" ? { maxWidth: "100%" } : {}), + }, + overlay: { + flex: 1, + backgroundColor: "rgba(0,0,0,0.85)", + alignItems: "center", + justifyContent: "center", + }, + previewWrap: { + width: "90%", + height: "80%", + alignItems: "center", + justifyContent: "center", + }, + previewImage: { + width: "100%", + height: "100%", + }, +}); diff --git a/features/agent/components/message-list/tool-call/write-tool-call.tsx b/features/agent/components/message-list/tool-call/write-tool-call.tsx index a27aae7..0bcae89 100644 --- a/features/agent/components/message-list/tool-call/write-tool-call.tsx +++ b/features/agent/components/message-list/tool-call/write-tool-call.tsx @@ -74,7 +74,7 @@ const styles = StyleSheet.create({ flexDirection: "row", alignItems: "center", gap: 6, - paddingVertical: 3, + paddingVertical: 2, }, titleRow: { flexDirection: "row", @@ -103,7 +103,7 @@ const styles = StyleSheet.create({ fontFamily: Fonts.mono, }, previewWrap: { - marginTop: 4, + marginTop: 8, marginLeft: 12, }, diff --git a/features/agent/store/index.ts b/features/agent/store/index.ts index 8f00c63..83bd216 100644 --- a/features/agent/store/index.ts +++ b/features/agent/store/index.ts @@ -564,7 +564,7 @@ function reduceStreamEvents( const lastMsg = { ...msgs[lastIdx] }; msgs[lastIdx] = lastMsg; - if (piEvent.message) { + if (piEvent.message?.role === "assistant") { const msg = piEvent.message; const content = Array.isArray(msg.content) ? msg.content : []; lastMsg.text = content @@ -663,6 +663,9 @@ function reduceStreamEvents( } case "message_end": { + if (piEvent.message?.role !== "assistant") { + break; + } const lastIdx = msgs.findLastIndex( (m) => m.role === "assistant" && m.isStreaming, ); diff --git a/features/agent/types.ts b/features/agent/types.ts index 493cd4a..cab5cee 100644 --- a/features/agent/types.ts +++ b/features/agent/types.ts @@ -41,12 +41,18 @@ export interface SubagentMeta { turns?: number; } +export interface ToolResultImage { + data: string; + mimeType: string; +} + export interface ToolCallInfo { id: string; name: string; arguments: string; status: "streaming" | "pending" | "running" | "complete" | "error"; result?: string; + resultImages?: ToolResultImage[]; isError?: boolean; partialResult?: string; progress?: SubagentProgress; @@ -70,6 +76,13 @@ export interface MessageUsageInfo { currency?: string; } +export interface TurnFileStats { + filesEdited: number; + filesCreated: number; + linesAdded: number; + linesRemoved: number; +} + export interface ChatMessage { id: string; role: "user" | "assistant" | "system"; @@ -85,6 +98,8 @@ export interface ChatMessage { responseId?: string; usage?: MessageUsageInfo; stopReason?: string; + turnDurationMs?: number; + turnFileStats?: TurnFileStats; systemKind?: "bashExecution" | "event"; command?: string; exitCode?: number; diff --git a/features/auth/store/index.ts b/features/auth/store/index.ts index 79d2a14..4f77e69 100644 --- a/features/auth/store/index.ts +++ b/features/auth/store/index.ts @@ -62,7 +62,7 @@ interface AuthState { remote: boolean; load: () => Promise; - loginToServer: (server: Server) => Promise<{ success: boolean; error?: string }>; + loginToServer: (server: Server, credentials: { username: string; password: string }) => Promise<{ success: boolean; error?: string }>; pairWithServer: ( baseUrl: string, qrId: string, @@ -424,12 +424,12 @@ export const useAuthStore = create((set, get) => { } }, - loginToServer: async (server: Server) => { + loginToServer: async (server: Server, credentials: { username: string; password: string }) => { const result = await apiLogin({ baseUrl: server.address, body: { - username: server.username, - password: server.password, + username: credentials.username, + password: credentials.password, }, }); diff --git a/features/chat/components/chat-sidebar.tsx b/features/chat/components/chat-sidebar.tsx index 4e78f5d..fe21b64 100644 --- a/features/chat/components/chat-sidebar.tsx +++ b/features/chat/components/chat-sidebar.tsx @@ -16,6 +16,7 @@ import { useColorScheme } from '@/hooks/use-color-scheme'; import { useChatSessions, usePiClient, useIsSessionActive } from '@pi-ui/client'; import type { SessionListItem } from '@pi-ui/client'; import { SessionActivityIndicator } from '@/features/workspace/components/session-activity-indicator'; +import { AnimatedListItem } from '@/components/ui/animated-list-item'; interface ChatSidebarProps { onNewSession: () => void; @@ -168,15 +169,16 @@ function SessionList({ return ( {sessions.map((session) => ( - + + + ))} ); diff --git a/features/navigation/components/session-sheet-content/index.tsx b/features/navigation/components/session-sheet-content/index.tsx index 01c90ff..15cb516 100644 --- a/features/navigation/components/session-sheet-content/index.tsx +++ b/features/navigation/components/session-sheet-content/index.tsx @@ -11,6 +11,7 @@ import { SquarePen, RefreshCw } from 'lucide-react-native'; import { Fonts } from '@/constants/theme'; import { SessionActivityIndicator } from '@/features/workspace/components/session-activity-indicator'; +import { AnimatedListItem } from '@/components/ui/animated-list-item'; export interface SessionItem { id: string; @@ -117,22 +118,23 @@ export function SessionSheetContent({ {emptyLabel} ) : ( sessions.map((session) => ( - onSelect(session.id)} - style={({ pressed }) => [ - styles.sessionItem, - session.id === selectedSessionId && { - backgroundColor: isDark ? 'rgba(255,255,255,0.08)' : 'rgba(0,0,0,0.06)', - }, - pressed && { opacity: 0.7 }, - ]} - > - - - {session.display_name ?? session.id} - - + + onSelect(session.id)} + style={({ pressed }) => [ + styles.sessionItem, + session.id === selectedSessionId && { + backgroundColor: isDark ? 'rgba(255,255,255,0.08)' : 'rgba(0,0,0,0.06)', + }, + pressed && { opacity: 0.7 }, + ]} + > + + + {session.display_name ?? session.id} + + + )) )} {hasNextPage && ( diff --git a/features/navigation/components/workspace-sheet/index.tsx b/features/navigation/components/workspace-sheet/index.tsx index 9f47d68..c6c9583 100644 --- a/features/navigation/components/workspace-sheet/index.tsx +++ b/features/navigation/components/workspace-sheet/index.tsx @@ -30,6 +30,7 @@ import { usePiClient } from '@pi-ui/client'; import { requestBrowserNotificationPermission } from '@/features/agent/browser-notifications'; import { NewWorkspaceDialog } from '@/features/workspace/components/new-workspace-dialog'; import { SessionActivityIndicator } from '@/features/workspace/components/session-activity-indicator'; +import { AnimatedListItem } from '@/components/ui/animated-list-item'; const SHEET_HEIGHT = 620; const TIMING_CONFIG = { duration: 280, easing: Easing.out(Easing.cubic) }; @@ -438,28 +439,29 @@ function SessionPage({ workspaceId, onSessionPress, onDismiss }: SessionPageProp ) : ( sessions.map((session) => ( - onSessionPress(session.id)} - style={({ pressed }) => [ - styles.sessionItem, - session.id === selectedSessionId && { - backgroundColor: isDark ? 'rgba(255,255,255,0.08)' : 'rgba(0,0,0,0.06)', - }, - pressed && { opacity: 0.7 }, - ]} - > - - + onSessionPress(session.id)} + style={({ pressed }) => [ + styles.sessionItem, + session.id === selectedSessionId && { + backgroundColor: isDark ? 'rgba(255,255,255,0.08)' : 'rgba(0,0,0,0.06)', + }, + pressed && { opacity: 0.7 }, + ]} > - {session.display_name ?? session.id} - - + + + {session.display_name ?? session.id} + + + )) )} {hasNextPage && ( diff --git a/features/servers/components/qr-scanner/index.tsx b/features/servers/components/qr-scanner/index.tsx index 9e1418f..30c7f98 100644 --- a/features/servers/components/qr-scanner/index.tsx +++ b/features/servers/components/qr-scanner/index.tsx @@ -83,8 +83,6 @@ export function QrScanner({ visible, onClose, onNeedNewWorkspace }: QrScannerPro id: serverId, name: existingServer?.name || params.hostname || ip, address, - username: existingServer?.username ?? "", - password: existingServer?.password ?? "", }); reset(); onClose(); diff --git a/features/servers/store/index.ts b/features/servers/store/index.ts index 9e2528d..bd0aa48 100644 --- a/features/servers/store/index.ts +++ b/features/servers/store/index.ts @@ -8,8 +8,6 @@ export interface Server { id: string; name: string; address: string; - username: string; - password: string; } interface ServersState { @@ -60,8 +58,12 @@ export const useServersStore = create((set, get) => ({ loaded: false, load: async () => { - const servers = (await readFromStore()).map((s) => ({ ...s, address: stripTrailingSlashes(s.address) })); + const raw = await readFromStore(); + const servers = raw.map(({ username, password, ...s }: any) => ({ ...s, address: stripTrailingSlashes(s.address) })); set({ servers, loaded: true }); + if (raw.some((s: any) => s.username || s.password)) { + await writeToStore(servers); + } }, addServer: async (server) => { diff --git a/features/workspace/components/session-sidebar/index.tsx b/features/workspace/components/session-sidebar/index.tsx index 248a10e..c72a659 100644 --- a/features/workspace/components/session-sidebar/index.tsx +++ b/features/workspace/components/session-sidebar/index.tsx @@ -18,6 +18,7 @@ import { useWorkspaceSessions as useSessions, usePiClient, useIsSessionActive } import type { SessionListItem } from '@pi-ui/client'; import { requestBrowserNotificationPermission } from "@/features/agent/browser-notifications"; import { SessionActivityIndicator } from "@/features/workspace/components/session-activity-indicator"; +import { AnimatedListItem } from "@/components/ui/animated-list-item"; export function SessionSidebar() { const colorScheme = useColorScheme() ?? "light"; @@ -192,15 +193,16 @@ function SessionList({ return ( {sessions.map((session) => ( - + + + ))} ); diff --git a/packages/pi-client/src/core/api-client.ts b/packages/pi-client/src/core/api-client.ts index a38566d..1ef9813 100644 --- a/packages/pi-client/src/core/api-client.ts +++ b/packages/pi-client/src/core/api-client.ts @@ -1,4 +1,3 @@ -import { client as defaultClient } from "../generated/client.gen"; import * as sdk from "../generated/sdk.gen"; import type { AgentSessionInfo, @@ -32,16 +31,10 @@ import type { FsEntry, FsUploadResponse, PathCompletion, + SessionHistoryResponse, } from "../generated/types.gen"; import type { ImageContent } from "../types/stream-events"; -interface AuthRetryOptions { - _authRetry?: boolean; - _authRetryRequest?: Request; - fetch?: (request: Request) => Promise; - url?: string; -} - function unwrapResult(result: { data?: unknown; error?: unknown }): T { if (result.error !== undefined && result.error !== null) { const errBody = result.error; @@ -69,64 +62,11 @@ function unwrapResult(result: { data?: unknown; error?: unknown }): T { export class ApiClient { private _serverUrl: string; private _accessToken: string; - private _onAuthError?: () => Promise; + private _onAuthError?: () => Promise; constructor(serverUrl: string, accessToken: string) { this._serverUrl = serverUrl; this._accessToken = accessToken; - defaultClient.setConfig({ baseUrl: serverUrl }); - defaultClient.interceptors.request.use((request: Request, opts: AuthRetryOptions) => { - if (this._accessToken) { - request.headers.set("Authorization", `Bearer ${this._accessToken}`); - } else { - request.headers.delete("Authorization"); - request.headers.delete("authorization"); - } - - try { - opts._authRetryRequest = request.clone(); - } catch { - opts._authRetryRequest = undefined; - } - - return request; - }); - - defaultClient.interceptors.response.use( - async (response: Response, request: Request, opts: AuthRetryOptions) => { - if (response.status !== 401 || opts._authRetry || !this._onAuthError) { - return response; - } - - const path = opts.url ?? ""; - if (path.includes("/auth/")) return response; - - const refreshed = await this._onAuthError(); - if (!refreshed) return response; - - const retrySource = opts._authRetryRequest; - const retryHeaders = new Headers(retrySource?.headers ?? request.headers); - retryHeaders.set("Authorization", `Bearer ${this._accessToken}`); - - const retryRequest = retrySource - ? new Request(retrySource, { headers: retryHeaders, signal: request.signal }) - : new Request(request.url, { - method: request.method, - headers: retryHeaders, - body: - request.method !== "GET" && request.method !== "HEAD" - ? request.body - : undefined, - signal: request.signal, - }); - - try { - return await (opts.fetch ?? fetch)(retryRequest); - } catch { - return response; - } - }, - ); } get serverUrl(): string { @@ -140,14 +80,13 @@ export class ApiClient { updateConfig(serverUrl: string, accessToken: string): void { this._serverUrl = serverUrl; this._accessToken = accessToken; - defaultClient.setConfig({ baseUrl: serverUrl }); } updateToken(accessToken: string): void { this._accessToken = accessToken; } - setAuthErrorHandler(handler: () => Promise): void { + setAuthErrorHandler(handler: () => Promise): void { this._onAuthError = handler; } @@ -175,11 +114,12 @@ export class ApiClient { return response; } - const refreshed = await this._onAuthError(); - if (!refreshed) return response; + const newToken = await this._onAuthError(); + if (!newToken) return response; + this._accessToken = newToken; const retryHeaders = new Headers(init?.headers); - retryHeaders.set("Authorization", `Bearer ${this._accessToken}`); + retryHeaders.set("Authorization", `Bearer ${newToken}`); return fetch(input, { ...init, @@ -1037,8 +977,9 @@ export class ApiClient { xhr.onabort = () => reject(new Error("Upload cancelled")); xhr.onload = async () => { if (xhr.status === 401 && allowRetry && this._onAuthError) { - const refreshed = await this._onAuthError(); - if (refreshed) { + const newToken = await this._onAuthError(); + if (newToken) { + this._accessToken = newToken; execute(false).then(resolve).catch(reject); return; } @@ -1260,6 +1201,48 @@ export class ApiClient { return unwrapResult<{ session_id: string; mode: AgentMode | null }>(result); } + // --------------------------------------------------------------------------- + // Session history (REST) + // --------------------------------------------------------------------------- + + async getSessionHistory( + sessionId: string, + params?: { before?: string; limit?: number }, + ): Promise { + const result = await sdk.sessionHistory({ + path: { session_id: sessionId }, + query: { before: params?.before, limit: params?.limit }, + }); + return unwrapResult(result); + } + + // --------------------------------------------------------------------------- + // Active session (per-connection) + // --------------------------------------------------------------------------- + + async setActiveSession( + connectionId: string, + sessionId: string | null, + fromEventId?: number, + fromDeltaEventId?: number, + ): Promise { + const url = this.buildApiUrl("/api/stream-active-session"); + const response = await this.authFetch(url, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + connection_id: connectionId, + session_id: sessionId, + from_event_id: fromEventId, + from_delta_event_id: fromDeltaEventId, + }), + }); + if (!response.ok) { + const body = await response.json().catch(() => null); + throw new Error((body as any)?.error ?? `setActiveSession failed (${response.status})`); + } + } + // --------------------------------------------------------------------------- // Stream URLs // --------------------------------------------------------------------------- diff --git a/packages/pi-client/src/core/index.ts b/packages/pi-client/src/core/index.ts index a6b2222..f29f26f 100644 --- a/packages/pi-client/src/core/index.ts +++ b/packages/pi-client/src/core/index.ts @@ -1,12 +1,6 @@ export { PiClient, type SessionListState } from "./pi-client"; export { ApiClient } from "./api-client"; export { StreamConnection, type StreamConnectionConfig } from "./stream-connection"; -export { - SessionStreamConnection, - type SessionStreamConfig, - type SessionStreamState, - type SessionStreamStatus, -} from "./session-stream-connection"; export { XhrEventSource, type EventSourceEvent } from "./event-source"; export { reduceStreamEvent, diff --git a/packages/pi-client/src/core/message-reducer.ts b/packages/pi-client/src/core/message-reducer.ts index 02f7a6b..98e3860 100644 --- a/packages/pi-client/src/core/message-reducer.ts +++ b/packages/pi-client/src/core/message-reducer.ts @@ -1,5 +1,5 @@ -import type { ChatMessage, ToolCallInfo, MessageUsageInfo, AgentMode, PendingExtensionUiRequest, SubagentMeta } from "../types/chat-message"; -import type { AgentStreamEvent, StreamEventEnvelope } from "../types/stream-events"; +import type { ChatMessage, ToolCallInfo, ToolResultImage, MessageUsageInfo, AgentMode, PendingExtensionUiRequest, SubagentMeta } from "../types/chat-message"; +import type { AgentStateData, StreamEventEnvelope } from "../types/stream-events"; export interface SessionState { messages: ChatMessage[]; @@ -11,6 +11,7 @@ export interface SessionState { oldestEntryId: string | null; mode: AgentMode; pendingExtensionUiRequest: PendingExtensionUiRequest | null; + agentState: AgentStateData | null; } export function createEmptySessionState(): SessionState { @@ -18,6 +19,7 @@ export function createEmptySessionState(): SessionState { messages: [], isStreaming: false, isReady: false, + agentState: null, isLoading: false, isLoadingOlderMessages: false, hasMoreMessages: false, @@ -37,6 +39,21 @@ function extractTextFromContent(content: unknown[] | undefined): string { .join(""); } +function extractImagesFromContent(content: unknown[] | undefined): ToolResultImage[] | undefined { + if (!Array.isArray(content)) return undefined; + const images = content + .filter((c): c is { type: string; data: string; mimeType: string } => + typeof c === "object" && + c !== null && + "type" in c && + (c as { type: string }).type === "image" && + "data" in c && + typeof (c as { data: unknown }).data === "string", + ) + .map((c) => ({ data: c.data, mimeType: c.mimeType ?? "image/png" })); + return images.length > 0 ? images : undefined; +} + function extractMessageEntryId(msg: Record): string | undefined { const rawId = msg["entryId"] ?? msg["entry_id"] ?? msg["id"] ?? msg["messageId"]; if (typeof rawId === "string" && rawId.trim()) return rawId; @@ -55,6 +72,30 @@ const CLEAR_PENDING_EVENTS = new Set([ "turn_end", "agent_end", "session_process_exited", ]); +function stampTurnEndFromBackend( + messages: ChatMessage[], + stats: { filesEdited?: number; filesCreated?: number; linesAdded?: number; linesRemoved?: number; durationMs?: number }, +): ChatMessage[] { + let lastAssistantIdx = -1; + for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i]!.role === "assistant") { lastAssistantIdx = i; break; } + } + if (lastAssistantIdx === -1) return messages; + const msg = messages[lastAssistantIdx]!; + if (msg.turnDurationMs !== undefined) return messages; + + const hasFileStats = (stats.filesEdited ?? 0) > 0 || (stats.filesCreated ?? 0) > 0; + const next = [...messages]; + next[lastAssistantIdx] = { + ...msg, + turnDurationMs: stats.durationMs ?? 0, + turnFileStats: hasFileStats + ? { filesEdited: stats.filesEdited ?? 0, filesCreated: stats.filesCreated ?? 0, linesAdded: stats.linesAdded ?? 0, linesRemoved: stats.linesRemoved ?? 0 } + : undefined, + }; + return next; +} + function findLastStreamingIndex(messages: ChatMessage[]): number { for (let i = messages.length - 1; i >= 0; i--) { if (messages[i]!.role === "assistant" && messages[i]!.isStreaming) return i; @@ -143,7 +184,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve const event = envelope.data; const eventType = envelope.type; - let { messages, isStreaming, mode, pendingExtensionUiRequest } = state; + let { messages, isStreaming, mode, pendingExtensionUiRequest, agentState } = state; if (CLEAR_PENDING_EVENTS.has(eventType)) { pendingExtensionUiRequest = null; @@ -173,8 +214,24 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve break; } + case "turn_end": { + break; + } + case "agent_end": { isStreaming = false; + let lastAssist: ChatMessage | undefined; + for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i]!.role === "assistant") { lastAssist = messages[i]; break; } + } + if (lastAssist?.stopReason === "stop") { + const raw = (event as unknown as Record)["turnStats"] as + | { filesEdited?: number; filesCreated?: number; linesAdded?: number; linesRemoved?: number; durationMs?: number } + | undefined; + if (raw) { + messages = stampTurnEndFromBackend(messages, raw); + } + } messages = updateLastStreaming(messages, (msg) => ({ ...msg, isStreaming: false })); break; } @@ -212,11 +269,6 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve case "message_update": { if (event.type !== "message_update") break; let idx = findLastStreamingIndex(messages); - if (idx === -1) { - for (let j = messages.length - 1; j >= 0; j--) { - if (messages[j]!.role === "assistant") { idx = j; break; } - } - } if (idx === -1) { messages = [...messages, { id: `assistant-${envelope.id}`, @@ -233,7 +285,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve const current = messages[idx]!; let updated = { ...current }; - if (event.message) { + if (event.message?.role === "assistant") { const msg = event.message; const content = Array.isArray(msg.content) ? msg.content : []; updated.text = content @@ -341,7 +393,17 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve case "message_end": { if (event.type !== "message_end") break; const endMsg = event.message as unknown as Record | undefined; - messages = updateLastStreaming(messages, (msg) => { + if (endMsg?.["role"] !== "assistant") { + break; + } + let endIdx = findLastStreamingIndex(messages); + if (endIdx === -1) { + for (let i = messages.length - 1; i >= 0; i--) { + if (messages[i]!.role === "assistant") { endIdx = i; break; } + } + } + if (endIdx !== -1) { + const msg = messages[endIdx]!; const updated: ChatMessage = { ...msg, isStreaming: false, @@ -362,8 +424,10 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve const toolCalls = buildToolCallsFromContent(content, msg.toolCalls, "pending"); if (toolCalls.length > 0) updated.toolCalls = toolCalls; } - return updated; - }); + const next = [...messages]; + next[endIdx] = updated; + messages = next; + } break; } @@ -395,6 +459,9 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve const resultText = event.result ? extractTextFromContent(event.result.content as unknown[]) : undefined; + const resultImages = event.result + ? extractImagesFromContent(event.result.content as unknown[]) + : undefined; const resultDetails = (event.result as any)?.details; const subagentMeta = resultDetails ? extractSubagentMeta(resultDetails) : undefined; const diff = resultDetails && typeof resultDetails.diff === "string" ? resultDetails.diff : undefined; @@ -402,6 +469,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve ...tc, status: event.isError ? "error" : "complete", result: resultText, + resultImages, isError: event.isError, ...(subagentMeta ? { subagentMeta } : {}), ...(diff ? { diff } : {}), @@ -427,12 +495,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve } case "agent_state": { - // agent_state is emitted on session touch/create with full session state - // from the backend. It carries isStreaming, mode, model info, etc. - const data = event as unknown as { - isStreaming?: boolean; - mode?: string; - }; + const data = event as unknown as AgentStateData; if (typeof data.isStreaming === "boolean") { isStreaming = data.isStreaming; if (!isStreaming) { @@ -444,6 +507,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve mode = data.mode; } } + agentState = data; break; } @@ -470,7 +534,7 @@ export function reduceStreamEvent(state: SessionState, envelope: StreamEventEnve } } - return { ...state, messages, isStreaming, mode, pendingExtensionUiRequest }; + return { ...state, messages, isStreaming, mode, pendingExtensionUiRequest, agentState }; } function updateToolCall( @@ -563,6 +627,9 @@ function convertSingleMessage(msg: Record, index: number): Chat const thinking = content.filter(c => c["type"] === "thinking").map(c => c["thinking"] as string ?? "").join(""); const toolCalls = buildToolCallsFromContent(content, undefined, "complete"); + const backendStats = msg["turnFileStats"] as Record | undefined; + const backendDuration = typeof msg["turnDurationMs"] === "number" ? msg["turnDurationMs"] as number : undefined; + return { id: stableId(msg, "assistant", index), entryId: extractMessageEntryId(msg), @@ -578,6 +645,13 @@ function convertSingleMessage(msg: Record, index: number): Chat responseId: msg["responseId"] as string | undefined, usage: extractUsage(msg), stopReason: msg["stopReason"] as ChatMessage["stopReason"], + turnDurationMs: backendDuration, + turnFileStats: backendStats ? { + filesEdited: backendStats["filesEdited"] ?? 0, + filesCreated: backendStats["filesCreated"] ?? 0, + linesAdded: backendStats["linesAdded"] ?? 0, + linesRemoved: backendStats["linesRemoved"] ?? 0, + } : undefined, }; } @@ -619,6 +693,7 @@ export function convertRawMessages(rawMessages: Record[]): ChatM ); if (!tc) continue; tc.result = extractTextFromContent(raw["content"] as unknown[] | undefined); + tc.resultImages = extractImagesFromContent(raw["content"] as unknown[] | undefined); tc.isError = raw["isError"] as boolean; tc.status = raw["isError"] ? "error" : "complete"; const details = raw["details"] as Record | undefined; diff --git a/packages/pi-client/src/core/pi-client.ts b/packages/pi-client/src/core/pi-client.ts index b33d9d8..f76ab71 100644 --- a/packages/pi-client/src/core/pi-client.ts +++ b/packages/pi-client/src/core/pi-client.ts @@ -1,10 +1,9 @@ -import { BehaviorSubject, Subject, Observable, filter, map, distinctUntilChanged, take } from "rxjs"; +import { BehaviorSubject, Subject, Observable, filter, map, distinctUntilChanged } from "rxjs"; import type { ConnectionState, PiClientConfig, SessionListItem } from "../types"; -import type { StreamEventEnvelope, ImageContent } from "../types/stream-events"; +import type { StreamEventEnvelope, ImageContent, AgentStateData } from "../types/stream-events"; import type { ChatMessage, AgentMode, PendingExtensionUiRequest } from "../types/chat-message"; import { ApiClient } from "./api-client"; import { StreamConnection } from "./stream-connection"; -import { SessionStreamConnection } from "./session-stream-connection"; import { reduceStreamEvent, createEmptySessionState, convertRawMessages, type SessionState } from "./message-reducer"; export interface SessionListState { @@ -21,8 +20,6 @@ export class PiClient { private readonly _stream: StreamConnection; private readonly _sessionStates = new Map>(); private readonly _sessionListStates = new Map>(); - private readonly _sessionStreams = new Map(); - private readonly _staleSessionIds = new Set(); private readonly _activeSessionIds$ = new BehaviorSubject>(new Set()); private readonly _config: PiClientConfig; private readonly _serverRestart$ = new Subject(); @@ -30,12 +27,9 @@ export class PiClient { private _instanceId: string | null = null; private _activeSessionIds = new Set(); private _viewedSessionId: string | null = null; - - /** - * Tracks the highest event id processed per session to deduplicate events - * that arrive on both the global stream and the session stream. - */ + private _pendingActiveSession: string | null | undefined = undefined; private readonly _highWaterMarks = new Map(); + private readonly _deltaHighWaterMarks = new Map(); constructor(config: PiClientConfig) { this._config = config; @@ -52,14 +46,27 @@ export class PiClient { }); this._stream.events$.subscribe((envelope) => { - if (__DEV__) console.log("[pi:global]", envelope.type, envelope.session_id, envelope.id, envelope.data); - this._processGlobalEvent(envelope); + if (__DEV__) console.log("[pi:stream]", envelope.type, envelope.session_id, envelope.id); + this._processEvent(envelope); }); this._stream.instanceId$.subscribe((instanceId) => { this._handleInstanceId(instanceId); }); + this._stream.connectionId$.subscribe(() => { + const sessionId = this._pendingActiveSession !== undefined + ? this._pendingActiveSession + : this._viewedSessionId; + this._pendingActiveSession = undefined; + + if (sessionId) { + this._fetchAndApplyHistory(sessionId).then(() => { + this._sendActiveSession(sessionId); + }); + } + }); + this._stream.activeSessions$.subscribe((sessionIds) => { this._activeSessionIds = new Set(sessionIds); this._activeSessionIds$.next(this._activeSessionIds); @@ -79,23 +86,11 @@ export class PiClient { } disconnect(): void { - for (const stream of this._sessionStreams.values()) { - stream.destroy(); - } - this._sessionStreams.clear(); this._stream.disconnect(); } reconnect(): void { this._stream.reconnect(); - if (!this._viewedSessionId) return; - const sessionStream = this._sessionStreams.get(this._viewedSessionId); - if (sessionStream) { - if (__DEV__) console.log("[pi:session]", "reconnect (explicit, full reload)", this._viewedSessionId); - sessionStream.connect(this._viewedSessionId); - return; - } - this._ensureSessionStream(this._viewedSessionId); } get serverRestart$(): Observable { @@ -117,11 +112,6 @@ export class PiClient { updateToken(accessToken: string): void { (this._config as { accessToken: string }).accessToken = accessToken; this.api.updateToken(accessToken); - for (const stream of this._sessionStreams.values()) { - if (stream.stateSnapshot.status === "disconnected") { - stream.reconnect(); - } - } } get events$(): Observable { @@ -140,18 +130,11 @@ export class PiClient { sessionId: string, params: { workspaceId?: string; sessionFile: string }, ): Promise { - // Close previous session stream if switching sessions - if (this._viewedSessionId && this._viewedSessionId !== sessionId) { - const prev = this._viewedSessionId; - if (__DEV__) console.log("[pi:session]", "auto-close previous", prev); - this._closeSessionStream(prev); - } this._viewedSessionId = sessionId; const subject = this._getOrCreateSessionSubject(sessionId); const current = subject.getValue(); - // Touch the session on the backend (register it as active) const touch = async () => { if (params.workspaceId) { await this.api.touchAgentSession(sessionId, { @@ -164,55 +147,29 @@ export class PiClient { }; if (current.isReady) { - // Session was already open — just re-touch and ensure stream is connected. - // Always do a full history reload to avoid stale cache. - try { await touch(); } catch { /* keep going */ } - this._ensureSessionStream(sessionId); + touch().catch(() => {}); + await this._fetchAndApplyHistory(sessionId); + this._setActiveSessionOnBackend(sessionId); return; } - // First open — show loading state subject.next({ ...current, isLoading: true }); try { await touch(); } catch { - subject.next({ ...current, isLoading: false }); + await this._fetchAndApplyHistory(sessionId); + this._setActiveSessionOnBackend(sessionId); return; } - // Connect the session stream. isReady will be set by _onSessionHistoryDone - // once history + buffered events have been fully received. - const sessionStream = this._getOrCreateSessionStream(sessionId); - sessionStream.connect(sessionId); - - // Fetch pending extension UI request in parallel - try { - const state = await this.api.getState(sessionId); - const pending = (state as Record)["pendingExtensionUiRequest"]; - if (pending && typeof pending === "object" && "id" in pending && "method" in pending) { - const latest = subject.getValue(); - subject.next({ - ...latest, - pendingExtensionUiRequest: pending as PendingExtensionUiRequest, - }); - } - } catch { - // state fetch failed — non-critical - } + await this._fetchAndApplyHistory(sessionId); + this._setActiveSessionOnBackend(sessionId); } closeSession(sessionId: string): void { if (__DEV__) console.log("[pi:close]", sessionId); if (this._viewedSessionId === sessionId) { this._viewedSessionId = null; - } - this._closeSessionStream(sessionId); - } - - private _closeSessionStream(sessionId: string): void { - const stream = this._sessionStreams.get(sessionId); - if (stream) { - if (__DEV__) console.log("[pi:close]", "disconnecting stream", sessionId, stream.stateSnapshot.status); - stream.disconnect(); + this._setActiveSessionOnBackend(null); } } @@ -240,6 +197,10 @@ export class PiClient { return this.session$(sessionId).pipe(map((s) => s.pendingExtensionUiRequest), distinctUntilChanged()); } + agentState$(sessionId: string): Observable { + return this.session$(sessionId).pipe(map((s) => s.agentState), distinctUntilChanged()); + } + getSessionSnapshot(sessionId: string): SessionState { return this._getOrCreateSessionSubject(sessionId).getValue(); } @@ -252,7 +213,7 @@ export class PiClient { // Load older messages (pagination) // --------------------------------------------------------------------------- - async loadOlderMessages(sessionId: string, limit = 20): Promise { + async loadOlderMessages(sessionId: string, limit = 50): Promise { const subject = this._getOrCreateSessionSubject(sessionId); const current = subject.getValue(); if (!current.hasMoreMessages || current.isLoadingOlderMessages) return; @@ -260,40 +221,31 @@ export class PiClient { subject.next({ ...current, isLoadingOlderMessages: true }); try { - const temp = new SessionStreamConnection({ - serverUrl: this._config.serverUrl, - getAccessToken: () => this._config.accessToken, + const result = await this.api.getSessionHistory(sessionId, { + before: current.oldestEntryId ?? undefined, + limit, }); - const result = await new Promise((resolve) => { - const collected: StreamEventEnvelope[] = []; - - const timeout = setTimeout(() => { - temp.destroy(); - resolve(collected); - }, 10_000); - - // Collect ALL history events until history_done - temp.historyEvents$.subscribe((envelope) => { - collected.push(envelope); - }); - - temp.historyDone$.pipe(take(1)).subscribe(() => { - clearTimeout(timeout); - temp.destroy(); - resolve(collected); + const rawMessages = result.messages as Record[]; + if (rawMessages.length > 0) { + const converted = convertRawMessages(rawMessages); + const latest = subject.getValue(); + const existingKeys = new Set(latest.messages.map((m) => m.entryId).filter(Boolean)); + const unique = converted.filter((m) => !m.entryId || !existingKeys.has(m.entryId)); + subject.next({ + ...latest, + messages: [...unique, ...latest.messages], + hasMoreMessages: result.has_more, + oldestEntryId: result.oldest_entry_id ?? null, + isLoadingOlderMessages: false, }); - - temp.connect(sessionId, undefined, current.oldestEntryId ?? undefined, limit, false); - }); - - if (result.length > 0) { - // Process all collected history events - for (const envelope of result) { - this._processHistoryEvent(sessionId, envelope, true); - } } else { - subject.next({ ...subject.getValue(), isLoadingOlderMessages: false }); + subject.next({ + ...subject.getValue(), + hasMoreMessages: result.has_more, + oldestEntryId: result.oldest_entry_id ?? null, + isLoadingOlderMessages: false, + }); } } catch { subject.next({ ...subject.getValue(), isLoadingOlderMessages: false }); @@ -459,101 +411,24 @@ export class PiClient { } // --------------------------------------------------------------------------- - // Internal — event processing + // Internal — event processing (single unified stream) // --------------------------------------------------------------------------- private _knownStreamSessionIds = new Set(); - /** - * Event types that the per-session stream (`/api/stream/{id}`) delivers. - * These are filtered by `is_session_only_event` on the backend. - * Everything else (status events) only comes via the global stream. - */ - private static readonly _SESSION_ONLY_EVENTS = new Set([ - "message_start", - "message_update", - "message_end", - "tool_execution_start", - "tool_execution_update", - "tool_execution_end", - "turn_start", - "auto_compaction_start", - "auto_compaction_end", - "auto_retry_start", - "auto_retry_end", - ]); - - /** - * Process an event from the GLOBAL stream (`/api/stream`). - * - * The backend delivers two categories of events: - * - * SESSION_ONLY events (message_*, tool_execution_*, turn_start, etc.) - * → Delivered by BOTH global stream and session stream - * → When session stream is active, skip here (session stream handles them - * including buffered events from the current turn) - * - * STATUS events (agent_state, session_state, agent_start, agent_end, etc.) - * → Delivered ONLY by global stream (backend filters them from session stream) - * → Always process here - * → Do NOT advance highWaterMark (would block session stream buffered events) - */ - private _processGlobalEvent(envelope: StreamEventEnvelope): void { - const sessionId = envelope.session_id; - - if (PiClient._SESSION_ONLY_EVENTS.has(envelope.type)) { - // Session stream handles these (including buffered events from current turn) - if (this._isSessionStreamConnectedOrLoading(sessionId)) { - return; - } - // No session stream — process normally with hwm tracking - this._applyLiveEvent(sessionId, envelope, true); - return; - } - - // Status event — only comes from global stream. - // Apply to reducer but do NOT advance hwm. If we advanced hwm here, - // a status event with id=151 would cause session stream buffered events - // (id=100-150) to be skipped as "already processed". - if (envelope.type === "agent_end" && !this._isSessionStreamConnectedOrLoading(sessionId)) { - this._staleSessionIds.add(sessionId); - } + private _processEvent(envelope: StreamEventEnvelope): void { + if (envelope.type === "history_messages") return; - this._applyLiveEvent(sessionId, envelope, false); - } - - /** - * Process an event from the per-session stream (after history_done). - * These are buffered events + live events. Always advances hwm. - */ - private _processSessionEvent(sessionId: string, envelope: StreamEventEnvelope): void { - if (__DEV__) console.log("[pi:sess-live]", sessionId, envelope.type, envelope.id, envelope.data); - this._applyLiveEvent(sessionId, envelope, true); - } - - /** - * Apply a live (non-history) event to session state. - * - * @param updateHwm - If true, advance the high-water-mark. Session stream - * events always advance it. Global stream status events must NOT advance - * it because their IDs can be higher than buffered events that haven't - * arrived yet from the session stream. - */ - private _applyLiveEvent(sessionId: string, envelope: StreamEventEnvelope, updateHwm: boolean): void { - if (envelope.type === "history_messages") { - this._processHistoryEvent(sessionId, envelope); - return; - } + const sessionId = envelope.session_id; - // Deduplicate by event id if (envelope.id > 0) { const hwm = this._highWaterMarks.get(sessionId) ?? 0; if (envelope.id <= hwm) { - if (__DEV__) console.log("[pi:dedupe]", "skip", sessionId, envelope.type, envelope.id, "<=", hwm); return; } - if (updateHwm) { - this._highWaterMarks.set(sessionId, envelope.id); + this._highWaterMarks.set(sessionId, envelope.id); + if (envelope.type === "message_update" || envelope.type === "tool_execution_update") { + this._deltaHighWaterMarks.set(sessionId, envelope.id); } } @@ -564,7 +439,6 @@ export class PiClient { subject.next(nextState); } - // Side-effects if (envelope.type === "turn_end") { this._fileSystemChanged$.next(); } else if (envelope.type === "tool_execution_end") { @@ -589,150 +463,82 @@ export class PiClient { } // --------------------------------------------------------------------------- - // Internal — history processing + // Internal — active session management // --------------------------------------------------------------------------- - private _processHistoryEvent(sessionId: string, envelope: StreamEventEnvelope, prepend = false): void { - const data = envelope.data as unknown as Record; - if (data["type"] !== "history_messages") return; - - const rawMessages = data["messages"] as Record[]; - const hasMore = data["has_more"] === true; - const oldestEntryId = (data["oldest_entry_id"] as string) ?? null; - - if (__DEV__) console.log("[pi:history]", sessionId, rawMessages?.length ?? 0, "messages", prepend ? "(prepend)" : "(replace)", "hasMore:", hasMore); - - const subject = this._getOrCreateSessionSubject(sessionId); - const current = subject.getValue(); - - if (!rawMessages || rawMessages.length === 0) { - subject.next({ ...current, hasMoreMessages: hasMore, oldestEntryId, isLoadingOlderMessages: false }); + private _setActiveSessionOnBackend(sessionId: string | null): void { + const connectionId = this._stream.connectionId; + if (!connectionId) { + this._pendingActiveSession = sessionId; return; } + this._pendingActiveSession = undefined; + this._sendActiveSession(sessionId); + } - const converted = convertRawMessages(rawMessages); - - if (prepend) { - // Loading older messages — prepend, dedup by entryId - const existingKeys = new Set(current.messages.map((m) => m.entryId).filter(Boolean)); - const unique = converted.filter((m) => !m.entryId || !existingKeys.has(m.entryId)); - subject.next({ - ...current, - messages: [...unique, ...current.messages], - hasMoreMessages: hasMore, - oldestEntryId, - isLoadingOlderMessages: false, - }); - return; + private _sendActiveSession(sessionId: string | null): void { + const connectionId = this._stream.connectionId; + if (!connectionId) return; + const fromEventId = sessionId ? this._highWaterMarks.get(sessionId) : undefined; + const fromDeltaEventId = sessionId ? this._deltaHighWaterMarks.get(sessionId) : undefined; + if (__DEV__) { + console.log( + "[pi:active-session]", + "set", + sessionId, + "conn=", + connectionId, + "from=", + fromEventId, + "fromDelta=", + fromDeltaEventId, + ); } - - // Full history load — REPLACE. History is the authoritative source for - // committed messages. Buffered events (which arrive right after - // history_done) will rebuild any in-progress streaming messages. - // Do NOT merge with current state — that causes duplicates because - // live events and history use different ID formats for the same message. - subject.next({ - ...current, - messages: converted, - hasMoreMessages: hasMore, - oldestEntryId, - isLoadingOlderMessages: false, + this.api.setActiveSession(connectionId, sessionId, fromEventId, fromDeltaEventId).catch((err) => { + if (__DEV__) console.warn("[pi:active-session]", "failed", err); }); } - /** - * Called when the session stream finishes sending history + enters connected state. - * This is the right moment to mark the session as ready. - */ - private _onSessionHistoryDone(sessionId: string): void { - if (__DEV__) console.log("[pi:session]", "history done, marking ready", sessionId); + private async _fetchAndApplyHistory(sessionId: string): Promise { const subject = this._getOrCreateSessionSubject(sessionId); - const current = subject.getValue(); - - // If the session is known to be active (streaming), mark isStreaming. - // The buffered events will arrive right after this and fill in the - // streaming message content. - const isActive = this._activeSessionIds.has(sessionId); - - subject.next({ - ...current, - isReady: true, - isLoading: false, - ...(isActive && !current.isStreaming ? { isStreaming: true } : {}), - }); - } - // --------------------------------------------------------------------------- - // Internal — deduplication helpers - // --------------------------------------------------------------------------- - - // (dedup is handled by simple replace on full history load, - // and entryId-based dedup on prepend) - - // --------------------------------------------------------------------------- - // Internal — session stream management - // --------------------------------------------------------------------------- - - private _isSessionStreamConnectedOrLoading(sessionId: string): boolean { - const stream = this._sessionStreams.get(sessionId); - if (!stream) return false; - const status = stream.stateSnapshot.status; - return status === "connected" || status === "loading_history" || status === "connecting"; - } - - private _ensureSessionStream(sessionId: string): void { - const isStale = this._staleSessionIds.has(sessionId); - this._staleSessionIds.delete(sessionId); - - const stream = this._sessionStreams.get(sessionId); - if (stream && stream.stateSnapshot.status === "connected" && !isStale) { - if (stream.msSinceLastConnected > 60_000) { - if (__DEV__) console.log("[pi:session]", "reconnect (long disconnect, full reload)", sessionId); - this._resetHighWaterMark(sessionId); - stream.connect(sessionId); - return; - } - if (__DEV__) console.log("[pi:session]", "already connected", sessionId); - return; - } - - const sessionStream = this._getOrCreateSessionStream(sessionId); - const wasLongDisconnect = sessionStream.msSinceLastConnected > 60_000; + try { + const result = await this.api.getSessionHistory(sessionId, { limit: 50 }); + const rawMessages = result.messages as Record[]; + const converted = convertRawMessages(rawMessages); + const current = subject.getValue(); - // Always do full history reload on open/stale/long-disconnect. - // Never use SKIP_HISTORY — it causes stale cache issues. - if (__DEV__) { - const reason = isStale ? "stale" : wasLongDisconnect ? "long disconnect" : "full reload"; - console.log("[pi:session]", `connect (${reason})`, sessionId); + subject.next({ + ...current, + messages: converted, + hasMoreMessages: result.has_more, + oldestEntryId: result.oldest_entry_id ?? null, + isReady: true, + isLoading: false, + isLoadingOlderMessages: false, + isStreaming: this._activeSessionIds.has(sessionId) ? true : current.isStreaming, + }); + } catch { + const current = subject.getValue(); + subject.next({ ...current, isReady: true, isLoading: false, isLoadingOlderMessages: false }); } - this._resetHighWaterMark(sessionId); - sessionStream.connect(sessionId); - } - - private _resetHighWaterMark(sessionId: string): void { - this._highWaterMarks.delete(sessionId); } private _handleInstanceId(instanceId: string): void { if (this._instanceId !== null && this._instanceId !== instanceId) { - for (const [id, subject] of this._sessionStates) { - if (!this._activeSessionIds.has(id)) { - subject.next({ ...createEmptySessionState(), isLoading: true }); - } - } - for (const stream of this._sessionStreams.values()) { - stream.disconnect(); + for (const [_id, subject] of this._sessionStates) { + subject.next({ ...createEmptySessionState(), isLoading: true }); } this._knownStreamSessionIds.clear(); - this._staleSessionIds.clear(); this._highWaterMarks.clear(); + this._deltaHighWaterMarks.clear(); this._serverRestart$.next(); } this._instanceId = instanceId; } // --------------------------------------------------------------------------- - // Internal — subject/stream factories + // Internal — subject factories // --------------------------------------------------------------------------- private _getOrCreateSessionSubject(sessionId: string): BehaviorSubject { @@ -760,45 +566,4 @@ export class PiClient { return subject; } - private _getOrCreateSessionStream(sessionId: string): SessionStreamConnection { - let stream = this._sessionStreams.get(sessionId); - if (!stream) { - stream = new SessionStreamConnection({ - serverUrl: this._config.serverUrl, - getAccessToken: () => this._config.accessToken, - getResumeCursor: (sid) => this._getSessionResumeCursor(sid), - onAuthError: this._config.onAuthError, - reconnectBaseMs: this._config.reconnectBaseMs, - reconnectMaxMs: this._config.reconnectMaxMs, - }); - - // History events → process as history (builds committed messages) - stream.historyEvents$.subscribe((envelope) => { - if (__DEV__) console.log("[pi:sess-history]", sessionId, envelope.type); - this._processHistoryEvent(sessionId, envelope); - }); - - // History done → mark session ready, buffered/live events follow - stream.historyDone$.subscribe(() => { - this._onSessionHistoryDone(sessionId); - }); - - // Live events (buffered events from current turn + new live events) - stream.events$.subscribe((envelope) => { - this._processSessionEvent(sessionId, envelope); - }); - - this._sessionStreams.set(sessionId, stream); - } - return stream; - } - - private _getSessionResumeCursor(sessionId: string): string | undefined { - const messages = this._getOrCreateSessionSubject(sessionId).getValue().messages; - for (let i = messages.length - 1; i >= 0; i--) { - const entryId = messages[i]?.entryId; - if (entryId) return entryId; - } - return messages.length > 0 ? "SKIP_HISTORY" : undefined; - } } diff --git a/packages/pi-client/src/core/session-stream-connection.ts b/packages/pi-client/src/core/session-stream-connection.ts deleted file mode 100644 index 1065159..0000000 --- a/packages/pi-client/src/core/session-stream-connection.ts +++ /dev/null @@ -1,297 +0,0 @@ -import { Subject, BehaviorSubject, Observable } from "rxjs"; -import type { StreamEventEnvelope } from "../types/stream-events"; -import { XhrEventSource } from "./event-source"; - -const RECONNECT_BASE_MS = 1000; -const RECONNECT_MAX_MS = 30_000; - -export type SessionStreamStatus = "idle" | "connecting" | "loading_history" | "connected" | "disconnected"; - -export interface SessionStreamState { - status: SessionStreamStatus; - sessionId: string | null; -} - -export interface SessionStreamConfig { - serverUrl: string; - getAccessToken: () => string; - getResumeCursor?: (sessionId: string) => string | undefined; - onAuthError?: () => void; - reconnectBaseMs?: number; - reconnectMaxMs?: number; -} - -function isStreamEventPayload(value: object): boolean { - const v = value as Record; - return ( - typeof v["id"] === "number" && - typeof v["session_id"] === "string" && - typeof v["type"] === "string" && - typeof v["timestamp"] === "number" && - typeof v["data"] === "object" && - v["data"] !== null - ); -} - -export class SessionStreamConnection { - private readonly _events$ = new Subject(); - private readonly _historyEvents$ = new Subject(); - private readonly _historyDone$ = new Subject(); - private readonly _state$ = new BehaviorSubject({ - status: "idle", - sessionId: null, - }); - - private readonly _config: SessionStreamConfig; - private _es: XhrEventSource | null = null; - private _sessionId: string | null = null; - private _lastMessageId: string | undefined; - private _before: string | undefined; - private _limit: number | undefined; - private _retryCount = 0; - private _autoReconnect = true; - private _reconnectTimer: ReturnType | null = null; - private _destroyed = false; - private _lastConnectedAt = 0; - private _wasEverConnected = false; - - constructor(config: SessionStreamConfig) { - this._config = config; - } - - get events$(): Observable { - return this._events$.asObservable(); - } - - get historyEvents$(): Observable { - return this._historyEvents$.asObservable(); - } - - get historyDone$(): Observable { - return this._historyDone$.asObservable(); - } - - get state$(): Observable { - return this._state$.asObservable(); - } - - get stateSnapshot(): SessionStreamState { - return this._state$.getValue(); - } - - get currentSessionId(): string | null { - return this._sessionId; - } - - /** How many ms since we were last connected. Returns Infinity if never connected. */ - get msSinceLastConnected(): number { - if (!this._wasEverConnected) return Infinity; - return Date.now() - this._lastConnectedAt; - } - - connect( - sessionId: string, - lastMessageId?: string, - before?: string, - limit?: number, - autoReconnect = true, - ): void { - if (this._destroyed) return; - - this._clearReconnectTimer(); - this._close(); - this._sessionId = sessionId; - this._lastMessageId = lastMessageId; - this._before = before; - this._limit = limit; - this._retryCount = 0; - this._autoReconnect = autoReconnect; - this._setState({ status: "connecting", sessionId }); - this._openSse(sessionId); - } - - disconnect(): void { - if (__DEV__) console.log("[pi:sess-stream]", "disconnect", this._sessionId); - this._clearReconnectTimer(); - this._close(); - this._sessionId = null; - this._setState({ status: "idle", sessionId: null }); - } - - reconnect(): void { - if (this._destroyed || !this._sessionId) return; - this._clearReconnectTimer(); - this._close(); - this._retryCount = 0; - this._setState({ status: "connecting", sessionId: this._sessionId }); - this._openSse(this._sessionId); - } - - destroy(): void { - this._destroyed = true; - this._clearReconnectTimer(); - this._close(); - this._events$.complete(); - this._historyEvents$.complete(); - this._historyDone$.complete(); - this._state$.complete(); - } - - private _openSse(sessionId: string): void { - if (this._destroyed) return; - - const url = this._buildUrl(sessionId); - const token = this._config.getAccessToken(); - - const es = new XhrEventSource(url, { - headers: { Authorization: `Bearer ${token}` }, - }); - this._es = es; - - let receivingHistory = true; - - es.addEventListener("open", () => { - if (this._destroyed || this._sessionId !== sessionId) { - es.close(); - return; - } - this._retryCount = 0; - this._lastConnectedAt = Date.now(); - this._wasEverConnected = true; - this._setState({ status: "loading_history", sessionId }); - }); - - es.addEventListener("message", (event) => { - if (this._destroyed || !event.data || this._sessionId !== sessionId) return; - - try { - const raw = JSON.parse(event.data) as Record; - if (raw["type"] === "session_stream_hello") return; - if (raw["type"] === "history_done") { - receivingHistory = false; - this._setState({ status: "connected", sessionId }); - this._historyDone$.next(); - return; - } - } catch { - // not a control event - } - - try { - const parsed = JSON.parse(event.data) as object; - if (typeof parsed === "object" && parsed !== null && isStreamEventPayload(parsed)) { - const envelope = parsed as StreamEventEnvelope; - if (receivingHistory) { - this._historyEvents$.next(envelope); - } else { - this._events$.next(envelope); - } - } - } catch { - // parse error - } - }); - - es.addEventListener("history", (event) => { - if (this._destroyed || !event.data || this._sessionId !== sessionId) return; - try { - const parsed = JSON.parse(event.data) as object; - if (typeof parsed === "object" && parsed !== null && isStreamEventPayload(parsed)) { - this._historyEvents$.next(parsed as StreamEventEnvelope); - } - } catch { - // parse error - } - }); - - es.addEventListener("history_done", () => { - if (this._destroyed || this._sessionId !== sessionId) return; - receivingHistory = false; - this._setState({ status: "connected", sessionId }); - this._historyDone$.next(); - }); - - es.addEventListener("error", (event) => { - if (this._destroyed || this._sessionId !== sessionId) return; - const status = event.xhrStatus ?? 0; - this._close(); - if (status === 401 || status === 403) { - this._setState({ status: "disconnected", sessionId }); - this._config.onAuthError?.(); - return; - } - this._scheduleReconnect(sessionId); - }); - - es.addEventListener("close", () => { - if (this._destroyed || this._sessionId !== sessionId) return; - this._close(); - this._scheduleReconnect(sessionId); - }); - } - - private _close(): void { - if (this._es) { - this._es.removeAllEventListeners(); - this._es.close(); - this._es = null; - } - } - - private _setState(state: SessionStreamState): void { - this._state$.next(state); - } - - private _scheduleReconnect(sessionId: string): void { - if (this._destroyed || !this._autoReconnect) { - this._setState({ status: "disconnected", sessionId }); - return; - } - if (this._reconnectTimer) return; - - this._retryCount += 1; - const baseMs = this._config.reconnectBaseMs ?? RECONNECT_BASE_MS; - const maxMs = this._config.reconnectMaxMs ?? RECONNECT_MAX_MS; - const delay = Math.min(baseMs * Math.pow(2, Math.max(0, this._retryCount - 1)), maxMs); - this._setState({ status: "disconnected", sessionId }); - this._reconnectTimer = setTimeout(() => { - this._reconnectTimer = null; - if (this._destroyed || this._sessionId !== sessionId) return; - this._setState({ status: "connecting", sessionId }); - this._openSse(sessionId); - }, delay); - } - - private _clearReconnectTimer(): void { - if (this._reconnectTimer) { - clearTimeout(this._reconnectTimer); - this._reconnectTimer = null; - } - } - - private _resolveLastMessageId(sessionId: string): string | undefined { - // After a long disconnect (>60s, e.g. sleep), do a full reload - if (this._wasEverConnected && this.msSinceLastConnected > 60_000) { - return undefined; // no cursor = full history - } - if (this._retryCount > 0 && !this._before) { - return this._config.getResumeCursor?.(sessionId) ?? this._lastMessageId; - } - return this._lastMessageId; - } - - private _buildUrl(sessionId: string): string { - const url = new URL(`${this._config.serverUrl}/api/stream/${encodeURIComponent(sessionId)}`); - const lastMessageId = this._resolveLastMessageId(sessionId); - if (lastMessageId) { - url.searchParams.set("last_message_id", lastMessageId); - } - if (this._before) { - url.searchParams.set("before", this._before); - } - if (this._limit) { - url.searchParams.set("limit", String(this._limit)); - } - return url.toString(); - } -} diff --git a/packages/pi-client/src/core/stream-connection.ts b/packages/pi-client/src/core/stream-connection.ts index d8cfd13..f001a4a 100644 --- a/packages/pi-client/src/core/stream-connection.ts +++ b/packages/pi-client/src/core/stream-connection.ts @@ -53,9 +53,11 @@ export class StreamConnection { disconnectedAt: null, }); private readonly _instanceId$ = new Subject(); + private readonly _connectionId$ = new Subject(); private readonly _activeSessions$ = new Subject(); private readonly _config: StreamConnectionConfig; + private _connectionId: string | null = null; private _lastEventId: number | null = null; private _retryCount = 0; private _es: XhrEventSource | null = null; @@ -78,6 +80,14 @@ export class StreamConnection { return this._instanceId$.asObservable(); } + get connectionId$(): Observable { + return this._connectionId$.asObservable(); + } + + get connectionId(): string | null { + return this._connectionId; + } + get activeSessions$(): Observable { return this._activeSessions$.asObservable(); } @@ -106,6 +116,7 @@ export class StreamConnection { this._events$.complete(); this._connection$.complete(); this._instanceId$.complete(); + this._connectionId$.complete(); this._activeSessions$.complete(); } @@ -146,6 +157,10 @@ export class StreamConnection { try { const raw = JSON.parse(event.data) as Record; if (raw["type"] === "server_hello" && typeof raw["instance_id"] === "string") { + if (typeof raw["connection_id"] === "string") { + this._connectionId = raw["connection_id"] as string; + this._connectionId$.next(this._connectionId); + } this._instanceId$.next(raw["instance_id"] as string); return; } diff --git a/packages/pi-client/src/generated/sdk.gen.ts b/packages/pi-client/src/generated/sdk.gen.ts index 8a6a480..71ba898 100644 --- a/packages/pi-client/src/generated/sdk.gen.ts +++ b/packages/pi-client/src/generated/sdk.gen.ts @@ -1,7 +1,7 @@ // This file is auto-generated by @hey-api/openapi-ts import { type Options as ClientOptions, type TDataShape, type Client, formDataBodySerializer } from './client'; -import type { AbortData, AbortResponses, AbortErrors, AbortBashData, AbortBashResponses, AbortBashErrors, AbortRetryData, AbortRetryResponses, AbortRetryErrors, BashData, BashResponses, BashErrors, GetCommandsData, GetCommandsResponses, GetCommandsErrors, CompactData, CompactResponses, CompactErrors, CycleModelData, CycleModelResponses, CycleModelErrors, CycleThinkingLevelData, CycleThinkingLevelResponses, CycleThinkingLevelErrors, ExportHtmlData, ExportHtmlResponses, ExportHtmlErrors, ExtensionUiResponseData, ExtensionUiResponseResponses, ExtensionUiResponseErrors, FollowUpData, FollowUpResponses, FollowUpErrors, ForkData, ForkResponses, ForkErrors, GetForkMessagesData, GetForkMessagesResponses, GetForkMessagesErrors, GetLastAssistantTextData, GetLastAssistantTextResponses, GetLastAssistantTextErrors, GetMessagesData, GetMessagesResponses, GetMessagesErrors, GetAvailableModelsData, GetAvailableModelsResponses, GetAvailableModelsErrors, NewSessionData, NewSessionResponses, NewSessionErrors, PromptData, PromptResponses, PromptErrors, RuntimeStatusData, RuntimeStatusResponses, RuntimeStatusErrors, GetSessionStatsData, GetSessionStatsResponses, GetSessionStatsErrors, ListSessionsData, ListSessionsResponses, ListSessionsErrors, CreateSessionData, CreateSessionResponses, CreateSessionErrors, KillSessionData, KillSessionResponses, KillSessionErrors, TouchSessionData, TouchSessionResponses, TouchSessionErrors, SetAutoCompactionData, SetAutoCompactionResponses, SetAutoCompactionErrors, SetAutoRetryData, SetAutoRetryResponses, SetAutoRetryErrors, SetFollowUpModeData, SetFollowUpModeResponses, SetFollowUpModeErrors, SetModelData, SetModelResponses, SetModelErrors, SetSessionNameData, SetSessionNameResponses, SetSessionNameErrors, SetSteeringModeData, SetSteeringModeResponses, SetSteeringModeErrors, SetThinkingLevelData, SetThinkingLevelResponses, SetThinkingLevelErrors, GetStateData, GetStateResponses, GetStateErrors, SteerData, SteerResponses, SteerErrors, SwitchSessionData, SwitchSessionResponses, SwitchSessionErrors, LoginData, LoginResponses, LoginErrors, LogoutData, LogoutResponses, LogoutErrors, PairData, PairResponses, PairErrors, RefreshData, RefreshResponses, RefreshErrors, CheckSessionData, CheckSessionResponses, CheckSessionErrors, ListSessions2Data, ListSessions2Responses, ListSessions2Errors, CreateSession2Data, CreateSession2Responses, CreateSession2Errors, DeleteSessionData, DeleteSessionResponses, DeleteSessionErrors, TouchSession2Data, TouchSession2Responses, TouchSession2Errors, GetCustomModelsData, GetCustomModelsResponses, SaveCustomModelsData, SaveCustomModelsResponses, CompleteData, CompleteResponses, CompleteErrors, DeleteData, DeleteResponses, DeleteErrors, DownloadData, DownloadResponses, DownloadErrors, ListData, ListResponses, ListErrors, MkdirData, MkdirResponses, MkdirErrors, ReadData, ReadResponses, ReadErrors, UploadData, UploadResponses, UploadErrors, WriteData, WriteResponses, WriteErrors, BranchesData, BranchesResponses, BranchesErrors, CheckoutData, CheckoutResponses, CheckoutErrors, CommitData, CommitResponses, CommitErrors, DiffData, DiffResponses, DiffErrors, DiffFileData, DiffFileResponses, DiffFileErrors, DiscardData, DiscardResponses, DiscardErrors, LogData, LogResponses, LogErrors, NestedReposData, NestedReposResponses, NestedReposErrors, StageData, StageResponses, StageErrors, StashDropData, StashDropResponses, StashDropErrors, StashListData, StashListResponses, StashListErrors, StashPushData, StashPushResponses, StashPushErrors, StashApplyData, StashApplyResponses, StashApplyErrors, StatusData, StatusResponses, StatusErrors, UnstageData, UnstageResponses, UnstageErrors, WorktreeRemoveData, WorktreeRemoveResponses, WorktreeRemoveErrors, WorktreeListData, WorktreeListResponses, WorktreeListErrors, WorktreeAddData, WorktreeAddResponses, WorktreeAddErrors, ListModesData, ListModesResponses, ListModesErrors, CreateModeData, CreateModeResponses, CreateModeErrors, DeleteModeData, DeleteModeResponses, DeleteModeErrors, UpdateModeData, UpdateModeResponses, UpdateModeErrors, InstallData, InstallResponses, InstallErrors, LogsData, LogsResponses, LogsErrors, Status2Data, Status2Responses, Status2Errors, UpdateData, UpdateResponses, UpdateErrors, GetSessionModeData, GetSessionModeResponses, GetSessionModeErrors, StreamData, StreamResponses, StreamErrors, GetConfigData, GetConfigResponses, GetConfigErrors, ListTasksData, ListTasksResponses, ListTasksErrors, GetLogsData, GetLogsResponses, GetLogsErrors, RemoveTaskData, RemoveTaskResponses, RemoveTaskErrors, RestartTaskData, RestartTaskResponses, RestartTaskErrors, StartTaskData, StartTaskResponses, StartTaskErrors, StopTaskData, StopTaskResponses, StopTaskErrors, List2Data, List2Responses, List2Errors, CreateData, CreateResponses, CreateErrors, SuggestWorkspacesData, SuggestWorkspacesResponses, SuggestWorkspacesErrors, Delete2Data, Delete2Responses, Delete2Errors, GetData, GetResponses, GetErrors, Update2Data, Update2Responses, Update2Errors, ArchiveData, ArchiveResponses, ArchiveErrors, SessionsListData, SessionsListResponses, SessionsListErrors, SessionsDeleteData, SessionsDeleteResponses, SessionsDeleteErrors, SessionsGetData, SessionsGetResponses, SessionsGetErrors, SessionsBranchData, SessionsBranchResponses, SessionsBranchErrors, SessionsChildrenData, SessionsChildrenResponses, SessionsChildrenErrors, SessionsLeafData, SessionsLeafResponses, SessionsLeafErrors, SessionsTreeData, SessionsTreeResponses, SessionsTreeErrors, UnarchiveData, UnarchiveResponses, UnarchiveErrors, HealthzData, HealthzResponses, VersionData, VersionResponses } from './types.gen'; +import type { AbortData, AbortResponses, AbortErrors, AbortBashData, AbortBashResponses, AbortBashErrors, AbortRetryData, AbortRetryResponses, AbortRetryErrors, BashData, BashResponses, BashErrors, GetCommandsData, GetCommandsResponses, GetCommandsErrors, CompactData, CompactResponses, CompactErrors, CycleModelData, CycleModelResponses, CycleModelErrors, CycleThinkingLevelData, CycleThinkingLevelResponses, CycleThinkingLevelErrors, ExportHtmlData, ExportHtmlResponses, ExportHtmlErrors, ExtensionUiResponseData, ExtensionUiResponseResponses, ExtensionUiResponseErrors, FollowUpData, FollowUpResponses, FollowUpErrors, ForkData, ForkResponses, ForkErrors, GetForkMessagesData, GetForkMessagesResponses, GetForkMessagesErrors, GetLastAssistantTextData, GetLastAssistantTextResponses, GetLastAssistantTextErrors, GetMessagesData, GetMessagesResponses, GetMessagesErrors, GetAvailableModelsData, GetAvailableModelsResponses, GetAvailableModelsErrors, NewSessionData, NewSessionResponses, NewSessionErrors, PromptData, PromptResponses, PromptErrors, RuntimeStatusData, RuntimeStatusResponses, RuntimeStatusErrors, GetSessionStatsData, GetSessionStatsResponses, GetSessionStatsErrors, ListSessionsData, ListSessionsResponses, ListSessionsErrors, CreateSessionData, CreateSessionResponses, CreateSessionErrors, KillSessionData, KillSessionResponses, KillSessionErrors, TouchSessionData, TouchSessionResponses, TouchSessionErrors, SetAutoCompactionData, SetAutoCompactionResponses, SetAutoCompactionErrors, SetAutoRetryData, SetAutoRetryResponses, SetAutoRetryErrors, SetFollowUpModeData, SetFollowUpModeResponses, SetFollowUpModeErrors, SetModelData, SetModelResponses, SetModelErrors, SetSessionNameData, SetSessionNameResponses, SetSessionNameErrors, SetSteeringModeData, SetSteeringModeResponses, SetSteeringModeErrors, SetThinkingLevelData, SetThinkingLevelResponses, SetThinkingLevelErrors, GetStateData, GetStateResponses, GetStateErrors, SteerData, SteerResponses, SteerErrors, SwitchSessionData, SwitchSessionResponses, SwitchSessionErrors, LoginData, LoginResponses, LoginErrors, LogoutData, LogoutResponses, LogoutErrors, PairData, PairResponses, PairErrors, RefreshData, RefreshResponses, RefreshErrors, CheckSessionData, CheckSessionResponses, CheckSessionErrors, ListSessions2Data, ListSessions2Responses, ListSessions2Errors, CreateSession2Data, CreateSession2Responses, CreateSession2Errors, DeleteSessionData, DeleteSessionResponses, DeleteSessionErrors, TouchSession2Data, TouchSession2Responses, TouchSession2Errors, GetCustomModelsData, GetCustomModelsResponses, SaveCustomModelsData, SaveCustomModelsResponses, CompleteData, CompleteResponses, CompleteErrors, DeleteData, DeleteResponses, DeleteErrors, DownloadData, DownloadResponses, DownloadErrors, ListData, ListResponses, ListErrors, MkdirData, MkdirResponses, MkdirErrors, ReadData, ReadResponses, ReadErrors, UploadData, UploadResponses, UploadErrors, WriteData, WriteResponses, WriteErrors, BranchesData, BranchesResponses, BranchesErrors, CheckoutData, CheckoutResponses, CheckoutErrors, CommitData, CommitResponses, CommitErrors, DiffData, DiffResponses, DiffErrors, DiffFileData, DiffFileResponses, DiffFileErrors, DiscardData, DiscardResponses, DiscardErrors, LogData, LogResponses, LogErrors, NestedReposData, NestedReposResponses, NestedReposErrors, StageData, StageResponses, StageErrors, StashDropData, StashDropResponses, StashDropErrors, StashListData, StashListResponses, StashListErrors, StashPushData, StashPushResponses, StashPushErrors, StashApplyData, StashApplyResponses, StashApplyErrors, StatusData, StatusResponses, StatusErrors, UnstageData, UnstageResponses, UnstageErrors, WorktreeRemoveData, WorktreeRemoveResponses, WorktreeRemoveErrors, WorktreeListData, WorktreeListResponses, WorktreeListErrors, WorktreeAddData, WorktreeAddResponses, WorktreeAddErrors, ListModesData, ListModesResponses, ListModesErrors, CreateModeData, CreateModeResponses, CreateModeErrors, DeleteModeData, DeleteModeResponses, DeleteModeErrors, UpdateModeData, UpdateModeResponses, UpdateModeErrors, InstallData, InstallResponses, InstallErrors, LogsData, LogsResponses, LogsErrors, Status2Data, Status2Responses, Status2Errors, UpdateData, UpdateResponses, UpdateErrors, SessionHistoryData, SessionHistoryResponses, SessionHistoryErrors, GetSessionModeData, GetSessionModeResponses, GetSessionModeErrors, StreamData, StreamResponses, StreamErrors, SetActiveSessionData, SetActiveSessionResponses, SetActiveSessionErrors, GetConfigData, GetConfigResponses, GetConfigErrors, ListTasksData, ListTasksResponses, ListTasksErrors, GetLogsData, GetLogsResponses, GetLogsErrors, RemoveTaskData, RemoveTaskResponses, RemoveTaskErrors, RestartTaskData, RestartTaskResponses, RestartTaskErrors, StartTaskData, StartTaskResponses, StartTaskErrors, StopTaskData, StopTaskResponses, StopTaskErrors, List2Data, List2Responses, List2Errors, CreateData, CreateResponses, CreateErrors, SuggestWorkspacesData, SuggestWorkspacesResponses, SuggestWorkspacesErrors, Delete2Data, Delete2Responses, Delete2Errors, GetData, GetResponses, GetErrors, Update2Data, Update2Responses, Update2Errors, ArchiveData, ArchiveResponses, ArchiveErrors, SessionsListData, SessionsListResponses, SessionsListErrors, SessionsDeleteData, SessionsDeleteResponses, SessionsDeleteErrors, SessionsGetData, SessionsGetResponses, SessionsGetErrors, SessionsBranchData, SessionsBranchResponses, SessionsBranchErrors, SessionsChildrenData, SessionsChildrenResponses, SessionsChildrenErrors, SessionsLeafData, SessionsLeafResponses, SessionsLeafErrors, SessionsTreeData, SessionsTreeResponses, SessionsTreeErrors, UnarchiveData, UnarchiveResponses, UnarchiveErrors, HealthzData, HealthzResponses, VersionData, VersionResponses } from './types.gen'; import { client as _heyApiClient } from './client.gen'; export type Options = ClientOptions & { @@ -1236,6 +1236,19 @@ export const update = (options?: Options(options: Options) => { + return (options.client ?? _heyApiClient).get({ + security: [ + { + scheme: 'bearer', + type: 'http' + } + ], + url: '/api/sessions/{session_id}/history', + ...options + }); +}; + export const getSessionMode = (options: Options) => { return (options.client ?? _heyApiClient).get({ security: [ @@ -1262,6 +1275,23 @@ export const stream = (options?: Options(options: Options) => { + return (options.client ?? _heyApiClient).post({ + security: [ + { + scheme: 'bearer', + type: 'http' + } + ], + url: '/api/stream-active-session', + ...options, + headers: { + 'Content-Type': 'application/json', + ...options.headers + } + }); +}; + export const getConfig = (options: Options) => { return (options.client ?? _heyApiClient).get({ security: [ diff --git a/packages/pi-client/src/generated/types.gen.ts b/packages/pi-client/src/generated/types.gen.ts index 66caf34..c6d360b 100644 --- a/packages/pi-client/src/generated/types.gen.ts +++ b/packages/pi-client/src/generated/types.gen.ts @@ -449,6 +449,17 @@ export type SessionHeader = { version: number; }; +export type SessionHistoryQuery = { + before?: string | null; + limit?: number | null; +}; + +export type SessionHistoryResponse = { + has_more: boolean; + messages: Array; + oldest_entry_id?: string | null; +}; + export type SessionInfo = { access_expires_at: string; access_token: string; @@ -482,6 +493,13 @@ export type SessionTreeNode = { timestamp: string; }; +export type SetActiveSessionRequest = { + connection_id: string; + from_delta_event_id?: number | null; + from_event_id?: number | null; + session_id?: string | null; +}; + /** * Request to start a task */ @@ -2634,6 +2652,47 @@ export type UpdateResponses = { export type UpdateResponse = UpdateResponses[keyof UpdateResponses]; +export type SessionHistoryData = { + body?: never; + path: { + /** + * Session ID + */ + session_id: string; + }; + query?: { + /** + * Load messages before this entry ID + */ + before?: string; + /** + * Max messages to return (default 20) + */ + limit?: number; + }; + url: '/api/sessions/{session_id}/history'; +}; + +export type SessionHistoryErrors = { + /** + * Unauthorized + */ + 401: unknown; + /** + * Session not found + */ + 404: unknown; +}; + +export type SessionHistoryResponses = { + /** + * Paginated session messages + */ + 200: SessionHistoryResponse; +}; + +export type SessionHistoryResponse2 = SessionHistoryResponses[keyof SessionHistoryResponses]; + export type GetSessionModeData = { body?: never; path: { @@ -2688,6 +2747,31 @@ export type StreamResponses = { 200: unknown; }; +export type SetActiveSessionData = { + body: SetActiveSessionRequest; + path?: never; + query?: never; + url: '/api/stream-active-session'; +}; + +export type SetActiveSessionErrors = { + /** + * Unauthorized + */ + 401: unknown; + /** + * Connection not found + */ + 404: unknown; +}; + +export type SetActiveSessionResponses = { + /** + * Active session updated + */ + 200: unknown; +}; + export type GetConfigData = { body?: never; path: { diff --git a/packages/pi-client/src/global.d.ts b/packages/pi-client/src/global.d.ts new file mode 100644 index 0000000..b867229 --- /dev/null +++ b/packages/pi-client/src/global.d.ts @@ -0,0 +1 @@ +declare const __DEV__: boolean; diff --git a/packages/pi-client/src/hooks/use-agent-config.ts b/packages/pi-client/src/hooks/use-agent-config.ts index 884750c..14b1d1c 100644 --- a/packages/pi-client/src/hooks/use-agent-config.ts +++ b/packages/pi-client/src/hooks/use-agent-config.ts @@ -28,7 +28,6 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { const retryTimerRef = useRef | null>(null); const sessionIdRef = useRef(sessionId); - // Track sessionId changes to cancel stale retries useEffect(() => { sessionIdRef.current = sessionId; }, [sessionId]); @@ -40,20 +39,32 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { } }, []); - const load = useCallback( + // Subscribe to agent_state from SSE + useEffect(() => { + if (!sessionId) { + setState(null); + return; + } + + const sub = client.agentState$(sessionId).subscribe((agentState) => { + if (agentState && sessionIdRef.current === sessionId) { + setState(agentState); + } + }); + + return () => sub.unsubscribe(); + }, [client, sessionId]); + + // Fetch available models via REST (still needed, not in SSE) + const loadModels = useCallback( async (attempt = 0) => { if (!sessionId) return; setIsLoading(true); setError(null); try { - const [stateResult, modelsResult] = await Promise.all([ - client.api.getState(sessionId), - client.api.getAvailableModels(sessionId), - ]); - // Ignore result if sessionId changed while request was in flight + const modelsResult = await client.api.getAvailableModels(sessionId); if (sessionIdRef.current !== sessionId) return; - setState(stateResult as unknown as AgentStateData); setModels((modelsResult.models ?? []) as unknown as ModelInfo[]); attemptRef.current = 0; setIsLoading(false); @@ -65,12 +76,12 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { attemptRef.current = nextAttempt; retryTimerRef.current = setTimeout(() => { if (sessionIdRef.current === sessionId) { - load(nextAttempt); + loadModels(nextAttempt); } }, RETRY_DELAY_MS); } else { const message = - err instanceof Error ? err.message : "Failed to load toolbar configuration"; + err instanceof Error ? err.message : "Failed to load available models"; setError(message); setIsLoading(false); attemptRef.current = 0; @@ -81,23 +92,22 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { ); useEffect(() => { - // Reset state when sessionId changes clearRetryTimer(); attemptRef.current = 0; setError(null); - load(); + loadModels(); return () => { clearRetryTimer(); }; - }, [load, clearRetryTimer]); + }, [loadModels, clearRetryTimer]); const retry = useCallback(() => { clearRetryTimer(); attemptRef.current = 0; setError(null); - load(0); - }, [load, clearRetryTimer]); + loadModels(0); + }, [loadModels, clearRetryTimer]); const setModel = useCallback( async (params: { provider: string; modelId: string }) => { @@ -126,10 +136,10 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { try { await client.setModel(sessionId, params); } catch { - load(); + loadModels(); } }, - [client, sessionId, load, models], + [client, sessionId, loadModels, models], ); const setThinkingLevel = useCallback( @@ -148,10 +158,10 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { try { await client.setThinkingLevel(sessionId, level); } catch { - load(); + loadModels(); } }, - [client, sessionId, load], + [client, sessionId, loadModels], ); const setMode = useCallback( @@ -170,10 +180,10 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { try { await client.prompt(sessionId, mode === "plan" ? "/plan" : "/chat"); } catch { - load(); + loadModels(); } }, - [client, sessionId, load], + [client, sessionId, loadModels], ); return { @@ -184,7 +194,7 @@ export function useAgentConfig(sessionId: string | null): AgentConfigHandle { setModel, setThinkingLevel, setMode, - reload: load, + reload: loadModels, retry, }; } diff --git a/packages/pi-client/src/hooks/use-agent-session.ts b/packages/pi-client/src/hooks/use-agent-session.ts index bd135f5..9377cf4 100644 --- a/packages/pi-client/src/hooks/use-agent-session.ts +++ b/packages/pi-client/src/hooks/use-agent-session.ts @@ -34,6 +34,7 @@ const EMPTY: SessionState = { oldestEntryId: null, mode: "chat", pendingExtensionUiRequest: null, + agentState: null, }; export function useAgentSession( diff --git a/packages/pi-client/src/hooks/use-chat-sessions.ts b/packages/pi-client/src/hooks/use-chat-sessions.ts index 4b9fc33..cd07d27 100644 --- a/packages/pi-client/src/hooks/use-chat-sessions.ts +++ b/packages/pi-client/src/hooks/use-chat-sessions.ts @@ -3,8 +3,26 @@ import { BehaviorSubject } from "rxjs"; import { usePiClient } from "./context"; import { useObservable } from "./use-observable"; import type { SessionListItem } from "../types"; +import type { StreamEventEnvelope } from "../types/stream-events"; const PAGE_SIZE = 20; +const REFRESH_DEBOUNCE_MS = 350; + +function shouldRefreshChatSessions(event: StreamEventEnvelope): boolean { + if (!event.session_id || !!event.workspace_id) return false; + if (event.type === "client_command") { + const commandType = (event.data as { type?: string } | undefined)?.type; + return commandType === "prompt" || commandType === "steer" || commandType === "follow_up"; + } + return ( + event.type === "message_start" || + event.type === "message_end" || + event.type === "turn_end" || + event.type === "agent_end" || + event.type === "session_process_exited" || + event.type === "session_idle_timeout" + ); +} export interface ChatSessionsState { sessions: SessionListItem[]; @@ -37,8 +55,10 @@ export interface ChatSessionsHandle extends ChatSessionsState { } export function useChatSessions(): ChatSessionsHandle { - const { api } = usePiClient(); + const client = usePiClient(); + const { api } = client; const state$ = useRef(new BehaviorSubject(INITIAL_STATE)); + const refreshTimerRef = useRef | null>(null); const emit = useCallback( (patch: Partial) => @@ -80,6 +100,10 @@ export function useChatSessions(): ChatSessionsHandle { const fetchNextPage = useCallback(() => { const s = state$.current.value; if (!s.hasNextPage || s.isFetchingNextPage) return; + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + refreshTimerRef.current = null; + } emit({ isFetchingNextPage: true }); loadPage(s.page + 1, true); }, [loadPage, emit]); @@ -89,6 +113,24 @@ export function useChatSessions(): ChatSessionsHandle { loadPage(1, false); }, [loadPage, emit]); + const scheduleRefetch = useCallback(() => { + const current = state$.current.value; + if (current.isLoading || current.isFetchingNextPage || current.page > 1) return; + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + } + refreshTimerRef.current = setTimeout(() => { + const latest = state$.current.value; + if (latest.isLoading || latest.isFetchingNextPage || latest.page > 1) { + refreshTimerRef.current = null; + return; + } + refreshTimerRef.current = null; + emit({ isRefetching: true }); + loadPage(1, false); + }, REFRESH_DEBOUNCE_MS); + }, [emit, loadPage]); + const deleteSession = useCallback( async (sessionId: string) => { await api.deleteChatSession(sessionId); @@ -97,6 +139,23 @@ export function useChatSessions(): ChatSessionsHandle { [api, refetch], ); + useEffect(() => { + const subscription = client.events$.subscribe((event) => { + if (!shouldRefreshChatSessions(event)) return; + scheduleRefetch(); + }); + return () => subscription.unsubscribe(); + }, [client, scheduleRefetch]); + + useEffect(() => { + return () => { + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + refreshTimerRef.current = null; + } + }; + }, []); + const snapshot = useObservable(state$.current, INITIAL_STATE); const publicState = useMemo( () => ({ diff --git a/packages/pi-client/src/hooks/use-workspace-sessions.ts b/packages/pi-client/src/hooks/use-workspace-sessions.ts index 0d02ebb..ede2583 100644 --- a/packages/pi-client/src/hooks/use-workspace-sessions.ts +++ b/packages/pi-client/src/hooks/use-workspace-sessions.ts @@ -3,8 +3,26 @@ import { BehaviorSubject } from "rxjs"; import { usePiClient } from "./context"; import { useObservable } from "./use-observable"; import type { SessionListItem } from "../types"; +import type { StreamEventEnvelope } from "../types/stream-events"; const PAGE_SIZE = 20; +const REFRESH_DEBOUNCE_MS = 350; + +function shouldRefreshWorkspaceSessions(event: StreamEventEnvelope, workspaceId: string): boolean { + if (event.workspace_id !== workspaceId || !event.session_id) return false; + if (event.type === "client_command") { + const commandType = (event.data as { type?: string } | undefined)?.type; + return commandType === "prompt" || commandType === "steer" || commandType === "follow_up"; + } + return ( + event.type === "message_start" || + event.type === "message_end" || + event.type === "turn_end" || + event.type === "agent_end" || + event.type === "session_process_exited" || + event.type === "session_idle_timeout" + ); +} export interface WorkspaceSessionsState { sessions: SessionListItem[]; @@ -40,9 +58,11 @@ export interface WorkspaceSessionsHandle extends WorkspaceSessionsState { export function useWorkspaceSessions( workspaceId: string | null, ): WorkspaceSessionsHandle { - const { api } = usePiClient(); + const client = usePiClient(); + const { api } = client; const state$ = useRef(new BehaviorSubject(INITIAL_STATE)); const workspaceIdRef = useRef(workspaceId); + const refreshTimerRef = useRef | null>(null); workspaceIdRef.current = workspaceId; const emit = useCallback( @@ -93,6 +113,10 @@ export function useWorkspaceSessions( const fetchNextPage = useCallback(() => { const s = state$.current.value; if (!s.hasNextPage || s.isFetchingNextPage) return; + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + refreshTimerRef.current = null; + } emit({ isFetchingNextPage: true }); loadPage(s.page + 1, true); }, [loadPage, emit]); @@ -102,6 +126,24 @@ export function useWorkspaceSessions( loadPage(1, false); }, [loadPage, emit]); + const scheduleRefetch = useCallback(() => { + const current = state$.current.value; + if (current.isLoading || current.isFetchingNextPage || current.page > 1) return; + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + } + refreshTimerRef.current = setTimeout(() => { + const latest = state$.current.value; + if (latest.isLoading || latest.isFetchingNextPage || latest.page > 1) { + refreshTimerRef.current = null; + return; + } + refreshTimerRef.current = null; + emit({ isRefetching: true }); + loadPage(1, false); + }, REFRESH_DEBOUNCE_MS); + }, [emit, loadPage]); + const deleteSession = useCallback( async (sessionId: string) => { const wid = workspaceIdRef.current; @@ -112,6 +154,24 @@ export function useWorkspaceSessions( [api, refetch], ); + useEffect(() => { + if (!workspaceId) return; + const subscription = client.events$.subscribe((event) => { + if (!shouldRefreshWorkspaceSessions(event, workspaceId)) return; + scheduleRefetch(); + }); + return () => subscription.unsubscribe(); + }, [client, workspaceId, scheduleRefetch]); + + useEffect(() => { + return () => { + if (refreshTimerRef.current) { + clearTimeout(refreshTimerRef.current); + refreshTimerRef.current = null; + } + }; + }, []); + const snapshot = useObservable(state$.current, INITIAL_STATE); const publicState = useMemo( () => ({ diff --git a/packages/pi-client/src/types/chat-message.ts b/packages/pi-client/src/types/chat-message.ts index 6008d7e..7485e2a 100644 --- a/packages/pi-client/src/types/chat-message.ts +++ b/packages/pi-client/src/types/chat-message.ts @@ -25,12 +25,18 @@ export interface SubagentMeta { turns?: number; } +export interface ToolResultImage { + data: string; + mimeType: string; +} + export interface ToolCallInfo { id: string; name: string; arguments: string; status: "streaming" | "pending" | "running" | "complete" | "error"; result?: string; + resultImages?: ToolResultImage[]; isError?: boolean; partialResult?: string; progress?: SubagentProgress; @@ -54,6 +60,13 @@ export interface MessageUsageInfo { currency?: string; } +export interface TurnFileStats { + filesEdited: number; + filesCreated: number; + linesAdded: number; + linesRemoved: number; +} + export interface ChatMessage { id: string; entryId?: string; @@ -70,6 +83,8 @@ export interface ChatMessage { responseId?: string; usage?: MessageUsageInfo; stopReason?: StopReason; + turnDurationMs?: number; + turnFileStats?: TurnFileStats; systemKind?: "bashExecution" | "event"; command?: string; exitCode?: number; diff --git a/packages/pi-client/src/types/index.ts b/packages/pi-client/src/types/index.ts index 9be587b..478826e 100644 --- a/packages/pi-client/src/types/index.ts +++ b/packages/pi-client/src/types/index.ts @@ -32,6 +32,7 @@ export type { PathCompletion, SessionDetail, SessionEntry, + SessionHistoryResponse, SessionListItem, SessionTreeNode, TaskDefinition, @@ -60,7 +61,7 @@ export interface PiClientConfig { serverUrl: string; accessToken: string; onAuthError?: () => void; - onApiAuthError?: () => Promise; + onApiAuthError?: () => Promise; transport?: "sse" | "ws"; reconnectBaseMs?: number; reconnectMaxMs?: number; diff --git a/scripts/test-subagent-viewer.mjs b/scripts/test-subagent-viewer.mjs index ac5a6a0..a310a19 100644 --- a/scripts/test-subagent-viewer.mjs +++ b/scripts/test-subagent-viewer.mjs @@ -178,11 +178,70 @@ function testMessageEndRebuildsToolCallsFromFullContent() { assert.equal(state.messages[0]?.toolCalls?.[1]?.id, 'final-2'); } +function testToolResultMessageEndDoesNotOverwriteAssistantText() { + const state = apply([ + envelope(1, 'message_start', { + type: 'message_start', + message: { role: 'assistant', content: [], timestamp: 1 }, + }), + envelope(2, 'message_update', { + type: 'message_update', + assistantMessageEvent: { + type: 'toolcall_start', + contentIndex: 0, + partial: { id: 'read-1', name: 'read' }, + }, + }), + envelope(3, 'message_end', { + type: 'message_end', + message: { + role: 'assistant', + stopReason: 'toolUse', + timestamp: 3, + content: [ + { + type: 'toolCall', + id: 'read-1', + name: 'read', + arguments: { path: 'index.html', offset: 6 }, + }, + ], + }, + }), + envelope(4, 'tool_execution_end', { + type: 'tool_execution_end', + toolCallId: 'read-1', + toolName: 'read', + isError: false, + result: { + content: [{ type: 'text', text: '' }], + }, + }), + envelope(5, 'message_end', { + type: 'message_end', + message: { + role: 'toolResult', + toolCallId: 'read-1', + toolName: 'read', + isError: false, + timestamp: 5, + content: [{ type: 'text', text: '' }], + }, + }), + ]); + + assert.equal(state.messages.length, 1); + assert.equal(state.messages[0]?.role, 'assistant'); + assert.equal(state.messages[0]?.text, ''); + assert.equal(state.messages[0]?.toolCalls?.[0]?.result, ''); +} + function run() { testHistorySubagentResult(); testStreamingPartialResult(); testParallelToolcallContentIndexRouting(); testMessageEndRebuildsToolCallsFromFullContent(); + testToolResultMessageEndDoesNotOverwriteAssistantText(); console.log('subagent viewer reducer tests passed'); }