| 1 | //! Runtime HTTP/SSE API for local DeepSeek automation. |
| 2 | |
| 3 | use std::collections::HashSet; |
| 4 | use std::convert::Infallible; |
| 5 | use std::fs; |
| 6 | use std::net::SocketAddr; |
| 7 | use std::path::PathBuf; |
| 8 | use std::process::Command; |
| 9 | use std::sync::Arc; |
| 10 | use std::time::Duration; |
| 11 | |
| 12 | use anyhow::{Context, Result, anyhow, bail}; |
| 13 | use async_stream::stream; |
| 14 | use axum::extract::{Path, Query, Request, State}; |
| 15 | use axum::http::{HeaderValue, Method, StatusCode, header}; |
| 16 | use axum::middleware::{self, Next}; |
| 17 | use axum::response::sse::{Event as SseEvent, KeepAlive, Sse}; |
| 18 | use axum::response::{IntoResponse, Response}; |
| 19 | use axum::routing::{get, post}; |
| 20 | use axum::{Json, Router}; |
| 21 | use chrono::Utc; |
| 22 | use serde::{Deserialize, Serialize}; |
| 23 | use serde_json::{Value, json}; |
| 24 | use tokio::net::TcpListener; |
| 25 | use tokio::sync::Mutex; |
| 26 | use tokio_util::sync::CancellationToken; |
| 27 | use tower_http::cors::{Any, CorsLayer}; |
| 28 | |
| 29 | use crate::automation_manager::{ |
| 30 | AutomationManager, AutomationRecord, AutomationRunRecord, AutomationSchedulerConfig, |
| 31 | CreateAutomationRequest, SharedAutomationManager, UpdateAutomationRequest, spawn_scheduler, |
| 32 | }; |
| 33 | use crate::config::{Config, DEFAULT_TEXT_MODEL}; |
| 34 | use crate::mcp::{McpConfig, McpPool}; |
| 35 | use crate::runtime_threads::{ |
| 36 | CompactThreadRequest, CreateThreadRequest, RuntimeThreadManager, RuntimeThreadManagerConfig, |
| 37 | SharedRuntimeThreadManager, StartTurnRequest, SteerTurnRequest, ThreadDetail, ThreadListFilter, |
| 38 | ThreadRecord, TurnItemKind, TurnRecord, UpdateThreadRequest, UsageGroupBy, |
| 39 | }; |
| 40 | use crate::session_manager::{SavedSession, SessionManager, SessionMetadata, default_sessions_dir}; |
| 41 | use crate::skills::SkillRegistry; |
| 42 | use crate::task_manager::{ |
| 43 | NewTaskRequest, SharedTaskManager, TaskManager, TaskManagerConfig, TaskRecord, TaskSummary, |
| 44 | }; |
| 45 | |
| 46 | #[derive(Clone)] |
| 47 | pub struct RuntimeApiState { |
| 48 | config: Config, |
| 49 | workspace: PathBuf, |
| 50 | task_manager: SharedTaskManager, |
| 51 | runtime_threads: SharedRuntimeThreadManager, |
| 52 | cors_origins: Vec<String>, |
| 53 | sessions_dir: PathBuf, |
| 54 | mcp_config_path: PathBuf, |
| 55 | automations: SharedAutomationManager, |
| 56 | runtime_token: Option<String>, |
| 57 | } |
| 58 | |
| 59 | #[derive(Debug, Clone)] |
| 60 | pub struct RuntimeApiOptions { |
| 61 | pub host: String, |
| 62 | pub port: u16, |
| 63 | pub workers: usize, |
| 64 | /// Additional CORS origins to allow on top of the built-in defaults |
| 65 | /// (`http://localhost:{3000,1420}`, `http://127.0.0.1:{3000,1420}`, |
| 66 | /// `tauri://localhost`). Populated by `--cors-origin` (repeatable), |
| 67 | /// `DEEPSEEK_CORS_ORIGINS` (comma-separated), and `[runtime_api] |
| 68 | /// cors_origins` in `config.toml`. Whalescale#255 / #561. |
| 69 | pub cors_origins: Vec<String>, |
| 70 | /// Optional bearer token required for `/v1/*` routes. If omitted here, |
| 71 | /// `run_http_server` also checks `DEEPSEEK_RUNTIME_TOKEN`. |
| 72 | pub auth_token: Option<String>, |
| 73 | } |
| 74 | |
| 75 | impl Default for RuntimeApiOptions { |
| 76 | fn default() -> Self { |
| 77 | Self { |
| 78 | host: "127.0.0.1".to_string(), |
| 79 | port: 7878, |
| 80 | workers: 2, |
| 81 | cors_origins: Vec::new(), |
| 82 | auth_token: None, |
| 83 | } |
| 84 | } |
| 85 | } |
| 86 | |
| 87 | #[derive(Debug, Deserialize)] |
| 88 | struct StreamTurnRequest { |
| 89 | prompt: String, |
| 90 | model: Option<String>, |
| 91 | mode: Option<String>, |
| 92 | workspace: Option<PathBuf>, |
| 93 | allow_shell: Option<bool>, |
| 94 | trust_mode: Option<bool>, |
| 95 | auto_approve: Option<bool>, |
| 96 | } |
| 97 | |
| 98 | #[derive(Debug, Serialize)] |
| 99 | struct HealthResponse { |
| 100 | status: &'static str, |
| 101 | service: &'static str, |
| 102 | mode: &'static str, |
| 103 | } |
| 104 | |
| 105 | #[derive(Debug, Serialize)] |
| 106 | struct SessionsResponse { |
| 107 | sessions: Vec<SessionMetadata>, |
| 108 | } |
| 109 | |
| 110 | #[derive(Debug, Serialize)] |
| 111 | struct SessionDetailResponse { |
| 112 | metadata: SessionMetadata, |
| 113 | messages: Vec<serde_json::Value>, |
| 114 | system_prompt: Option<String>, |
| 115 | } |
| 116 | |
| 117 | #[derive(Debug, Deserialize)] |
| 118 | struct ResumeSessionRequest { |
| 119 | model: Option<String>, |
| 120 | mode: Option<String>, |
| 121 | } |
| 122 | |
| 123 | #[derive(Debug, Serialize)] |
| 124 | struct ResumeSessionResponse { |
| 125 | thread_id: String, |
| 126 | session_id: String, |
| 127 | message_count: usize, |
| 128 | summary: String, |
| 129 | } |
| 130 | |
| 131 | #[derive(Debug, Serialize)] |
| 132 | struct TasksResponse { |
| 133 | tasks: Vec<TaskSummary>, |
| 134 | counts: crate::task_manager::TaskCounts, |
| 135 | } |
| 136 | |
| 137 | #[derive(Debug, Deserialize)] |
| 138 | struct SessionsQuery { |
| 139 | limit: Option<usize>, |
| 140 | search: Option<String>, |
| 141 | } |
| 142 | |
| 143 | #[derive(Debug, Deserialize)] |
| 144 | struct TasksQuery { |
| 145 | limit: Option<usize>, |
| 146 | } |
| 147 | |
| 148 | #[derive(Debug, Deserialize)] |
| 149 | struct ThreadsQuery { |
| 150 | limit: Option<usize>, |
| 151 | include_archived: Option<bool>, |
| 152 | /// When `true`, returns archived threads only (overrides `include_archived`). |
| 153 | /// Whalescale#260 / #563. |
| 154 | archived_only: Option<bool>, |
| 155 | } |
| 156 | |
| 157 | #[derive(Debug, Deserialize)] |
| 158 | struct ThreadSummaryQuery { |
| 159 | limit: Option<usize>, |
| 160 | search: Option<String>, |
| 161 | include_archived: Option<bool>, |
| 162 | /// When `true`, returns archived threads only (overrides `include_archived`). |
| 163 | /// Whalescale#260 / #563. |
| 164 | archived_only: Option<bool>, |
| 165 | } |
| 166 | |
| 167 | fn resolve_thread_filter( |
| 168 | include_archived: Option<bool>, |
| 169 | archived_only: Option<bool>, |
| 170 | ) -> ThreadListFilter { |
| 171 | if archived_only.unwrap_or(false) { |
| 172 | ThreadListFilter::ArchivedOnly |
| 173 | } else if include_archived.unwrap_or(false) { |
| 174 | ThreadListFilter::IncludeArchived |
| 175 | } else { |
| 176 | ThreadListFilter::ActiveOnly |
| 177 | } |
| 178 | } |
| 179 | |
| 180 | #[derive(Debug, Serialize)] |
| 181 | struct ThreadSummary { |
| 182 | id: String, |
| 183 | title: String, |
| 184 | preview: String, |
| 185 | model: String, |
| 186 | mode: String, |
| 187 | archived: bool, |
| 188 | updated_at: chrono::DateTime<Utc>, |
| 189 | latest_turn_id: Option<String>, |
| 190 | latest_turn_status: Option<String>, |
| 191 | } |
| 192 | |
| 193 | #[derive(Debug, Serialize)] |
| 194 | struct WorkspaceStatusResponse { |
| 195 | workspace: PathBuf, |
| 196 | git_repo: bool, |
| 197 | branch: Option<String>, |
| 198 | staged: usize, |
| 199 | unstaged: usize, |
| 200 | untracked: usize, |
| 201 | ahead: Option<u32>, |
| 202 | behind: Option<u32>, |
| 203 | } |
| 204 | |
| 205 | #[derive(Debug, Serialize)] |
| 206 | struct SkillEntry { |
| 207 | name: String, |
| 208 | description: String, |
| 209 | path: PathBuf, |
| 210 | } |
| 211 | |
| 212 | #[derive(Debug, Serialize)] |
| 213 | struct SkillsResponse { |
| 214 | directory: PathBuf, |
| 215 | warnings: Vec<String>, |
| 216 | skills: Vec<SkillEntry>, |
| 217 | } |
| 218 | |
| 219 | #[derive(Debug, Serialize)] |
| 220 | struct McpServerEntry { |
| 221 | name: String, |
| 222 | enabled: bool, |
| 223 | required: bool, |
| 224 | command: Option<String>, |
| 225 | url: Option<String>, |
| 226 | connected: bool, |
| 227 | enabled_tools: Vec<String>, |
| 228 | disabled_tools: Vec<String>, |
| 229 | } |
| 230 | |
| 231 | #[derive(Debug, Serialize)] |
| 232 | struct McpServersResponse { |
| 233 | servers: Vec<McpServerEntry>, |
| 234 | } |
| 235 | |
| 236 | #[derive(Debug, Deserialize)] |
| 237 | struct McpToolsQuery { |
| 238 | server: Option<String>, |
| 239 | } |
| 240 | |
| 241 | #[derive(Debug, Serialize)] |
| 242 | struct McpToolEntry { |
| 243 | server: String, |
| 244 | name: String, |
| 245 | prefixed_name: String, |
| 246 | description: Option<String>, |
| 247 | input_schema: Value, |
| 248 | } |
| 249 | |
| 250 | #[derive(Debug, Serialize)] |
| 251 | struct McpToolsResponse { |
| 252 | tools: Vec<McpToolEntry>, |
| 253 | } |
| 254 | |
| 255 | #[derive(Debug, Deserialize)] |
| 256 | struct AutomationRunsQuery { |
| 257 | limit: Option<usize>, |
| 258 | } |
| 259 | |
| 260 | #[derive(Debug, Deserialize)] |
| 261 | struct ThreadEventsQuery { |
| 262 | since_seq: Option<u64>, |
| 263 | } |
| 264 | |
| 265 | #[derive(Debug, Serialize)] |
| 266 | struct StartTurnResponse { |
| 267 | thread: ThreadRecord, |
| 268 | turn: TurnRecord, |
| 269 | } |
| 270 | |
| 271 | /// Start the runtime API server. |
| 272 | pub async fn run_http_server( |
| 273 | config: Config, |
| 274 | workspace: PathBuf, |
| 275 | options: RuntimeApiOptions, |
| 276 | ) -> Result<()> { |
| 277 | if options.port == 0 { |
| 278 | bail!("Port must be > 0"); |
| 279 | } |
| 280 | |
| 281 | let task_cfg = TaskManagerConfig::from_runtime( |
| 282 | &config, |
| 283 | workspace.clone(), |
| 284 | config.default_text_model.clone(), |
| 285 | Some(options.workers), |
| 286 | ); |
| 287 | let runtime_threads = Arc::new(RuntimeThreadManager::open( |
| 288 | config.clone(), |
| 289 | workspace.clone(), |
| 290 | RuntimeThreadManagerConfig::from_task_data_dir(task_cfg.data_dir.clone()), |
| 291 | )?); |
| 292 | let task_manager = |
| 293 | TaskManager::start_with_runtime_manager(task_cfg, config.clone(), runtime_threads.clone()) |
| 294 | .await?; |
| 295 | let automations = Arc::new(Mutex::new(AutomationManager::default_location()?)); |
| 296 | runtime_threads.attach_automation_manager(automations.clone()); |
| 297 | let scheduler_cancel = CancellationToken::new(); |
| 298 | let scheduler_handle = spawn_scheduler( |
| 299 | automations.clone(), |
| 300 | task_manager.clone(), |
| 301 | scheduler_cancel.clone(), |
| 302 | AutomationSchedulerConfig::default(), |
| 303 | ); |
| 304 | |
| 305 | let sessions_dir = default_sessions_dir().unwrap_or_else(|_| { |
| 306 | dirs::home_dir() |
| 307 | .map(|h| h.join(".deepseek").join("sessions")) |
| 308 | .unwrap_or_else(|| PathBuf::from(".deepseek").join("sessions")) |
| 309 | }); |
| 310 | let runtime_token = options |
| 311 | .auth_token |
| 312 | .clone() |
| 313 | .or_else(|| std::env::var("DEEPSEEK_RUNTIME_TOKEN").ok()) |
| 314 | .filter(|token| !token.trim().is_empty()); |
| 315 | let auth_enabled = runtime_token.is_some(); |
| 316 | let state = RuntimeApiState { |
| 317 | config: config.clone(), |
| 318 | workspace, |
| 319 | task_manager, |
| 320 | runtime_threads, |
| 321 | cors_origins: options.cors_origins.clone(), |
| 322 | sessions_dir, |
| 323 | mcp_config_path: config.mcp_config_path(), |
| 324 | automations, |
| 325 | runtime_token, |
| 326 | }; |
| 327 | let app = build_router(state); |
| 328 | |
| 329 | let addr: SocketAddr = format!("{}:{}", options.host, options.port) |
| 330 | .parse() |
| 331 | .with_context(|| format!("Invalid bind address '{}:{}'", options.host, options.port))?; |
| 332 | let listener = TcpListener::bind(addr) |
| 333 | .await |
| 334 | .with_context(|| format!("Failed to bind {addr}"))?; |
| 335 | |
| 336 | println!("Runtime API listening on http://{addr}"); |
| 337 | println!("Security: this server is local-first. Do not expose it to untrusted networks."); |
| 338 | if auth_enabled { |
| 339 | println!("Runtime API auth: bearer token required for /v1/* routes."); |
| 340 | } |
| 341 | let serve_result = axum::serve(listener, app) |
| 342 | .await |
| 343 | .map_err(|e| anyhow!("Runtime API server error: {e}")); |
| 344 | scheduler_cancel.cancel(); |
| 345 | scheduler_handle.abort(); |
| 346 | serve_result |
| 347 | } |
| 348 | |
| 349 | pub fn build_router(state: RuntimeApiState) -> Router { |
| 350 | let api_routes = Router::new() |
| 351 | .route("/v1/sessions", get(list_sessions)) |
| 352 | .route("/v1/sessions/{id}", get(get_session).delete(delete_session)) |
| 353 | .route( |
| 354 | "/v1/sessions/{id}/resume-thread", |
| 355 | post(resume_session_thread), |
| 356 | ) |
| 357 | .route("/v1/workspace/status", get(workspace_status)) |
| 358 | .route("/v1/stream", post(stream_turn)) |
| 359 | .route("/v1/threads", get(list_threads).post(create_thread)) |
| 360 | .route("/v1/threads/summary", get(list_threads_summary)) |
| 361 | .route("/v1/threads/{id}", get(get_thread).patch(update_thread)) |
| 362 | .route("/v1/threads/{id}/resume", post(resume_thread)) |
| 363 | .route("/v1/threads/{id}/fork", post(fork_thread)) |
| 364 | .route("/v1/threads/{id}/turns", post(start_thread_turn)) |
| 365 | .route( |
| 366 | "/v1/threads/{id}/turns/{turn_id}/steer", |
| 367 | post(steer_thread_turn), |
| 368 | ) |
| 369 | .route( |
| 370 | "/v1/threads/{id}/turns/{turn_id}/interrupt", |
| 371 | post(interrupt_thread_turn), |
| 372 | ) |
| 373 | .route("/v1/threads/{id}/compact", post(compact_thread)) |
| 374 | .route("/v1/threads/{id}/events", get(stream_thread_events)) |
| 375 | .route("/v1/tasks", get(list_tasks).post(create_task)) |
| 376 | .route("/v1/tasks/{id}", get(get_task)) |
| 377 | .route("/v1/tasks/{id}/cancel", post(cancel_task)) |
| 378 | .route("/v1/skills", get(list_skills)) |
| 379 | .route("/v1/apps/mcp/servers", get(list_mcp_servers)) |
| 380 | .route("/v1/apps/mcp/tools", get(list_mcp_tools)) |
| 381 | .route( |
| 382 | "/v1/automations", |
| 383 | get(list_automations).post(create_automation), |
| 384 | ) |
| 385 | .route( |
| 386 | "/v1/automations/{id}", |
| 387 | get(get_automation) |
| 388 | .patch(update_automation) |
| 389 | .delete(delete_automation), |
| 390 | ) |
| 391 | .route("/v1/automations/{id}/run", post(run_automation)) |
| 392 | .route("/v1/automations/{id}/pause", post(pause_automation)) |
| 393 | .route("/v1/automations/{id}/resume", post(resume_automation)) |
| 394 | .route("/v1/automations/{id}/runs", get(list_automation_runs)) |
| 395 | .route("/v1/usage", get(get_usage)) |
| 396 | .route_layer(middleware::from_fn_with_state( |
| 397 | state.clone(), |
| 398 | require_runtime_token, |
| 399 | )); |
| 400 | |
| 401 | Router::new() |
| 402 | .route("/health", get(health)) |
| 403 | .merge(api_routes) |
| 404 | .layer(cors_layer(&state.cors_origins)) |
| 405 | .with_state(state) |
| 406 | } |
| 407 | |
| 408 | async fn require_runtime_token( |
| 409 | State(state): State<RuntimeApiState>, |
| 410 | req: Request, |
| 411 | next: Next, |
| 412 | ) -> Response { |
| 413 | let Some(expected) = state.runtime_token.as_deref() else { |
| 414 | return next.run(req).await; |
| 415 | }; |
| 416 | let authorized = req |
| 417 | .headers() |
| 418 | .get(header::AUTHORIZATION) |
| 419 | .and_then(|value| value.to_str().ok()) |
| 420 | .and_then(|raw| raw.strip_prefix("Bearer ")) |
| 421 | .is_some_and(|token| token == expected) |
| 422 | || req |
| 423 | .headers() |
| 424 | .get("x-deepseek-runtime-token") |
| 425 | .and_then(|value| value.to_str().ok()) |
| 426 | .is_some_and(|token| token == expected) |
| 427 | || token_from_query(req.uri().query()).is_some_and(|token| token == expected); |
| 428 | |
| 429 | if authorized { |
| 430 | next.run(req).await |
| 431 | } else { |
| 432 | ( |
| 433 | StatusCode::UNAUTHORIZED, |
| 434 | Json(json!({ |
| 435 | "error": { |
| 436 | "message": "runtime API bearer token required", |
| 437 | "status": StatusCode::UNAUTHORIZED.as_u16(), |
| 438 | } |
| 439 | })), |
| 440 | ) |
| 441 | .into_response() |
| 442 | } |
| 443 | } |
| 444 | |
| 445 | fn token_from_query(query: Option<&str>) -> Option<&str> { |
| 446 | query.and_then(|query| { |
| 447 | query.split('&').find_map(|pair| { |
| 448 | let (key, value) = pair.split_once('=')?; |
| 449 | (key == "token").then_some(value) |
| 450 | }) |
| 451 | }) |
| 452 | } |
| 453 | |
| 454 | async fn health() -> Json<HealthResponse> { |
| 455 | Json(HealthResponse { |
| 456 | status: "ok", |
| 457 | service: "deepseek-runtime-api", |
| 458 | mode: "local", |
| 459 | }) |
| 460 | } |
| 461 | |
| 462 | async fn list_sessions( |
| 463 | State(state): State<RuntimeApiState>, |
| 464 | Query(query): Query<SessionsQuery>, |
| 465 | ) -> Result<Json<SessionsResponse>, ApiError> { |
| 466 | let manager = SessionManager::new(state.sessions_dir.clone()) |
| 467 | .map_err(|e| ApiError::internal(format!("Failed to open sessions dir: {e}")))?; |
| 468 | let mut sessions = if let Some(search) = query.search { |
| 469 | manager |
| 470 | .search_sessions(&search) |
| 471 | .map_err(|e| ApiError::internal(format!("Failed to search sessions: {e}")))? |
| 472 | } else { |
| 473 | manager |
| 474 | .list_sessions() |
| 475 | .map_err(|e| ApiError::internal(format!("Failed to list sessions: {e}")))? |
| 476 | }; |
| 477 | let limit = query.limit.unwrap_or(50).clamp(1, 500); |
| 478 | sessions.truncate(limit); |
| 479 | Ok(Json(SessionsResponse { sessions })) |
| 480 | } |
| 481 | |
| 482 | async fn get_session( |
| 483 | State(state): State<RuntimeApiState>, |
| 484 | Path(id): Path<String>, |
| 485 | ) -> Result<Json<SessionDetailResponse>, ApiError> { |
| 486 | let manager = SessionManager::new(state.sessions_dir.clone()) |
| 487 | .map_err(|e| ApiError::internal(format!("Failed to open sessions dir: {e}")))?; |
| 488 | let session = manager |
| 489 | .load_session(&id) |
| 490 | .map_err(|e| map_session_err(&id, e, "read"))?; |
| 491 | Ok(Json(session_to_detail(session))) |
| 492 | } |
| 493 | |
| 494 | async fn resume_session_thread( |
| 495 | State(state): State<RuntimeApiState>, |
| 496 | Path(id): Path<String>, |
| 497 | Json(req): Json<ResumeSessionRequest>, |
| 498 | ) -> Result<(StatusCode, Json<ResumeSessionResponse>), ApiError> { |
| 499 | let manager = SessionManager::new(state.sessions_dir.clone()) |
| 500 | .map_err(|e| ApiError::internal(format!("Failed to open sessions dir: {e}")))?; |
| 501 | let session = manager |
| 502 | .load_session(&id) |
| 503 | .map_err(|e| map_session_err(&id, e, "read"))?; |
| 504 | |
| 505 | let model = req.model.unwrap_or_else(|| session.metadata.model.clone()); |
| 506 | let mode = req.mode.unwrap_or_else(|| { |
| 507 | session |
| 508 | .metadata |
| 509 | .mode |
| 510 | .clone() |
| 511 | .unwrap_or_else(|| "agent".to_string()) |
| 512 | }); |
| 513 | |
| 514 | let thread = state |
| 515 | .runtime_threads |
| 516 | .create_thread(CreateThreadRequest { |
| 517 | model: Some(model), |
| 518 | workspace: Some(state.workspace.clone()), |
| 519 | mode: Some(mode), |
| 520 | allow_shell: None, |
| 521 | trust_mode: None, |
| 522 | auto_approve: None, |
| 523 | archived: false, |
| 524 | system_prompt: session.system_prompt.clone(), |
| 525 | task_id: None, |
| 526 | }) |
| 527 | .await |
| 528 | .map_err(|e| ApiError::internal(format!("Failed to create thread: {e}")))?; |
| 529 | |
| 530 | let msg_count = session.messages.len(); |
| 531 | state |
| 532 | .runtime_threads |
| 533 | .seed_thread_from_messages(&thread.id, &session.messages) |
| 534 | .await |
| 535 | .map_err(|e| ApiError::internal(format!("Failed to seed thread history: {e}")))?; |
| 536 | |
| 537 | let summary = format!( |
| 538 | "Resumed session '{}' ({} messages) into thread {}", |
| 539 | session.metadata.title, msg_count, thread.id |
| 540 | ); |
| 541 | |
| 542 | Ok(( |
| 543 | StatusCode::CREATED, |
| 544 | Json(ResumeSessionResponse { |
| 545 | thread_id: thread.id, |
| 546 | session_id: id, |
| 547 | message_count: msg_count, |
| 548 | summary, |
| 549 | }), |
| 550 | )) |
| 551 | } |
| 552 | |
| 553 | async fn delete_session( |
| 554 | State(state): State<RuntimeApiState>, |
| 555 | Path(id): Path<String>, |
| 556 | ) -> Result<StatusCode, ApiError> { |
| 557 | let manager = SessionManager::new(state.sessions_dir.clone()) |
| 558 | .map_err(|e| ApiError::internal(format!("Failed to open sessions dir: {e}")))?; |
| 559 | manager |
| 560 | .delete_session(&id) |
| 561 | .map_err(|e| map_session_err(&id, e, "delete"))?; |
| 562 | Ok(StatusCode::NO_CONTENT) |
| 563 | } |
| 564 | |
| 565 | fn session_to_detail(session: SavedSession) -> SessionDetailResponse { |
| 566 | let messages: Vec<serde_json::Value> = session |
| 567 | .messages |
| 568 | .iter() |
| 569 | .map(|msg| { |
| 570 | let content_blocks: Vec<serde_json::Value> = msg |
| 571 | .content |
| 572 | .iter() |
| 573 | .map(|block| match block { |
| 574 | crate::models::ContentBlock::Text { text, .. } => { |
| 575 | json!({ "type": "text", "text": text }) |
| 576 | } |
| 577 | crate::models::ContentBlock::Thinking { thinking, .. } => { |
| 578 | json!({ "type": "thinking", "text": thinking }) |
| 579 | } |
| 580 | _ => json!({ "type": "other" }), |
| 581 | }) |
| 582 | .collect(); |
| 583 | json!({ |
| 584 | "role": msg.role, |
| 585 | "content": content_blocks, |
| 586 | }) |
| 587 | }) |
| 588 | .collect(); |
| 589 | SessionDetailResponse { |
| 590 | metadata: session.metadata, |
| 591 | messages, |
| 592 | system_prompt: session.system_prompt, |
| 593 | } |
| 594 | } |
| 595 | |
| 596 | fn map_session_err(id: &str, err: std::io::Error, action: &str) -> ApiError { |
| 597 | match err.kind() { |
| 598 | std::io::ErrorKind::NotFound => ApiError::not_found(format!("Session '{id}' not found")), |
| 599 | std::io::ErrorKind::InvalidData => { |
| 600 | ApiError::bad_request(format!("Failed to parse session '{id}': {err}")) |
| 601 | } |
| 602 | std::io::ErrorKind::InvalidInput => { |
| 603 | ApiError::bad_request(format!("Invalid session id '{id}'")) |
| 604 | } |
| 605 | _ => ApiError::internal(format!("Failed to {action} session '{id}': {err}")), |
| 606 | } |
| 607 | } |
| 608 | |
| 609 | async fn create_task( |
| 610 | State(state): State<RuntimeApiState>, |
| 611 | Json(mut req): Json<NewTaskRequest>, |
| 612 | ) -> Result<(StatusCode, Json<TaskRecord>), ApiError> { |
| 613 | if req.prompt.trim().is_empty() { |
| 614 | return Err(ApiError::bad_request("prompt is required")); |
| 615 | } |
| 616 | if req.workspace.is_none() { |
| 617 | req.workspace = Some(state.workspace.clone()); |
| 618 | } |
| 619 | if req.model.is_none() { |
| 620 | req.model = Some( |
| 621 | state |
| 622 | .config |
| 623 | .default_text_model |
| 624 | .clone() |
| 625 | .unwrap_or_else(|| DEFAULT_TEXT_MODEL.to_string()), |
| 626 | ); |
| 627 | } |
| 628 | let task = state |
| 629 | .task_manager |
| 630 | .add_task(req) |
| 631 | .await |
| 632 | .map_err(|e| ApiError::bad_request(e.to_string()))?; |
| 633 | Ok((StatusCode::CREATED, Json(task))) |
| 634 | } |
| 635 | |
| 636 | async fn create_thread( |
| 637 | State(state): State<RuntimeApiState>, |
| 638 | Json(mut req): Json<CreateThreadRequest>, |
| 639 | ) -> Result<(StatusCode, Json<ThreadRecord>), ApiError> { |
| 640 | if req.model.as_ref().is_none_or(|m| m.trim().is_empty()) { |
| 641 | req.model = Some( |
| 642 | state |
| 643 | .config |
| 644 | .default_text_model |
| 645 | .clone() |
| 646 | .unwrap_or_else(|| DEFAULT_TEXT_MODEL.to_string()), |
| 647 | ); |
| 648 | } |
| 649 | if req.workspace.is_none() { |
| 650 | req.workspace = Some(state.workspace.clone()); |
| 651 | } |
| 652 | if req.mode.as_ref().is_none_or(|m| m.trim().is_empty()) { |
| 653 | req.mode = Some("agent".to_string()); |
| 654 | } |
| 655 | |
| 656 | let thread = state |
| 657 | .runtime_threads |
| 658 | .create_thread(req) |
| 659 | .await |
| 660 | .map_err(|e| ApiError::bad_request(e.to_string()))?; |
| 661 | Ok((StatusCode::CREATED, Json(thread))) |
| 662 | } |
| 663 | |
| 664 | async fn list_threads( |
| 665 | State(state): State<RuntimeApiState>, |
| 666 | Query(query): Query<ThreadsQuery>, |
| 667 | ) -> Result<Json<Vec<ThreadRecord>>, ApiError> { |
| 668 | let filter = resolve_thread_filter(query.include_archived, query.archived_only); |
| 669 | let threads = state |
| 670 | .runtime_threads |
| 671 | .list_threads(filter, query.limit) |
| 672 | .await |
| 673 | .map_err(|e| ApiError::internal(e.to_string()))?; |
| 674 | Ok(Json(threads)) |
| 675 | } |
| 676 | |
| 677 | async fn list_threads_summary( |
| 678 | State(state): State<RuntimeApiState>, |
| 679 | Query(query): Query<ThreadSummaryQuery>, |
| 680 | ) -> Result<Json<Vec<ThreadSummary>>, ApiError> { |
| 681 | let limit = query.limit.unwrap_or(50).clamp(1, 500); |
| 682 | let search = query.search.as_deref().map(str::to_ascii_lowercase); |
| 683 | let filter = resolve_thread_filter(query.include_archived, query.archived_only); |
| 684 | let threads = state |
| 685 | .runtime_threads |
| 686 | .list_threads(filter, Some(limit)) |
| 687 | .await |
| 688 | .map_err(|e| ApiError::internal(e.to_string()))?; |
| 689 | |
| 690 | let mut summaries = Vec::new(); |
| 691 | for thread in threads { |
| 692 | let detail = state |
| 693 | .runtime_threads |
| 694 | .get_thread_detail(&thread.id) |
| 695 | .await |
| 696 | .map_err(map_thread_err)?; |
| 697 | let latest_turn = detail.turns.last(); |
| 698 | let latest_status = |
| 699 | latest_turn.map(|turn| format!("{:?}", turn.status).to_ascii_lowercase()); |
| 700 | |
| 701 | let title = thread |
| 702 | .title |
| 703 | .as_deref() |
| 704 | .map(str::trim) |
| 705 | .filter(|t| !t.is_empty()) |
| 706 | .map(|t| truncate_text(t, 72)) |
| 707 | .unwrap_or_else(|| { |
| 708 | latest_turn |
| 709 | .map(|turn| { |
| 710 | if turn.input_summary.trim().is_empty() { |
| 711 | "New Thread".to_string() |
| 712 | } else { |
| 713 | truncate_text(&turn.input_summary, 72) |
| 714 | } |
| 715 | }) |
| 716 | .unwrap_or_else(|| "New Thread".to_string()) |
| 717 | }); |
| 718 | |
| 719 | let preview = detail |
| 720 | .items |
| 721 | .iter() |
| 722 | .rev() |
| 723 | .find_map(|item| match item.kind { |
| 724 | TurnItemKind::AgentMessage | TurnItemKind::UserMessage => { |
| 725 | let text = item.detail.clone().unwrap_or_else(|| item.summary.clone()); |
| 726 | if text.trim().is_empty() { |
| 727 | None |
| 728 | } else { |
| 729 | Some(truncate_text(&text, 140)) |
| 730 | } |
| 731 | } |
| 732 | _ => None, |
| 733 | }) |
| 734 | .unwrap_or_else(|| title.clone()); |
| 735 | |
| 736 | if let Some(search) = &search { |
| 737 | let haystack = format!( |
| 738 | "{} {} {} {}", |
| 739 | thread.id.to_ascii_lowercase(), |
| 740 | title.to_ascii_lowercase(), |
| 741 | preview.to_ascii_lowercase(), |
| 742 | thread.model.to_ascii_lowercase() |
| 743 | ); |
| 744 | if !haystack.contains(search) { |
| 745 | continue; |
| 746 | } |
| 747 | } |
| 748 | |
| 749 | summaries.push(ThreadSummary { |
| 750 | id: thread.id, |
| 751 | title, |
| 752 | preview, |
| 753 | model: thread.model, |
| 754 | mode: thread.mode, |
| 755 | archived: thread.archived, |
| 756 | updated_at: thread.updated_at, |
| 757 | latest_turn_id: thread.latest_turn_id, |
| 758 | latest_turn_status: latest_status, |
| 759 | }); |
| 760 | } |
| 761 | |
| 762 | if summaries.len() > limit { |
| 763 | summaries.truncate(limit); |
| 764 | } |
| 765 | |
| 766 | Ok(Json(summaries)) |
| 767 | } |
| 768 | |
| 769 | async fn workspace_status( |
| 770 | State(state): State<RuntimeApiState>, |
| 771 | ) -> Result<Json<WorkspaceStatusResponse>, ApiError> { |
| 772 | Ok(Json(collect_workspace_status(&state.workspace))) |
| 773 | } |
| 774 | |
| 775 | async fn list_skills( |
| 776 | State(state): State<RuntimeApiState>, |
| 777 | ) -> Result<Json<SkillsResponse>, ApiError> { |
| 778 | let skills_dir = resolve_skills_dir(&state.config, &state.workspace); |
| 779 | let registry = SkillRegistry::discover(&skills_dir); |
| 780 | let skills = registry |
| 781 | .list() |
| 782 | .iter() |
| 783 | .map(|skill| SkillEntry { |
| 784 | name: skill.name.clone(), |
| 785 | description: skill.description.clone(), |
| 786 | path: skills_dir.join(&skill.name).join("SKILL.md"), |
| 787 | }) |
| 788 | .collect(); |
| 789 | Ok(Json(SkillsResponse { |
| 790 | directory: skills_dir, |
| 791 | warnings: registry.warnings().to_vec(), |
| 792 | skills, |
| 793 | })) |
| 794 | } |
| 795 | |
| 796 | async fn list_mcp_servers( |
| 797 | State(state): State<RuntimeApiState>, |
| 798 | ) -> Result<Json<McpServersResponse>, ApiError> { |
| 799 | let config = load_mcp_config_or_default(&state.mcp_config_path)?; |
| 800 | let mut pool = McpPool::new(config.clone()); |
| 801 | let _errors = pool.connect_all().await; |
| 802 | let connected: HashSet<String> = pool |
| 803 | .connected_servers() |
| 804 | .into_iter() |
| 805 | .map(str::to_string) |
| 806 | .collect(); |
| 807 | |
| 808 | let mut servers = Vec::new(); |
| 809 | for (name, server_cfg) in config.servers { |
| 810 | servers.push(McpServerEntry { |
| 811 | name: name.clone(), |
| 812 | enabled: server_cfg.is_enabled(), |
| 813 | required: server_cfg.required, |
| 814 | command: server_cfg.command.clone(), |
| 815 | url: server_cfg.url.clone(), |
| 816 | connected: connected.contains(&name), |
| 817 | enabled_tools: server_cfg.enabled_tools.clone(), |
| 818 | disabled_tools: server_cfg.disabled_tools.clone(), |
| 819 | }); |
| 820 | } |
| 821 | servers.sort_by(|a, b| a.name.cmp(&b.name)); |
| 822 | |
| 823 | Ok(Json(McpServersResponse { servers })) |
| 824 | } |
| 825 | |
| 826 | async fn list_mcp_tools( |
| 827 | State(state): State<RuntimeApiState>, |
| 828 | Query(query): Query<McpToolsQuery>, |
| 829 | ) -> Result<Json<McpToolsResponse>, ApiError> { |
| 830 | let mut pool = McpPool::from_config_path(&state.mcp_config_path) |
| 831 | .map_err(|e| ApiError::internal(format!("Failed to load MCP config: {e}")))?; |
| 832 | let _errors = pool.connect_all().await; |
| 833 | |
| 834 | let mut tools = Vec::new(); |
| 835 | for (prefixed_name, tool) in pool.all_tools() { |
| 836 | let Some(rest) = prefixed_name.strip_prefix("mcp_") else { |
| 837 | continue; |
| 838 | }; |
| 839 | let Some((server, name)) = rest.split_once('_') else { |
| 840 | continue; |
| 841 | }; |
| 842 | |
| 843 | if let Some(filter) = query.server.as_deref() |
| 844 | && server != filter |
| 845 | { |
| 846 | continue; |
| 847 | } |
| 848 | |
| 849 | tools.push(McpToolEntry { |
| 850 | server: server.to_string(), |
| 851 | name: name.to_string(), |
| 852 | prefixed_name, |
| 853 | description: tool.description.clone(), |
| 854 | input_schema: tool.input_schema.clone(), |
| 855 | }); |
| 856 | } |
| 857 | |
| 858 | tools.sort_by(|a, b| a.server.cmp(&b.server).then_with(|| a.name.cmp(&b.name))); |
| 859 | |
| 860 | Ok(Json(McpToolsResponse { tools })) |
| 861 | } |
| 862 | |
| 863 | async fn list_automations( |
| 864 | State(state): State<RuntimeApiState>, |
| 865 | ) -> Result<Json<Vec<AutomationRecord>>, ApiError> { |
| 866 | let manager = state.automations.lock().await; |
| 867 | let automations = manager |
| 868 | .list_automations() |
| 869 | .map_err(|e| ApiError::internal(format!("Failed to list automations: {e}")))?; |
| 870 | Ok(Json(automations)) |
| 871 | } |
| 872 | |
| 873 | async fn create_automation( |
| 874 | State(state): State<RuntimeApiState>, |
| 875 | Json(req): Json<CreateAutomationRequest>, |
| 876 | ) -> Result<(StatusCode, Json<AutomationRecord>), ApiError> { |
| 877 | let manager = state.automations.lock().await; |
| 878 | let automation = manager |
| 879 | .create_automation(req) |
| 880 | .map_err(|e| ApiError::bad_request(e.to_string()))?; |
| 881 | Ok((StatusCode::CREATED, Json(automation))) |
| 882 | } |
| 883 | |
| 884 | async fn get_automation( |
| 885 | State(state): State<RuntimeApiState>, |
| 886 | Path(id): Path<String>, |
| 887 | ) -> Result<Json<AutomationRecord>, ApiError> { |
| 888 | let manager = state.automations.lock().await; |
| 889 | let automation = manager.get_automation(&id).map_err(map_automation_err)?; |
| 890 | Ok(Json(automation)) |
| 891 | } |
| 892 | |
| 893 | async fn update_automation( |
| 894 | State(state): State<RuntimeApiState>, |
| 895 | Path(id): Path<String>, |
| 896 | Json(req): Json<UpdateAutomationRequest>, |
| 897 | ) -> Result<Json<AutomationRecord>, ApiError> { |
| 898 | let manager = state.automations.lock().await; |
| 899 | let automation = manager |
| 900 | .update_automation(&id, req) |
| 901 | .map_err(map_automation_err)?; |
| 902 | Ok(Json(automation)) |
| 903 | } |
| 904 | |
| 905 | async fn delete_automation( |
| 906 | State(state): State<RuntimeApiState>, |
| 907 | Path(id): Path<String>, |
| 908 | ) -> Result<Json<AutomationRecord>, ApiError> { |
| 909 | let manager = state.automations.lock().await; |
| 910 | let automation = manager.delete_automation(&id).map_err(map_automation_err)?; |
| 911 | Ok(Json(automation)) |
| 912 | } |
| 913 | |
| 914 | async fn run_automation( |
| 915 | State(state): State<RuntimeApiState>, |
| 916 | Path(id): Path<String>, |
| 917 | ) -> Result<Json<AutomationRunRecord>, ApiError> { |
| 918 | let manager = state.automations.lock().await; |
| 919 | let run = manager |
| 920 | .run_now(&id, &state.task_manager) |
| 921 | .await |
| 922 | .map_err(map_automation_err)?; |
| 923 | Ok(Json(run)) |
| 924 | } |
| 925 | |
| 926 | async fn pause_automation( |
| 927 | State(state): State<RuntimeApiState>, |
| 928 | Path(id): Path<String>, |
| 929 | ) -> Result<Json<AutomationRecord>, ApiError> { |
| 930 | let manager = state.automations.lock().await; |
| 931 | let automation = manager.pause_automation(&id).map_err(map_automation_err)?; |
| 932 | Ok(Json(automation)) |
| 933 | } |
| 934 | |
| 935 | async fn resume_automation( |
| 936 | State(state): State<RuntimeApiState>, |
| 937 | Path(id): Path<String>, |
| 938 | ) -> Result<Json<AutomationRecord>, ApiError> { |
| 939 | let manager = state.automations.lock().await; |
| 940 | let automation = manager.resume_automation(&id).map_err(map_automation_err)?; |
| 941 | Ok(Json(automation)) |
| 942 | } |
| 943 | |
| 944 | async fn list_automation_runs( |
| 945 | State(state): State<RuntimeApiState>, |
| 946 | Path(id): Path<String>, |
| 947 | Query(query): Query<AutomationRunsQuery>, |
| 948 | ) -> Result<Json<Vec<AutomationRunRecord>>, ApiError> { |
| 949 | let manager = state.automations.lock().await; |
| 950 | let runs = manager |
| 951 | .list_runs(&id, query.limit) |
| 952 | .map_err(map_automation_err)?; |
| 953 | Ok(Json(runs)) |
| 954 | } |
| 955 | |
| 956 | async fn get_thread( |
| 957 | State(state): State<RuntimeApiState>, |
| 958 | Path(id): Path<String>, |
| 959 | ) -> Result<Json<ThreadDetail>, ApiError> { |
| 960 | let detail = state |
| 961 | .runtime_threads |
| 962 | .get_thread_detail(&id) |
| 963 | .await |
| 964 | .map_err(map_thread_err)?; |
| 965 | Ok(Json(detail)) |
| 966 | } |
| 967 | |
| 968 | async fn update_thread( |
| 969 | State(state): State<RuntimeApiState>, |
| 970 | Path(id): Path<String>, |
| 971 | Json(req): Json<UpdateThreadRequest>, |
| 972 | ) -> Result<Json<ThreadRecord>, ApiError> { |
| 973 | let thread = state |
| 974 | .runtime_threads |
| 975 | .update_thread(&id, req) |
| 976 | .await |
| 977 | .map_err(map_thread_err)?; |
| 978 | Ok(Json(thread)) |
| 979 | } |
| 980 | |
| 981 | async fn resume_thread( |
| 982 | State(state): State<RuntimeApiState>, |
| 983 | Path(id): Path<String>, |
| 984 | ) -> Result<Json<ThreadRecord>, ApiError> { |
| 985 | let thread = state |
| 986 | .runtime_threads |
| 987 | .resume_thread(&id) |
| 988 | .await |
| 989 | .map_err(map_thread_err)?; |
| 990 | Ok(Json(thread)) |
| 991 | } |
| 992 | |
| 993 | async fn fork_thread( |
| 994 | State(state): State<RuntimeApiState>, |
| 995 | Path(id): Path<String>, |
| 996 | ) -> Result<(StatusCode, Json<ThreadRecord>), ApiError> { |
| 997 | let thread = state |
| 998 | .runtime_threads |
| 999 | .fork_thread(&id) |
| 1000 | .await |
| 1001 | .map_err(map_thread_err)?; |
| 1002 | Ok((StatusCode::CREATED, Json(thread))) |
| 1003 | } |
| 1004 | |
| 1005 | async fn start_thread_turn( |
| 1006 | State(state): State<RuntimeApiState>, |
| 1007 | Path(id): Path<String>, |
| 1008 | Json(req): Json<StartTurnRequest>, |
| 1009 | ) -> Result<(StatusCode, Json<StartTurnResponse>), ApiError> { |
| 1010 | let turn = state |
| 1011 | .runtime_threads |
| 1012 | .start_turn(&id, req) |
| 1013 | .await |
| 1014 | .map_err(map_thread_err)?; |
| 1015 | let thread = state |
| 1016 | .runtime_threads |
| 1017 | .get_thread(&id) |
| 1018 | .await |
| 1019 | .map_err(map_thread_err)?; |
| 1020 | Ok(( |
| 1021 | StatusCode::CREATED, |
| 1022 | Json(StartTurnResponse { thread, turn }), |
| 1023 | )) |
| 1024 | } |
| 1025 | |
| 1026 | async fn steer_thread_turn( |
| 1027 | State(state): State<RuntimeApiState>, |
| 1028 | Path((id, turn_id)): Path<(String, String)>, |
| 1029 | Json(req): Json<SteerTurnRequest>, |
| 1030 | ) -> Result<Json<TurnRecord>, ApiError> { |
| 1031 | let turn = state |
| 1032 | .runtime_threads |
| 1033 | .steer_turn(&id, &turn_id, req) |
| 1034 | .await |
| 1035 | .map_err(map_thread_err)?; |
| 1036 | Ok(Json(turn)) |
| 1037 | } |
| 1038 | |
| 1039 | async fn interrupt_thread_turn( |
| 1040 | State(state): State<RuntimeApiState>, |
| 1041 | Path((id, turn_id)): Path<(String, String)>, |
| 1042 | ) -> Result<Json<TurnRecord>, ApiError> { |
| 1043 | let turn = state |
| 1044 | .runtime_threads |
| 1045 | .interrupt_turn(&id, &turn_id) |
| 1046 | .await |
| 1047 | .map_err(map_thread_err)?; |
| 1048 | Ok(Json(turn)) |
| 1049 | } |
| 1050 | |
| 1051 | async fn compact_thread( |
| 1052 | State(state): State<RuntimeApiState>, |
| 1053 | Path(id): Path<String>, |
| 1054 | Json(req): Json<CompactThreadRequest>, |
| 1055 | ) -> Result<(StatusCode, Json<StartTurnResponse>), ApiError> { |
| 1056 | let turn = state |
| 1057 | .runtime_threads |
| 1058 | .compact_thread(&id, req) |
| 1059 | .await |
| 1060 | .map_err(map_thread_err)?; |
| 1061 | let thread = state |
| 1062 | .runtime_threads |
| 1063 | .get_thread(&id) |
| 1064 | .await |
| 1065 | .map_err(map_thread_err)?; |
| 1066 | Ok(( |
| 1067 | StatusCode::ACCEPTED, |
| 1068 | Json(StartTurnResponse { thread, turn }), |
| 1069 | )) |
| 1070 | } |
| 1071 | |
| 1072 | async fn list_tasks( |
| 1073 | State(state): State<RuntimeApiState>, |
| 1074 | Query(query): Query<TasksQuery>, |
| 1075 | ) -> Result<Json<TasksResponse>, ApiError> { |
| 1076 | let tasks = state.task_manager.list_tasks(query.limit).await; |
| 1077 | let counts = state.task_manager.counts().await; |
| 1078 | Ok(Json(TasksResponse { tasks, counts })) |
| 1079 | } |
| 1080 | |
| 1081 | async fn get_task( |
| 1082 | State(state): State<RuntimeApiState>, |
| 1083 | Path(id): Path<String>, |
| 1084 | ) -> Result<Json<TaskRecord>, ApiError> { |
| 1085 | let task = state |
| 1086 | .task_manager |
| 1087 | .get_task(&id) |
| 1088 | .await |
| 1089 | .map_err(map_task_err)?; |
| 1090 | Ok(Json(task)) |
| 1091 | } |
| 1092 | |
| 1093 | async fn cancel_task( |
| 1094 | State(state): State<RuntimeApiState>, |
| 1095 | Path(id): Path<String>, |
| 1096 | ) -> Result<Json<TaskRecord>, ApiError> { |
| 1097 | let task = state |
| 1098 | .task_manager |
| 1099 | .cancel_task(&id) |
| 1100 | .await |
| 1101 | .map_err(map_task_err)?; |
| 1102 | Ok(Json(task)) |
| 1103 | } |
| 1104 | |
| 1105 | async fn stream_thread_events( |
| 1106 | State(state): State<RuntimeApiState>, |
| 1107 | Path(id): Path<String>, |
| 1108 | Query(query): Query<ThreadEventsQuery>, |
| 1109 | ) -> Result<Sse<impl futures_util::Stream<Item = Result<SseEvent, Infallible>>>, ApiError> { |
| 1110 | let _ = state |
| 1111 | .runtime_threads |
| 1112 | .get_thread(&id) |
| 1113 | .await |
| 1114 | .map_err(map_thread_err)?; |
| 1115 | |
| 1116 | let backlog = state |
| 1117 | .runtime_threads |
| 1118 | .events_since(&id, query.since_seq) |
| 1119 | .map_err(|e| ApiError::internal(e.to_string()))?; |
| 1120 | let mut last_seq = query.since_seq.unwrap_or(0); |
| 1121 | if let Some(last) = backlog.last() { |
| 1122 | last_seq = last.seq; |
| 1123 | } |
| 1124 | |
| 1125 | let mut live = state.runtime_threads.subscribe_events(); |
| 1126 | let thread_id = id.clone(); |
| 1127 | let stream = stream! { |
| 1128 | for event in backlog { |
| 1129 | let event_name = event.event.clone(); |
| 1130 | yield Ok(sse_json(&event_name, runtime_event_payload(event))); |
| 1131 | } |
| 1132 | loop { |
| 1133 | let incoming = live.recv().await; |
| 1134 | let Ok(event) = incoming else { |
| 1135 | break; |
| 1136 | }; |
| 1137 | if event.thread_id != thread_id { |
| 1138 | continue; |
| 1139 | } |
| 1140 | if event.seq <= last_seq { |
| 1141 | continue; |
| 1142 | } |
| 1143 | last_seq = event.seq; |
| 1144 | let event_name = event.event.clone(); |
| 1145 | yield Ok(sse_json(&event_name, runtime_event_payload(event))); |
| 1146 | } |
| 1147 | }; |
| 1148 | |
| 1149 | Ok(Sse::new(stream).keep_alive( |
| 1150 | KeepAlive::new() |
| 1151 | .interval(Duration::from_secs(15)) |
| 1152 | .text("keepalive"), |
| 1153 | )) |
| 1154 | } |
| 1155 | |
| 1156 | async fn stream_turn( |
| 1157 | State(state): State<RuntimeApiState>, |
| 1158 | Json(req): Json<StreamTurnRequest>, |
| 1159 | ) -> Result<Sse<impl futures_util::Stream<Item = Result<SseEvent, Infallible>>>, ApiError> { |
| 1160 | if req.prompt.trim().is_empty() { |
| 1161 | return Err(ApiError::bad_request("prompt is required")); |
| 1162 | } |
| 1163 | |
| 1164 | let model = req.model.clone().unwrap_or_else(|| { |
| 1165 | state |
| 1166 | .config |
| 1167 | .default_text_model |
| 1168 | .clone() |
| 1169 | .unwrap_or_else(|| DEFAULT_TEXT_MODEL.to_string()) |
| 1170 | }); |
| 1171 | let workspace = req |
| 1172 | .workspace |
| 1173 | .clone() |
| 1174 | .unwrap_or_else(|| state.workspace.clone()); |
| 1175 | let mode = req.mode.clone().unwrap_or_else(|| "agent".to_string()); |
| 1176 | let allow_shell = req.allow_shell.unwrap_or(state.config.allow_shell()); |
| 1177 | let trust_mode = req.trust_mode.unwrap_or(false); |
| 1178 | let auto_approve = req.auto_approve.unwrap_or(false); |
| 1179 | let prompt = req.prompt; |
| 1180 | |
| 1181 | let thread = state |
| 1182 | .runtime_threads |
| 1183 | .create_thread(CreateThreadRequest { |
| 1184 | model: Some(model.clone()), |
| 1185 | workspace: Some(workspace.clone()), |
| 1186 | mode: Some(mode.clone()), |
| 1187 | allow_shell: Some(allow_shell), |
| 1188 | trust_mode: Some(trust_mode), |
| 1189 | auto_approve: Some(auto_approve), |
| 1190 | archived: true, |
| 1191 | system_prompt: None, |
| 1192 | task_id: None, |
| 1193 | }) |
| 1194 | .await |
| 1195 | .map_err(|e| ApiError::internal(format!("Failed to create stream thread: {e}")))?; |
| 1196 | |
| 1197 | let turn = state |
| 1198 | .runtime_threads |
| 1199 | .start_turn( |
| 1200 | &thread.id, |
| 1201 | StartTurnRequest { |
| 1202 | prompt, |
| 1203 | input_summary: None, |
| 1204 | model: Some(model.clone()), |
| 1205 | mode: Some(mode.clone()), |
| 1206 | allow_shell: Some(allow_shell), |
| 1207 | trust_mode: Some(trust_mode), |
| 1208 | auto_approve: Some(auto_approve), |
| 1209 | }, |
| 1210 | ) |
| 1211 | .await |
| 1212 | .map_err(|e| ApiError::internal(format!("Failed to start stream turn: {e}")))?; |
| 1213 | |
| 1214 | let backlog = state |
| 1215 | .runtime_threads |
| 1216 | .events_since(&thread.id, None) |
| 1217 | .map_err(|e| ApiError::internal(format!("Failed to load stream backlog: {e}")))?; |
| 1218 | let mut live = state.runtime_threads.subscribe_events(); |
| 1219 | let thread_id = thread.id.clone(); |
| 1220 | let turn_id = turn.id.clone(); |
| 1221 | |
| 1222 | let stream = stream! { |
| 1223 | yield Ok(sse_json("turn.started", json!({ |
| 1224 | "thread_id": thread.id, |
| 1225 | "turn_id": turn.id, |
| 1226 | "model": model, |
| 1227 | "mode": mode, |
| 1228 | "workspace": workspace, |
| 1229 | }))); |
| 1230 | |
| 1231 | for event in backlog { |
| 1232 | if event.thread_id != thread_id || event.turn_id.as_deref() != Some(&turn_id) { |
| 1233 | continue; |
| 1234 | } |
| 1235 | if let Some(mapped) = map_compat_stream_event(&event) { |
| 1236 | yield Ok(mapped); |
| 1237 | } |
| 1238 | if event.event == "turn.completed" { |
| 1239 | yield Ok(sse_json("done", json!({}))); |
| 1240 | return; |
| 1241 | } |
| 1242 | } |
| 1243 | |
| 1244 | loop { |
| 1245 | let incoming = live.recv().await; |
| 1246 | let Ok(event) = incoming else { |
| 1247 | yield Ok(sse_json("error", json!({ "message": "event channel closed" }))); |
| 1248 | break; |
| 1249 | }; |
| 1250 | if event.thread_id != thread_id || event.turn_id.as_deref() != Some(&turn_id) { |
| 1251 | continue; |
| 1252 | } |
| 1253 | if let Some(mapped) = map_compat_stream_event(&event) { |
| 1254 | yield Ok(mapped); |
| 1255 | } |
| 1256 | if event.event == "turn.completed" { |
| 1257 | break; |
| 1258 | } |
| 1259 | } |
| 1260 | |
| 1261 | yield Ok(sse_json("done", json!({}))); |
| 1262 | }; |
| 1263 | |
| 1264 | Ok(Sse::new(stream).keep_alive( |
| 1265 | KeepAlive::new() |
| 1266 | .interval(Duration::from_secs(15)) |
| 1267 | .text("keepalive"), |
| 1268 | )) |
| 1269 | } |
| 1270 | |
| 1271 | fn runtime_event_payload(event: crate::runtime_threads::RuntimeEventRecord) -> serde_json::Value { |
| 1272 | json!({ |
| 1273 | "seq": event.seq, |
| 1274 | "timestamp": event.timestamp, |
| 1275 | "thread_id": event.thread_id, |
| 1276 | "turn_id": event.turn_id, |
| 1277 | "item_id": event.item_id, |
| 1278 | "event": event.event, |
| 1279 | "payload": event.payload, |
| 1280 | }) |
| 1281 | } |
| 1282 | |
| 1283 | fn map_compat_stream_event(event: &crate::runtime_threads::RuntimeEventRecord) -> Option<SseEvent> { |
| 1284 | let payload = &event.payload; |
| 1285 | match event.event.as_str() { |
| 1286 | "item.delta" => { |
| 1287 | let kind = payload |
| 1288 | .get("kind") |
| 1289 | .and_then(|v| v.as_str()) |
| 1290 | .unwrap_or_default(); |
| 1291 | if kind == "agent_message" { |
| 1292 | let content = payload |
| 1293 | .get("delta") |
| 1294 | .and_then(|v| v.as_str()) |
| 1295 | .unwrap_or_default(); |
| 1296 | Some(sse_json("message.delta", json!({ "content": content }))) |
| 1297 | } else if kind == "tool_call" { |
| 1298 | let output = payload |
| 1299 | .get("delta") |
| 1300 | .and_then(|v| v.as_str()) |
| 1301 | .unwrap_or_default(); |
| 1302 | Some(sse_json("tool.progress", json!({ "output": output }))) |
| 1303 | } else { |
| 1304 | None |
| 1305 | } |
| 1306 | } |
| 1307 | "item.started" => { |
| 1308 | let tool = payload.get("tool")?; |
| 1309 | let id = tool.get("id").cloned().unwrap_or(Value::Null); |
| 1310 | let name = tool.get("name").cloned().unwrap_or(Value::Null); |
| 1311 | let input = tool.get("input").cloned().unwrap_or(Value::Null); |
| 1312 | Some(sse_json( |
| 1313 | "tool.started", |
| 1314 | json!({ |
| 1315 | "id": id, |
| 1316 | "name": name, |
| 1317 | "input": input, |
| 1318 | }), |
| 1319 | )) |
| 1320 | } |
| 1321 | "item.completed" | "item.failed" => { |
| 1322 | let item = payload.get("item")?; |
| 1323 | let kind = item |
| 1324 | .get("kind") |
| 1325 | .and_then(|v| v.as_str()) |
| 1326 | .unwrap_or_default(); |
| 1327 | if kind == "tool_call" || kind == "file_change" || kind == "command_execution" { |
| 1328 | let id = item.get("id").cloned().unwrap_or(Value::Null); |
| 1329 | let success = event.event == "item.completed"; |
| 1330 | let output = item.get("detail").cloned().unwrap_or_else(|| { |
| 1331 | Value::String( |
| 1332 | item.get("summary") |
| 1333 | .and_then(|v| v.as_str()) |
| 1334 | .unwrap_or_default() |
| 1335 | .to_string(), |
| 1336 | ) |
| 1337 | }); |
| 1338 | Some(sse_json( |
| 1339 | "tool.completed", |
| 1340 | json!({ |
| 1341 | "id": id, |
| 1342 | "success": success, |
| 1343 | "output": output, |
| 1344 | }), |
| 1345 | )) |
| 1346 | } else if kind == "status" { |
| 1347 | let message = item |
| 1348 | .get("detail") |
| 1349 | .and_then(|v| v.as_str()) |
| 1350 | .or_else(|| item.get("summary").and_then(|v| v.as_str())) |
| 1351 | .unwrap_or_default(); |
| 1352 | Some(sse_json("status", json!({ "message": message }))) |
| 1353 | } else if kind == "error" { |
| 1354 | let message = item |
| 1355 | .get("detail") |
| 1356 | .and_then(|v| v.as_str()) |
| 1357 | .or_else(|| item.get("summary").and_then(|v| v.as_str())) |
| 1358 | .unwrap_or_default(); |
| 1359 | Some(sse_json("error", json!({ "message": message }))) |
| 1360 | } else { |
| 1361 | None |
| 1362 | } |
| 1363 | } |
| 1364 | "approval.required" => Some(sse_json("approval.required", payload.clone())), |
| 1365 | "sandbox.denied" => Some(sse_json("sandbox.denied", payload.clone())), |
| 1366 | "turn.completed" => { |
| 1367 | let usage = payload |
| 1368 | .get("turn") |
| 1369 | .and_then(|turn| turn.get("usage")) |
| 1370 | .cloned() |
| 1371 | .unwrap_or(json!(null)); |
| 1372 | Some(sse_json("turn.completed", json!({ "usage": usage }))) |
| 1373 | } |
| 1374 | _ => None, |
| 1375 | } |
| 1376 | } |
| 1377 | |
| 1378 | fn sse_json(event: &str, payload: serde_json::Value) -> SseEvent { |
| 1379 | let data = serde_json::to_string(&payload).unwrap_or_else(|_| "{}".to_string()); |
| 1380 | SseEvent::default().event(event).data(data) |
| 1381 | } |
| 1382 | |
| 1383 | fn truncate_text(text: &str, max_chars: usize) -> String { |
| 1384 | let char_count = text.chars().count(); |
| 1385 | if char_count <= max_chars { |
| 1386 | return text.to_string(); |
| 1387 | } |
| 1388 | let truncated: String = text.chars().take(max_chars.saturating_sub(3)).collect(); |
| 1389 | format!("{truncated}...") |
| 1390 | } |
| 1391 | |
| 1392 | fn collect_workspace_status(workspace: &std::path::Path) -> WorkspaceStatusResponse { |
| 1393 | let mut status = WorkspaceStatusResponse { |
| 1394 | workspace: workspace.to_path_buf(), |
| 1395 | git_repo: false, |
| 1396 | branch: None, |
| 1397 | staged: 0, |
| 1398 | unstaged: 0, |
| 1399 | untracked: 0, |
| 1400 | ahead: None, |
| 1401 | behind: None, |
| 1402 | }; |
| 1403 | |
| 1404 | let Some(repo_check) = run_git(workspace, &["rev-parse", "--is-inside-work-tree"]) else { |
| 1405 | return status; |
| 1406 | }; |
| 1407 | if repo_check.trim() != "true" { |
| 1408 | return status; |
| 1409 | } |
| 1410 | |
| 1411 | status.git_repo = true; |
| 1412 | status.branch = run_git(workspace, &["rev-parse", "--abbrev-ref", "HEAD"]) |
| 1413 | .map(|s| s.trim().to_string()) |
| 1414 | .filter(|s| !s.is_empty()); |
| 1415 | |
| 1416 | if let Some(porcelain) = run_git(workspace, &["status", "--porcelain=v1"]) { |
| 1417 | for line in porcelain.lines() { |
| 1418 | if line.starts_with("??") { |
| 1419 | status.untracked += 1; |
| 1420 | continue; |
| 1421 | } |
| 1422 | let chars: Vec<char> = line.chars().collect(); |
| 1423 | if chars.len() >= 2 { |
| 1424 | if chars[0] != ' ' { |
| 1425 | status.staged += 1; |
| 1426 | } |
| 1427 | if chars[1] != ' ' { |
| 1428 | status.unstaged += 1; |
| 1429 | } |
| 1430 | } |
| 1431 | } |
| 1432 | } |
| 1433 | |
| 1434 | if let Some(counts) = run_git( |
| 1435 | workspace, |
| 1436 | &["rev-list", "--left-right", "--count", "@{upstream}...HEAD"], |
| 1437 | ) { |
| 1438 | let mut parts = counts.split_whitespace(); |
| 1439 | if let (Some(behind), Some(ahead)) = (parts.next(), parts.next()) { |
| 1440 | status.behind = behind.parse::<u32>().ok(); |
| 1441 | status.ahead = ahead.parse::<u32>().ok(); |
| 1442 | } |
| 1443 | } |
| 1444 | |
| 1445 | status |
| 1446 | } |
| 1447 | |
| 1448 | fn run_git(workspace: &std::path::Path, args: &[&str]) -> Option<String> { |
| 1449 | let output = Command::new("git") |
| 1450 | .args(args) |
| 1451 | .current_dir(workspace) |
| 1452 | .output() |
| 1453 | .ok()?; |
| 1454 | if !output.status.success() { |
| 1455 | return None; |
| 1456 | } |
| 1457 | String::from_utf8(output.stdout).ok() |
| 1458 | } |
| 1459 | |
| 1460 | fn resolve_skills_dir(config: &Config, workspace: &std::path::Path) -> PathBuf { |
| 1461 | let agents_skills = workspace.join(".agents").join("skills"); |
| 1462 | if agents_skills.exists() { |
| 1463 | return agents_skills; |
| 1464 | } |
| 1465 | let local_skills = workspace.join("skills"); |
| 1466 | if local_skills.exists() { |
| 1467 | return local_skills; |
| 1468 | } |
| 1469 | config.skills_dir() |
| 1470 | } |
| 1471 | |
| 1472 | fn load_mcp_config_or_default(path: &std::path::Path) -> Result<McpConfig, ApiError> { |
| 1473 | if !path.exists() { |
| 1474 | return Ok(McpConfig::default()); |
| 1475 | } |
| 1476 | let raw = fs::read_to_string(path).map_err(|e| { |
| 1477 | ApiError::internal(format!("Failed to read MCP config {}: {e}", path.display())) |
| 1478 | })?; |
| 1479 | serde_json::from_str::<McpConfig>(&raw).map_err(|e| { |
| 1480 | ApiError::internal(format!( |
| 1481 | "Failed to parse MCP config {}: {e}", |
| 1482 | path.display() |
| 1483 | )) |
| 1484 | }) |
| 1485 | } |
| 1486 | |
| 1487 | #[derive(Debug, Deserialize)] |
| 1488 | struct UsageQuery { |
| 1489 | /// ISO-8601 lower bound (inclusive). When omitted, no lower bound. |
| 1490 | since: Option<String>, |
| 1491 | /// ISO-8601 upper bound (inclusive). When omitted, no upper bound. |
| 1492 | until: Option<String>, |
| 1493 | /// Bucket key. One of `day` (default), `model`, `provider`, `thread`. |
| 1494 | group_by: Option<String>, |
| 1495 | } |
| 1496 | |
| 1497 | fn parse_iso8601(raw: &str, field: &str) -> Result<chrono::DateTime<Utc>, ApiError> { |
| 1498 | chrono::DateTime::parse_from_rfc3339(raw) |
| 1499 | .map(|dt| dt.with_timezone(&Utc)) |
| 1500 | .map_err(|e| ApiError::bad_request(format!("Invalid {field} (expected RFC 3339): {e}"))) |
| 1501 | } |
| 1502 | |
| 1503 | async fn get_usage( |
| 1504 | State(state): State<RuntimeApiState>, |
| 1505 | Query(query): Query<UsageQuery>, |
| 1506 | ) -> Result<Json<Value>, ApiError> { |
| 1507 | let since = match query.since.as_deref() { |
| 1508 | Some(raw) => Some(parse_iso8601(raw, "since")?), |
| 1509 | None => None, |
| 1510 | }; |
| 1511 | let until = match query.until.as_deref() { |
| 1512 | Some(raw) => Some(parse_iso8601(raw, "until")?), |
| 1513 | None => None, |
| 1514 | }; |
| 1515 | if let (Some(s), Some(u)) = (since, until) |
| 1516 | && s > u |
| 1517 | { |
| 1518 | return Err(ApiError::bad_request("since must be <= until".to_string())); |
| 1519 | } |
| 1520 | let group_by = match query.group_by.as_deref().unwrap_or("day") { |
| 1521 | "day" => UsageGroupBy::Day, |
| 1522 | "model" => UsageGroupBy::Model, |
| 1523 | "provider" => UsageGroupBy::Provider, |
| 1524 | "thread" => UsageGroupBy::Thread, |
| 1525 | other => { |
| 1526 | return Err(ApiError::bad_request(format!( |
| 1527 | "Unsupported group_by '{other}': expected one of day, model, provider, thread" |
| 1528 | ))); |
| 1529 | } |
| 1530 | }; |
| 1531 | |
| 1532 | let aggregation = state |
| 1533 | .runtime_threads |
| 1534 | .aggregate_usage(since, until, group_by) |
| 1535 | .await |
| 1536 | .map_err(|e| ApiError::internal(e.to_string()))?; |
| 1537 | Ok(Json(json!(aggregation))) |
| 1538 | } |
| 1539 | |
| 1540 | /// Built-in dev origins always allowed by the runtime API (whalescale#255). |
| 1541 | const DEFAULT_CORS_ORIGINS: &[&str] = &[ |
| 1542 | "http://localhost:3000", |
| 1543 | "http://127.0.0.1:3000", |
| 1544 | "http://localhost:1420", |
| 1545 | "http://127.0.0.1:1420", |
| 1546 | "tauri://localhost", |
| 1547 | ]; |
| 1548 | |
| 1549 | fn cors_layer(extra_origins: &[String]) -> CorsLayer { |
| 1550 | let mut origins: Vec<HeaderValue> = DEFAULT_CORS_ORIGINS |
| 1551 | .iter() |
| 1552 | .filter_map(|o| HeaderValue::from_str(o).ok()) |
| 1553 | .collect(); |
| 1554 | for raw in extra_origins { |
| 1555 | let trimmed = raw.trim(); |
| 1556 | if trimmed.is_empty() { |
| 1557 | continue; |
| 1558 | } |
| 1559 | match HeaderValue::from_str(trimmed) { |
| 1560 | Ok(value) if !origins.contains(&value) => origins.push(value), |
| 1561 | Ok(_) => {} |
| 1562 | Err(err) => tracing::warn!( |
| 1563 | "Ignoring invalid CORS origin '{trimmed}': {err}; expected scheme://host[:port]" |
| 1564 | ), |
| 1565 | } |
| 1566 | } |
| 1567 | CorsLayer::new() |
| 1568 | .allow_origin(origins) |
| 1569 | .allow_methods([ |
| 1570 | Method::GET, |
| 1571 | Method::POST, |
| 1572 | Method::PATCH, |
| 1573 | Method::DELETE, |
| 1574 | Method::OPTIONS, |
| 1575 | ]) |
| 1576 | .allow_headers(Any) |
| 1577 | } |
| 1578 | |
| 1579 | fn map_task_err(err: anyhow::Error) -> ApiError { |
| 1580 | let message = err.to_string(); |
| 1581 | if message.contains("not found") { |
| 1582 | ApiError::not_found(message) |
| 1583 | } else { |
| 1584 | ApiError::bad_request(message) |
| 1585 | } |
| 1586 | } |
| 1587 | |
| 1588 | fn map_automation_err(err: anyhow::Error) -> ApiError { |
| 1589 | let message = err.to_string(); |
| 1590 | if message.contains("Failed to read automation") |
| 1591 | || message.contains("No such file or directory") |
| 1592 | { |
| 1593 | ApiError::not_found(message) |
| 1594 | } else { |
| 1595 | ApiError::bad_request(message) |
| 1596 | } |
| 1597 | } |
| 1598 | |
| 1599 | fn map_thread_err(err: anyhow::Error) -> ApiError { |
| 1600 | let message = err.to_string(); |
| 1601 | if message.contains("not found") { |
| 1602 | ApiError::not_found(message) |
| 1603 | } else if message.contains("already has an active turn") |
| 1604 | || message.contains("No active turn") |
| 1605 | || message.contains("is not active") |
| 1606 | { |
| 1607 | ApiError { |
| 1608 | status: StatusCode::CONFLICT, |
| 1609 | message, |
| 1610 | } |
| 1611 | } else { |
| 1612 | ApiError::bad_request(message) |
| 1613 | } |
| 1614 | } |
| 1615 | |
| 1616 | #[derive(Debug, Clone)] |
| 1617 | struct ApiError { |
| 1618 | status: StatusCode, |
| 1619 | message: String, |
| 1620 | } |
| 1621 | |
| 1622 | impl ApiError { |
| 1623 | fn bad_request(message: impl Into<String>) -> Self { |
| 1624 | Self { |
| 1625 | status: StatusCode::BAD_REQUEST, |
| 1626 | message: message.into(), |
| 1627 | } |
| 1628 | } |
| 1629 | |
| 1630 | fn not_found(message: impl Into<String>) -> Self { |
| 1631 | Self { |
| 1632 | status: StatusCode::NOT_FOUND, |
| 1633 | message: message.into(), |
| 1634 | } |
| 1635 | } |
| 1636 | |
| 1637 | fn internal(message: impl Into<String>) -> Self { |
| 1638 | Self { |
| 1639 | status: StatusCode::INTERNAL_SERVER_ERROR, |
| 1640 | message: message.into(), |
| 1641 | } |
| 1642 | } |
| 1643 | } |
| 1644 | |
| 1645 | impl IntoResponse for ApiError { |
| 1646 | fn into_response(self) -> Response { |
| 1647 | ( |
| 1648 | self.status, |
| 1649 | Json(json!({ |
| 1650 | "error": { |
| 1651 | "message": self.message, |
| 1652 | "status": self.status.as_u16(), |
| 1653 | } |
| 1654 | })), |
| 1655 | ) |
| 1656 | .into_response() |
| 1657 | } |
| 1658 | } |
| 1659 | |
| 1660 | #[cfg(test)] |
| 1661 | mod tests { |
| 1662 | use super::*; |
| 1663 | use crate::core::events::{Event as EngineEvent, TurnOutcomeStatus}; |
| 1664 | use crate::core::ops::Op; |
| 1665 | use crate::models::Usage; |
| 1666 | use crate::runtime_threads::RuntimeEventRecord; |
| 1667 | use anyhow::{Context, bail}; |
| 1668 | use futures_util::StreamExt; |
| 1669 | use std::fs; |
| 1670 | use std::sync::Arc; |
| 1671 | use tokio::sync::{Mutex, mpsc}; |
| 1672 | use tokio::time::sleep; |
| 1673 | use uuid::Uuid; |
| 1674 | |
| 1675 | struct MockExecutor; |
| 1676 | |
| 1677 | #[async_trait::async_trait] |
| 1678 | impl crate::task_manager::TaskExecutor for MockExecutor { |
| 1679 | async fn execute( |
| 1680 | &self, |
| 1681 | _task: crate::task_manager::ExecutionTask, |
| 1682 | events: mpsc::UnboundedSender<crate::task_manager::TaskExecutionEvent>, |
| 1683 | cancel: tokio_util::sync::CancellationToken, |
| 1684 | ) -> crate::task_manager::TaskExecutionResult { |
| 1685 | let _ = events.send(crate::task_manager::TaskExecutionEvent::Status { |
| 1686 | message: "started".to_string(), |
| 1687 | }); |
| 1688 | sleep(Duration::from_millis(100)).await; |
| 1689 | if cancel.is_cancelled() { |
| 1690 | return crate::task_manager::TaskExecutionResult { |
| 1691 | status: crate::task_manager::TaskStatus::Canceled, |
| 1692 | result_text: None, |
| 1693 | error: None, |
| 1694 | }; |
| 1695 | } |
| 1696 | crate::task_manager::TaskExecutionResult { |
| 1697 | status: crate::task_manager::TaskStatus::Completed, |
| 1698 | result_text: Some("ok".to_string()), |
| 1699 | error: None, |
| 1700 | } |
| 1701 | } |
| 1702 | } |
| 1703 | |
| 1704 | async fn spawn_test_server_with_root( |
| 1705 | root: PathBuf, |
| 1706 | sessions_dir: PathBuf, |
| 1707 | ) -> Result< |
| 1708 | Option<( |
| 1709 | SocketAddr, |
| 1710 | SharedRuntimeThreadManager, |
| 1711 | tokio::task::JoinHandle<()>, |
| 1712 | )>, |
| 1713 | > { |
| 1714 | spawn_test_server_with_root_and_token(root, sessions_dir, None).await |
| 1715 | } |
| 1716 | |
| 1717 | async fn spawn_test_server_with_root_and_token( |
| 1718 | root: PathBuf, |
| 1719 | sessions_dir: PathBuf, |
| 1720 | runtime_token: Option<String>, |
| 1721 | ) -> Result< |
| 1722 | Option<( |
| 1723 | SocketAddr, |
| 1724 | SharedRuntimeThreadManager, |
| 1725 | tokio::task::JoinHandle<()>, |
| 1726 | )>, |
| 1727 | > { |
| 1728 | fs::create_dir_all(&sessions_dir)?; |
| 1729 | let manager = TaskManager::start_with_executor( |
| 1730 | TaskManagerConfig { |
| 1731 | data_dir: root.join("tasks"), |
| 1732 | worker_count: 1, |
| 1733 | default_workspace: PathBuf::from("."), |
| 1734 | default_model: DEFAULT_TEXT_MODEL.to_string(), |
| 1735 | default_mode: "agent".to_string(), |
| 1736 | allow_shell: false, |
| 1737 | trust_mode: false, |
| 1738 | max_subagents: 2, |
| 1739 | }, |
| 1740 | Arc::new(MockExecutor), |
| 1741 | ) |
| 1742 | .await?; |
| 1743 | let mut config = Config::default(); |
| 1744 | config.capacity = Some(crate::config::CapacityConfig { |
| 1745 | enabled: Some(false), |
| 1746 | low_risk_max: None, |
| 1747 | medium_risk_max: None, |
| 1748 | severe_min_slack: None, |
| 1749 | severe_violation_ratio: None, |
| 1750 | refresh_cooldown_turns: None, |
| 1751 | replan_cooldown_turns: None, |
| 1752 | max_replay_per_turn: None, |
| 1753 | min_turns_before_guardrail: None, |
| 1754 | profile_window: None, |
| 1755 | deepseek_v3_2_chat_prior: None, |
| 1756 | deepseek_v3_2_reasoner_prior: None, |
| 1757 | deepseek_v4_pro_prior: None, |
| 1758 | deepseek_v4_flash_prior: None, |
| 1759 | fallback_default_prior: None, |
| 1760 | }); |
| 1761 | let runtime_threads: SharedRuntimeThreadManager = Arc::new(RuntimeThreadManager::open( |
| 1762 | config, |
| 1763 | PathBuf::from("."), |
| 1764 | RuntimeThreadManagerConfig::from_task_data_dir(root.join("runtime")), |
| 1765 | )?); |
| 1766 | runtime_threads.attach_task_manager(manager.clone()); |
| 1767 | let automations = Arc::new(Mutex::new(AutomationManager::open( |
| 1768 | root.join("automations"), |
| 1769 | )?)); |
| 1770 | runtime_threads.attach_automation_manager(automations.clone()); |
| 1771 | |
| 1772 | let state = RuntimeApiState { |
| 1773 | config: Config::default(), |
| 1774 | workspace: PathBuf::from("."), |
| 1775 | task_manager: manager, |
| 1776 | runtime_threads: runtime_threads.clone(), |
| 1777 | cors_origins: Vec::new(), |
| 1778 | sessions_dir, |
| 1779 | mcp_config_path: root.join("mcp.json"), |
| 1780 | automations, |
| 1781 | runtime_token, |
| 1782 | }; |
| 1783 | let app = build_router(state); |
| 1784 | let listener = match TcpListener::bind("127.0.0.1:0").await { |
| 1785 | Ok(listener) => listener, |
| 1786 | Err(err) if err.kind() == std::io::ErrorKind::PermissionDenied => return Ok(None), |
| 1787 | Err(err) => return Err(err.into()), |
| 1788 | }; |
| 1789 | let addr = listener.local_addr()?; |
| 1790 | let handle = tokio::spawn(async move { |
| 1791 | let _ = axum::serve(listener, app).await; |
| 1792 | }); |
| 1793 | Ok(Some((addr, runtime_threads, handle))) |
| 1794 | } |
| 1795 | |
| 1796 | async fn spawn_test_server() -> Result< |
| 1797 | Option<( |
| 1798 | SocketAddr, |
| 1799 | SharedRuntimeThreadManager, |
| 1800 | tokio::task::JoinHandle<()>, |
| 1801 | )>, |
| 1802 | > { |
| 1803 | let root = std::env::temp_dir().join(format!("deepseek-runtime-api-{}", Uuid::new_v4())); |
| 1804 | let sessions_dir = root.join("sessions"); |
| 1805 | spawn_test_server_with_root(root, sessions_dir).await |
| 1806 | } |
| 1807 | |
| 1808 | async fn read_first_sse_frame(resp: reqwest::Response) -> Result<String> { |
| 1809 | let mut stream = resp.bytes_stream(); |
| 1810 | let mut buf = Vec::new(); |
| 1811 | loop { |
| 1812 | let next = tokio::time::timeout(Duration::from_secs(2), stream.next()) |
| 1813 | .await |
| 1814 | .context("timed out waiting for SSE frame")? |
| 1815 | .context("SSE stream ended unexpectedly")??; |
| 1816 | buf.extend_from_slice(&next); |
| 1817 | |
| 1818 | let text = String::from_utf8_lossy(&buf); |
| 1819 | if let Some(idx) = text.find("\n\n").or_else(|| text.find("\r\n\r\n")) { |
| 1820 | return Ok(text[..idx].to_string()); |
| 1821 | } |
| 1822 | |
| 1823 | if buf.len() > 64 * 1024 { |
| 1824 | bail!("SSE frame exceeded 64KB without delimiter"); |
| 1825 | } |
| 1826 | } |
| 1827 | } |
| 1828 | |
| 1829 | fn parse_sse_frame(frame: &str) -> Result<(String, serde_json::Value)> { |
| 1830 | let mut event_name: Option<String> = None; |
| 1831 | let mut data_lines = Vec::new(); |
| 1832 | for line in frame.lines() { |
| 1833 | if let Some(rest) = line.strip_prefix("event:") { |
| 1834 | event_name = Some(rest.trim().to_string()); |
| 1835 | } else if let Some(rest) = line.strip_prefix("data:") { |
| 1836 | data_lines.push(rest.trim_start().to_string()); |
| 1837 | } |
| 1838 | } |
| 1839 | let event_name = event_name.context("missing SSE event field")?; |
| 1840 | let payload = if data_lines.is_empty() { |
| 1841 | json!({}) |
| 1842 | } else { |
| 1843 | serde_json::from_str(&data_lines.join("\n")) |
| 1844 | .with_context(|| format!("invalid SSE data payload: {}", data_lines.join("\n")))? |
| 1845 | }; |
| 1846 | Ok((event_name, payload)) |
| 1847 | } |
| 1848 | |
| 1849 | async fn wait_for_terminal_turn_status( |
| 1850 | client: &reqwest::Client, |
| 1851 | addr: SocketAddr, |
| 1852 | thread_id: &str, |
| 1853 | turn_id: &str, |
| 1854 | timeout: Duration, |
| 1855 | ) -> Result<String> { |
| 1856 | let deadline = tokio::time::Instant::now() + timeout; |
| 1857 | loop { |
| 1858 | let detail: serde_json::Value = client |
| 1859 | .get(format!("http://{addr}/v1/threads/{thread_id}")) |
| 1860 | .send() |
| 1861 | .await? |
| 1862 | .error_for_status()? |
| 1863 | .json() |
| 1864 | .await?; |
| 1865 | let status = detail["turns"] |
| 1866 | .as_array() |
| 1867 | .and_then(|turns| turns.iter().find(|turn| turn["id"] == turn_id)) |
| 1868 | .and_then(|turn| turn.get("status")) |
| 1869 | .and_then(Value::as_str) |
| 1870 | .unwrap_or_default() |
| 1871 | .to_string(); |
| 1872 | if matches!( |
| 1873 | status.as_str(), |
| 1874 | "completed" | "failed" | "interrupted" | "canceled" |
| 1875 | ) { |
| 1876 | return Ok(status); |
| 1877 | } |
| 1878 | if tokio::time::Instant::now() >= deadline { |
| 1879 | bail!("timed out waiting for terminal turn status for {turn_id}"); |
| 1880 | } |
| 1881 | sleep(Duration::from_millis(25)).await; |
| 1882 | } |
| 1883 | } |
| 1884 | |
| 1885 | #[tokio::test] |
| 1886 | async fn health_and_tasks_endpoints_work() -> Result<()> { |
| 1887 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 1888 | return Ok(()); |
| 1889 | }; |
| 1890 | let client = reqwest::Client::new(); |
| 1891 | |
| 1892 | let health: serde_json::Value = client |
| 1893 | .get(format!("http://{addr}/health")) |
| 1894 | .send() |
| 1895 | .await? |
| 1896 | .error_for_status()? |
| 1897 | .json() |
| 1898 | .await?; |
| 1899 | assert_eq!(health["status"], "ok"); |
| 1900 | |
| 1901 | let created: serde_json::Value = client |
| 1902 | .post(format!("http://{addr}/v1/tasks")) |
| 1903 | .json(&json!({ "prompt": "hello task" })) |
| 1904 | .send() |
| 1905 | .await? |
| 1906 | .error_for_status()? |
| 1907 | .json() |
| 1908 | .await?; |
| 1909 | let id = created["id"].as_str().expect("task id").to_string(); |
| 1910 | |
| 1911 | let listed: serde_json::Value = client |
| 1912 | .get(format!("http://{addr}/v1/tasks")) |
| 1913 | .send() |
| 1914 | .await? |
| 1915 | .error_for_status()? |
| 1916 | .json() |
| 1917 | .await?; |
| 1918 | assert!( |
| 1919 | listed["tasks"] |
| 1920 | .as_array() |
| 1921 | .is_some_and(|tasks| !tasks.is_empty()) |
| 1922 | ); |
| 1923 | |
| 1924 | let detail: serde_json::Value = client |
| 1925 | .get(format!("http://{addr}/v1/tasks/{id}")) |
| 1926 | .send() |
| 1927 | .await? |
| 1928 | .error_for_status()? |
| 1929 | .json() |
| 1930 | .await?; |
| 1931 | assert_eq!(detail["id"], id); |
| 1932 | |
| 1933 | let _cancelled: serde_json::Value = client |
| 1934 | .post(format!("http://{addr}/v1/tasks/{id}/cancel")) |
| 1935 | .send() |
| 1936 | .await? |
| 1937 | .error_for_status()? |
| 1938 | .json() |
| 1939 | .await?; |
| 1940 | |
| 1941 | handle.abort(); |
| 1942 | Ok(()) |
| 1943 | } |
| 1944 | |
| 1945 | #[tokio::test] |
| 1946 | async fn runtime_token_guard_protects_v1_routes() -> Result<()> { |
| 1947 | let root = std::env::temp_dir().join(format!("deepseek-runtime-api-{}", Uuid::new_v4())); |
| 1948 | let sessions_dir = root.join("sessions"); |
| 1949 | let token = "local-test-token".to_string(); |
| 1950 | let Some((addr, _runtime_threads, handle)) = |
| 1951 | spawn_test_server_with_root_and_token(root, sessions_dir, Some(token.clone())).await? |
| 1952 | else { |
| 1953 | return Ok(()); |
| 1954 | }; |
| 1955 | let client = reqwest::Client::new(); |
| 1956 | |
| 1957 | let health = client |
| 1958 | .get(format!("http://{addr}/health")) |
| 1959 | .send() |
| 1960 | .await? |
| 1961 | .error_for_status()?; |
| 1962 | assert_eq!(health.status(), StatusCode::OK); |
| 1963 | |
| 1964 | let unauthorized = client |
| 1965 | .get(format!("http://{addr}/v1/threads/summary")) |
| 1966 | .send() |
| 1967 | .await?; |
| 1968 | assert_eq!(unauthorized.status(), StatusCode::UNAUTHORIZED); |
| 1969 | |
| 1970 | let bearer = client |
| 1971 | .get(format!("http://{addr}/v1/threads/summary")) |
| 1972 | .bearer_auth(&token) |
| 1973 | .send() |
| 1974 | .await? |
| 1975 | .error_for_status()?; |
| 1976 | assert_eq!(bearer.status(), StatusCode::OK); |
| 1977 | |
| 1978 | let query_token = client |
| 1979 | .get(format!("http://{addr}/v1/threads/summary?token={token}")) |
| 1980 | .send() |
| 1981 | .await? |
| 1982 | .error_for_status()?; |
| 1983 | assert_eq!(query_token.status(), StatusCode::OK); |
| 1984 | |
| 1985 | handle.abort(); |
| 1986 | Ok(()) |
| 1987 | } |
| 1988 | |
| 1989 | #[tokio::test] |
| 1990 | async fn workspace_and_automation_endpoints_work() -> Result<()> { |
| 1991 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 1992 | return Ok(()); |
| 1993 | }; |
| 1994 | let client = reqwest::Client::new(); |
| 1995 | |
| 1996 | let workspace: serde_json::Value = client |
| 1997 | .get(format!("http://{addr}/v1/workspace/status")) |
| 1998 | .send() |
| 1999 | .await? |
| 2000 | .error_for_status()? |
| 2001 | .json() |
| 2002 | .await?; |
| 2003 | assert!(workspace.get("workspace").is_some()); |
| 2004 | |
| 2005 | let created: serde_json::Value = client |
| 2006 | .post(format!("http://{addr}/v1/automations")) |
| 2007 | .json(&json!({ |
| 2008 | "name": "Smoke automation", |
| 2009 | "prompt": "automation smoke test", |
| 2010 | "rrule": "FREQ=HOURLY;INTERVAL=2", |
| 2011 | "status": "active" |
| 2012 | })) |
| 2013 | .send() |
| 2014 | .await? |
| 2015 | .error_for_status()? |
| 2016 | .json() |
| 2017 | .await?; |
| 2018 | let automation_id = created["id"] |
| 2019 | .as_str() |
| 2020 | .context("missing automation id")? |
| 2021 | .to_string(); |
| 2022 | |
| 2023 | let listed: serde_json::Value = client |
| 2024 | .get(format!("http://{addr}/v1/automations")) |
| 2025 | .send() |
| 2026 | .await? |
| 2027 | .error_for_status()? |
| 2028 | .json() |
| 2029 | .await?; |
| 2030 | assert!( |
| 2031 | listed |
| 2032 | .as_array() |
| 2033 | .is_some_and(|items| items.iter().any(|item| item["id"] == automation_id)) |
| 2034 | ); |
| 2035 | |
| 2036 | let run_now: serde_json::Value = client |
| 2037 | .post(format!("http://{addr}/v1/automations/{automation_id}/run")) |
| 2038 | .send() |
| 2039 | .await? |
| 2040 | .error_for_status()? |
| 2041 | .json() |
| 2042 | .await?; |
| 2043 | assert_eq!(run_now["automation_id"], automation_id); |
| 2044 | |
| 2045 | let paused: serde_json::Value = client |
| 2046 | .post(format!( |
| 2047 | "http://{addr}/v1/automations/{automation_id}/pause" |
| 2048 | )) |
| 2049 | .send() |
| 2050 | .await? |
| 2051 | .error_for_status()? |
| 2052 | .json() |
| 2053 | .await?; |
| 2054 | assert_eq!(paused["status"], "paused"); |
| 2055 | |
| 2056 | let resumed: serde_json::Value = client |
| 2057 | .post(format!( |
| 2058 | "http://{addr}/v1/automations/{automation_id}/resume" |
| 2059 | )) |
| 2060 | .send() |
| 2061 | .await? |
| 2062 | .error_for_status()? |
| 2063 | .json() |
| 2064 | .await?; |
| 2065 | assert_eq!(resumed["status"], "active"); |
| 2066 | |
| 2067 | let updated: serde_json::Value = client |
| 2068 | .patch(format!("http://{addr}/v1/automations/{automation_id}")) |
| 2069 | .json(&json!({ |
| 2070 | "name": "Smoke automation edited", |
| 2071 | "rrule": "FREQ=WEEKLY;BYDAY=MO,WE;BYHOUR=10;BYMINUTE=15" |
| 2072 | })) |
| 2073 | .send() |
| 2074 | .await? |
| 2075 | .error_for_status()? |
| 2076 | .json() |
| 2077 | .await?; |
| 2078 | assert_eq!(updated["name"], "Smoke automation edited"); |
| 2079 | |
| 2080 | let runs: serde_json::Value = client |
| 2081 | .get(format!( |
| 2082 | "http://{addr}/v1/automations/{automation_id}/runs?limit=5" |
| 2083 | )) |
| 2084 | .send() |
| 2085 | .await? |
| 2086 | .error_for_status()? |
| 2087 | .json() |
| 2088 | .await?; |
| 2089 | assert!( |
| 2090 | runs.as_array().is_some_and(|items| !items.is_empty()), |
| 2091 | "expected at least one run entry" |
| 2092 | ); |
| 2093 | |
| 2094 | let _deleted: serde_json::Value = client |
| 2095 | .delete(format!("http://{addr}/v1/automations/{automation_id}")) |
| 2096 | .send() |
| 2097 | .await? |
| 2098 | .error_for_status()? |
| 2099 | .json() |
| 2100 | .await?; |
| 2101 | |
| 2102 | let missing_status = client |
| 2103 | .get(format!("http://{addr}/v1/automations/{automation_id}")) |
| 2104 | .send() |
| 2105 | .await? |
| 2106 | .status(); |
| 2107 | assert_eq!(missing_status, StatusCode::NOT_FOUND); |
| 2108 | |
| 2109 | handle.abort(); |
| 2110 | Ok(()) |
| 2111 | } |
| 2112 | |
| 2113 | #[tokio::test] |
| 2114 | async fn stream_requires_prompt() -> Result<()> { |
| 2115 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 2116 | return Ok(()); |
| 2117 | }; |
| 2118 | let client = reqwest::Client::new(); |
| 2119 | |
| 2120 | let resp = client |
| 2121 | .post(format!("http://{addr}/v1/stream")) |
| 2122 | .json(&json!({ "prompt": "" })) |
| 2123 | .send() |
| 2124 | .await?; |
| 2125 | assert_eq!(resp.status(), StatusCode::BAD_REQUEST); |
| 2126 | handle.abort(); |
| 2127 | Ok(()) |
| 2128 | } |
| 2129 | |
| 2130 | #[tokio::test] |
| 2131 | async fn thread_endpoints_expose_lifecycle_contract() -> Result<()> { |
| 2132 | let Some((addr, runtime_threads, handle)) = spawn_test_server().await? else { |
| 2133 | return Ok(()); |
| 2134 | }; |
| 2135 | let client = reqwest::Client::new(); |
| 2136 | |
| 2137 | let created: serde_json::Value = client |
| 2138 | .post(format!("http://{addr}/v1/threads")) |
| 2139 | .json(&json!({})) |
| 2140 | .send() |
| 2141 | .await? |
| 2142 | .error_for_status()? |
| 2143 | .json() |
| 2144 | .await?; |
| 2145 | let thread_id = created["id"] |
| 2146 | .as_str() |
| 2147 | .context("missing thread id")? |
| 2148 | .to_string(); |
| 2149 | |
| 2150 | let archived: serde_json::Value = client |
| 2151 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 2152 | .json(&json!({ "archived": true })) |
| 2153 | .send() |
| 2154 | .await? |
| 2155 | .error_for_status()? |
| 2156 | .json() |
| 2157 | .await?; |
| 2158 | assert_eq!(archived["id"], thread_id); |
| 2159 | assert_eq!(archived["archived"], true); |
| 2160 | |
| 2161 | let listed: serde_json::Value = client |
| 2162 | .get(format!("http://{addr}/v1/threads")) |
| 2163 | .send() |
| 2164 | .await? |
| 2165 | .error_for_status()? |
| 2166 | .json() |
| 2167 | .await?; |
| 2168 | assert!( |
| 2169 | listed |
| 2170 | .as_array() |
| 2171 | .is_some_and(|threads| threads.iter().all(|t| t["id"] != thread_id)) |
| 2172 | ); |
| 2173 | |
| 2174 | let listed_all: serde_json::Value = client |
| 2175 | .get(format!( |
| 2176 | "http://{addr}/v1/threads/summary?include_archived=true&limit=100" |
| 2177 | )) |
| 2178 | .send() |
| 2179 | .await? |
| 2180 | .error_for_status()? |
| 2181 | .json() |
| 2182 | .await?; |
| 2183 | assert!( |
| 2184 | listed_all |
| 2185 | .as_array() |
| 2186 | .is_some_and(|threads| threads.iter().any(|t| t["id"] == thread_id)) |
| 2187 | ); |
| 2188 | |
| 2189 | let unarchived: serde_json::Value = client |
| 2190 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 2191 | .json(&json!({ "archived": false })) |
| 2192 | .send() |
| 2193 | .await? |
| 2194 | .error_for_status()? |
| 2195 | .json() |
| 2196 | .await?; |
| 2197 | assert_eq!(unarchived["archived"], false); |
| 2198 | |
| 2199 | let invalid_patch = client |
| 2200 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 2201 | .json(&json!({})) |
| 2202 | .send() |
| 2203 | .await?; |
| 2204 | assert_eq!(invalid_patch.status(), StatusCode::BAD_REQUEST); |
| 2205 | |
| 2206 | let missing_patch = client |
| 2207 | .patch(format!("http://{addr}/v1/threads/thr_missing")) |
| 2208 | .json(&json!({ "archived": true })) |
| 2209 | .send() |
| 2210 | .await?; |
| 2211 | assert_eq!(missing_patch.status(), StatusCode::NOT_FOUND); |
| 2212 | |
| 2213 | let detail: serde_json::Value = client |
| 2214 | .get(format!("http://{addr}/v1/threads/{thread_id}")) |
| 2215 | .send() |
| 2216 | .await? |
| 2217 | .error_for_status()? |
| 2218 | .json() |
| 2219 | .await?; |
| 2220 | assert_eq!(detail["thread"]["id"], thread_id); |
| 2221 | |
| 2222 | let resumed: serde_json::Value = client |
| 2223 | .post(format!("http://{addr}/v1/threads/{thread_id}/resume")) |
| 2224 | .send() |
| 2225 | .await? |
| 2226 | .error_for_status()? |
| 2227 | .json() |
| 2228 | .await?; |
| 2229 | assert_eq!(resumed["id"], thread_id); |
| 2230 | |
| 2231 | let forked: serde_json::Value = client |
| 2232 | .post(format!("http://{addr}/v1/threads/{thread_id}/fork")) |
| 2233 | .send() |
| 2234 | .await? |
| 2235 | .error_for_status()? |
| 2236 | .json() |
| 2237 | .await?; |
| 2238 | let forked_id = forked["id"].as_str().context("missing forked id")?; |
| 2239 | assert_ne!(forked_id, thread_id); |
| 2240 | |
| 2241 | // Install a mock engine so the turn completes without calling the real API. |
| 2242 | // The mock handles both SendMessage and CompactContext ops so the |
| 2243 | // compact endpoint tested later also works. |
| 2244 | let harness = crate::core::engine::mock_engine_handle(); |
| 2245 | runtime_threads |
| 2246 | .install_test_engine(&thread_id, harness.handle.clone()) |
| 2247 | .await?; |
| 2248 | let mut rx_op = harness.rx_op; |
| 2249 | let tx_event = harness.tx_event; |
| 2250 | tokio::spawn(async move { |
| 2251 | while let Some(op) = rx_op.recv().await { |
| 2252 | match op { |
| 2253 | Op::SendMessage { .. } => { |
| 2254 | let _ = tx_event |
| 2255 | .send(EngineEvent::TurnStarted { |
| 2256 | turn_id: "mock_lifecycle".to_string(), |
| 2257 | }) |
| 2258 | .await; |
| 2259 | let _ = tx_event |
| 2260 | .send(EngineEvent::MessageStarted { index: 0 }) |
| 2261 | .await; |
| 2262 | let _ = tx_event |
| 2263 | .send(EngineEvent::MessageDelta { |
| 2264 | index: 0, |
| 2265 | content: "mock reply".to_string(), |
| 2266 | }) |
| 2267 | .await; |
| 2268 | let _ = tx_event |
| 2269 | .send(EngineEvent::MessageComplete { index: 0 }) |
| 2270 | .await; |
| 2271 | let _ = tx_event |
| 2272 | .send(EngineEvent::TurnComplete { |
| 2273 | usage: Usage { |
| 2274 | input_tokens: 10, |
| 2275 | output_tokens: 5, |
| 2276 | ..Usage::default() |
| 2277 | }, |
| 2278 | status: TurnOutcomeStatus::Completed, |
| 2279 | error: None, |
| 2280 | }) |
| 2281 | .await; |
| 2282 | } |
| 2283 | Op::CompactContext => { |
| 2284 | let _ = tx_event |
| 2285 | .send(EngineEvent::TurnComplete { |
| 2286 | usage: Usage { |
| 2287 | input_tokens: 0, |
| 2288 | output_tokens: 0, |
| 2289 | ..Usage::default() |
| 2290 | }, |
| 2291 | status: TurnOutcomeStatus::Completed, |
| 2292 | error: None, |
| 2293 | }) |
| 2294 | .await; |
| 2295 | } |
| 2296 | _ => {} |
| 2297 | } |
| 2298 | } |
| 2299 | }); |
| 2300 | |
| 2301 | let turn_start: serde_json::Value = client |
| 2302 | .post(format!("http://{addr}/v1/threads/{thread_id}/turns")) |
| 2303 | .json(&json!({ "prompt": "thread endpoint test" })) |
| 2304 | .send() |
| 2305 | .await? |
| 2306 | .error_for_status()? |
| 2307 | .json() |
| 2308 | .await?; |
| 2309 | let turn_id = turn_start["turn"]["id"] |
| 2310 | .as_str() |
| 2311 | .context("missing turn id")? |
| 2312 | .to_string(); |
| 2313 | |
| 2314 | let _ = wait_for_terminal_turn_status( |
| 2315 | &client, |
| 2316 | addr, |
| 2317 | &thread_id, |
| 2318 | &turn_id, |
| 2319 | Duration::from_secs(2), |
| 2320 | ) |
| 2321 | .await?; |
| 2322 | |
| 2323 | let steer_resp = client |
| 2324 | .post(format!( |
| 2325 | "http://{addr}/v1/threads/{thread_id}/turns/{turn_id}/steer" |
| 2326 | )) |
| 2327 | .json(&json!({ "prompt": "late steer" })) |
| 2328 | .send() |
| 2329 | .await?; |
| 2330 | assert_eq!(steer_resp.status(), StatusCode::CONFLICT); |
| 2331 | |
| 2332 | let interrupt_resp = client |
| 2333 | .post(format!( |
| 2334 | "http://{addr}/v1/threads/{thread_id}/turns/{turn_id}/interrupt" |
| 2335 | )) |
| 2336 | .send() |
| 2337 | .await?; |
| 2338 | assert_eq!(interrupt_resp.status(), StatusCode::CONFLICT); |
| 2339 | |
| 2340 | let compact_start: serde_json::Value = client |
| 2341 | .post(format!("http://{addr}/v1/threads/{thread_id}/compact")) |
| 2342 | .json(&json!({ "reason": "test manual compact" })) |
| 2343 | .send() |
| 2344 | .await? |
| 2345 | .error_for_status()? |
| 2346 | .json() |
| 2347 | .await?; |
| 2348 | assert_eq!(compact_start["thread"]["id"], thread_id); |
| 2349 | |
| 2350 | let events_resp = client |
| 2351 | .get(format!( |
| 2352 | "http://{addr}/v1/threads/{thread_id}/events?since_seq=0" |
| 2353 | )) |
| 2354 | .send() |
| 2355 | .await? |
| 2356 | .error_for_status()?; |
| 2357 | let content_type = events_resp |
| 2358 | .headers() |
| 2359 | .get(reqwest::header::CONTENT_TYPE) |
| 2360 | .and_then(|v| v.to_str().ok()) |
| 2361 | .unwrap_or_default() |
| 2362 | .to_string(); |
| 2363 | assert!(content_type.starts_with("text/event-stream")); |
| 2364 | let chunk_text = read_first_sse_frame(events_resp).await?; |
| 2365 | assert!( |
| 2366 | chunk_text.contains("event:"), |
| 2367 | "expected SSE event chunk, got: {chunk_text}" |
| 2368 | ); |
| 2369 | |
| 2370 | handle.abort(); |
| 2371 | Ok(()) |
| 2372 | } |
| 2373 | |
| 2374 | #[tokio::test] |
| 2375 | async fn events_endpoint_respects_since_seq_cursor() -> Result<()> { |
| 2376 | let Some((addr, runtime_threads, handle)) = spawn_test_server().await? else { |
| 2377 | return Ok(()); |
| 2378 | }; |
| 2379 | let client = reqwest::Client::new(); |
| 2380 | |
| 2381 | let created: serde_json::Value = client |
| 2382 | .post(format!("http://{addr}/v1/threads")) |
| 2383 | .json(&json!({})) |
| 2384 | .send() |
| 2385 | .await? |
| 2386 | .error_for_status()? |
| 2387 | .json() |
| 2388 | .await?; |
| 2389 | let thread_id = created["id"] |
| 2390 | .as_str() |
| 2391 | .context("missing thread id")? |
| 2392 | .to_string(); |
| 2393 | |
| 2394 | // Install a mock engine so the turn completes without calling the real API. |
| 2395 | let harness = crate::core::engine::mock_engine_handle(); |
| 2396 | runtime_threads |
| 2397 | .install_test_engine(&thread_id, harness.handle.clone()) |
| 2398 | .await?; |
| 2399 | let mut rx_op = harness.rx_op; |
| 2400 | let tx_event = harness.tx_event; |
| 2401 | tokio::spawn(async move { |
| 2402 | if !matches!(rx_op.recv().await, Some(Op::SendMessage { .. })) { |
| 2403 | return; |
| 2404 | } |
| 2405 | let _ = tx_event |
| 2406 | .send(EngineEvent::TurnStarted { |
| 2407 | turn_id: "mock_cursor".to_string(), |
| 2408 | }) |
| 2409 | .await; |
| 2410 | let _ = tx_event |
| 2411 | .send(EngineEvent::MessageStarted { index: 0 }) |
| 2412 | .await; |
| 2413 | let _ = tx_event |
| 2414 | .send(EngineEvent::MessageComplete { index: 0 }) |
| 2415 | .await; |
| 2416 | let _ = tx_event |
| 2417 | .send(EngineEvent::TurnComplete { |
| 2418 | usage: Usage { |
| 2419 | input_tokens: 5, |
| 2420 | output_tokens: 3, |
| 2421 | ..Usage::default() |
| 2422 | }, |
| 2423 | status: TurnOutcomeStatus::Completed, |
| 2424 | error: None, |
| 2425 | }) |
| 2426 | .await; |
| 2427 | }); |
| 2428 | |
| 2429 | let started: serde_json::Value = client |
| 2430 | .post(format!("http://{addr}/v1/threads/{thread_id}/turns")) |
| 2431 | .json(&json!({ "prompt": "cursor replay test" })) |
| 2432 | .send() |
| 2433 | .await? |
| 2434 | .error_for_status()? |
| 2435 | .json() |
| 2436 | .await?; |
| 2437 | let turn_id = started["turn"]["id"] |
| 2438 | .as_str() |
| 2439 | .context("missing turn id")? |
| 2440 | .to_string(); |
| 2441 | |
| 2442 | let _ = wait_for_terminal_turn_status( |
| 2443 | &client, |
| 2444 | addr, |
| 2445 | &thread_id, |
| 2446 | &turn_id, |
| 2447 | Duration::from_secs(2), |
| 2448 | ) |
| 2449 | .await?; |
| 2450 | |
| 2451 | let resp_a = client |
| 2452 | .get(format!( |
| 2453 | "http://{addr}/v1/threads/{thread_id}/events?since_seq=0" |
| 2454 | )) |
| 2455 | .send() |
| 2456 | .await? |
| 2457 | .error_for_status()?; |
| 2458 | let frame_a = read_first_sse_frame(resp_a).await?; |
| 2459 | let (_event_a, payload_a) = parse_sse_frame(&frame_a)?; |
| 2460 | let seq_a = payload_a |
| 2461 | .get("seq") |
| 2462 | .and_then(Value::as_u64) |
| 2463 | .context("missing seq in first replay frame")?; |
| 2464 | |
| 2465 | let resp_b = client |
| 2466 | .get(format!( |
| 2467 | "http://{addr}/v1/threads/{thread_id}/events?since_seq={seq_a}" |
| 2468 | )) |
| 2469 | .send() |
| 2470 | .await? |
| 2471 | .error_for_status()?; |
| 2472 | let frame_b = read_first_sse_frame(resp_b).await?; |
| 2473 | let (_event_b, payload_b) = parse_sse_frame(&frame_b)?; |
| 2474 | let seq_b = payload_b |
| 2475 | .get("seq") |
| 2476 | .and_then(Value::as_u64) |
| 2477 | .context("missing seq in second replay frame")?; |
| 2478 | assert!( |
| 2479 | seq_b > seq_a, |
| 2480 | "expected seq after cursor: {seq_b} <= {seq_a}" |
| 2481 | ); |
| 2482 | assert_eq!(payload_b["thread_id"], thread_id); |
| 2483 | |
| 2484 | handle.abort(); |
| 2485 | Ok(()) |
| 2486 | } |
| 2487 | |
| 2488 | #[tokio::test] |
| 2489 | async fn steer_and_interrupt_endpoints_work_on_active_turn() -> Result<()> { |
| 2490 | let Some((addr, runtime_threads, handle)) = spawn_test_server().await? else { |
| 2491 | return Ok(()); |
| 2492 | }; |
| 2493 | let client = reqwest::Client::new(); |
| 2494 | |
| 2495 | let created: serde_json::Value = client |
| 2496 | .post(format!("http://{addr}/v1/threads")) |
| 2497 | .json(&json!({})) |
| 2498 | .send() |
| 2499 | .await? |
| 2500 | .error_for_status()? |
| 2501 | .json() |
| 2502 | .await?; |
| 2503 | let thread_id = created["id"] |
| 2504 | .as_str() |
| 2505 | .context("missing thread id")? |
| 2506 | .to_string(); |
| 2507 | |
| 2508 | let harness = crate::core::engine::mock_engine_handle(); |
| 2509 | runtime_threads |
| 2510 | .install_test_engine(&thread_id, harness.handle.clone()) |
| 2511 | .await?; |
| 2512 | let mut rx_op = harness.rx_op; |
| 2513 | let mut rx_steer = harness.rx_steer; |
| 2514 | let tx_event = harness.tx_event; |
| 2515 | let cancel_token = harness.cancel_token; |
| 2516 | tokio::spawn(async move { |
| 2517 | if !matches!(rx_op.recv().await, Some(Op::SendMessage { .. })) { |
| 2518 | return; |
| 2519 | } |
| 2520 | let _ = tx_event |
| 2521 | .send(EngineEvent::TurnStarted { |
| 2522 | turn_id: "engine_turn_api".to_string(), |
| 2523 | }) |
| 2524 | .await; |
| 2525 | let _ = tx_event |
| 2526 | .send(EngineEvent::MessageStarted { index: 0 }) |
| 2527 | .await; |
| 2528 | if let Some(steer_text) = rx_steer.recv().await { |
| 2529 | let _ = tx_event |
| 2530 | .send(EngineEvent::MessageDelta { |
| 2531 | index: 0, |
| 2532 | content: format!("steer:{steer_text}"), |
| 2533 | }) |
| 2534 | .await; |
| 2535 | } |
| 2536 | cancel_token.cancelled().await; |
| 2537 | sleep(Duration::from_millis(60)).await; |
| 2538 | let _ = tx_event |
| 2539 | .send(EngineEvent::TurnComplete { |
| 2540 | usage: Usage { |
| 2541 | input_tokens: 2, |
| 2542 | output_tokens: 1, |
| 2543 | ..Usage::default() |
| 2544 | }, |
| 2545 | status: TurnOutcomeStatus::Completed, |
| 2546 | error: None, |
| 2547 | }) |
| 2548 | .await; |
| 2549 | }); |
| 2550 | |
| 2551 | let turn_start: serde_json::Value = client |
| 2552 | .post(format!("http://{addr}/v1/threads/{thread_id}/turns")) |
| 2553 | .json(&json!({ "prompt": "active controls" })) |
| 2554 | .send() |
| 2555 | .await? |
| 2556 | .error_for_status()? |
| 2557 | .json() |
| 2558 | .await?; |
| 2559 | let turn_id = turn_start["turn"]["id"] |
| 2560 | .as_str() |
| 2561 | .context("missing turn id")? |
| 2562 | .to_string(); |
| 2563 | |
| 2564 | let steer_resp: serde_json::Value = client |
| 2565 | .post(format!( |
| 2566 | "http://{addr}/v1/threads/{thread_id}/turns/{turn_id}/steer" |
| 2567 | )) |
| 2568 | .json(&json!({ "prompt": "please steer" })) |
| 2569 | .send() |
| 2570 | .await? |
| 2571 | .error_for_status()? |
| 2572 | .json() |
| 2573 | .await?; |
| 2574 | assert_eq!(steer_resp["id"], turn_id); |
| 2575 | assert_eq!(steer_resp["steer_count"], 1); |
| 2576 | |
| 2577 | let interrupt_resp: serde_json::Value = client |
| 2578 | .post(format!( |
| 2579 | "http://{addr}/v1/threads/{thread_id}/turns/{turn_id}/interrupt" |
| 2580 | )) |
| 2581 | .send() |
| 2582 | .await? |
| 2583 | .error_for_status()? |
| 2584 | .json() |
| 2585 | .await?; |
| 2586 | assert_eq!(interrupt_resp["id"], turn_id); |
| 2587 | |
| 2588 | let terminal = wait_for_terminal_turn_status( |
| 2589 | &client, |
| 2590 | addr, |
| 2591 | &thread_id, |
| 2592 | &turn_id, |
| 2593 | Duration::from_secs(3), |
| 2594 | ) |
| 2595 | .await?; |
| 2596 | assert_eq!(terminal, "interrupted"); |
| 2597 | |
| 2598 | let events = runtime_threads.events_since(&thread_id, None)?; |
| 2599 | assert!(events.iter().any(|ev| ev.event == "turn.steered")); |
| 2600 | assert!( |
| 2601 | events |
| 2602 | .iter() |
| 2603 | .any(|ev| ev.event == "turn.interrupt_requested") |
| 2604 | ); |
| 2605 | assert!(events.iter().any(|ev| { |
| 2606 | ev.event == "turn.completed" |
| 2607 | && ev |
| 2608 | .payload |
| 2609 | .get("turn") |
| 2610 | .and_then(|turn| turn.get("status")) |
| 2611 | .and_then(Value::as_str) |
| 2612 | == Some("interrupted") |
| 2613 | })); |
| 2614 | |
| 2615 | handle.abort(); |
| 2616 | Ok(()) |
| 2617 | } |
| 2618 | |
| 2619 | #[tokio::test] |
| 2620 | async fn stream_compat_mapping_handles_expected_runtime_events() -> Result<()> { |
| 2621 | let agent_delta = RuntimeEventRecord { |
| 2622 | schema_version: 1, |
| 2623 | seq: 1, |
| 2624 | timestamp: chrono::Utc::now(), |
| 2625 | thread_id: "thr_test".to_string(), |
| 2626 | turn_id: Some("turn_test".to_string()), |
| 2627 | item_id: Some("item_test".to_string()), |
| 2628 | event: "item.delta".to_string(), |
| 2629 | payload: json!({ |
| 2630 | "kind": "agent_message", |
| 2631 | "delta": "hello", |
| 2632 | }), |
| 2633 | }; |
| 2634 | let mapped = map_compat_stream_event(&agent_delta).context("missing mapped SSE event")?; |
| 2635 | let stream = async_stream::stream! { |
| 2636 | yield Ok::<_, Infallible>(mapped); |
| 2637 | }; |
| 2638 | let body = |
| 2639 | axum::body::to_bytes(Sse::new(stream).into_response().into_body(), usize::MAX).await?; |
| 2640 | let text = String::from_utf8_lossy(&body); |
| 2641 | assert!(text.contains("event: message.delta")); |
| 2642 | assert!(text.contains("\"content\":\"hello\"")); |
| 2643 | |
| 2644 | let tool_start = RuntimeEventRecord { |
| 2645 | schema_version: 1, |
| 2646 | seq: 2, |
| 2647 | timestamp: chrono::Utc::now(), |
| 2648 | thread_id: "thr_test".to_string(), |
| 2649 | turn_id: Some("turn_test".to_string()), |
| 2650 | item_id: Some("item_tool".to_string()), |
| 2651 | event: "item.started".to_string(), |
| 2652 | payload: json!({ |
| 2653 | "tool": { "id": "tool_1", "name": "exec_shell", "input": { "cmd": "pwd" } } |
| 2654 | }), |
| 2655 | }; |
| 2656 | let mapped = map_compat_stream_event(&tool_start).context("missing tool.started event")?; |
| 2657 | let stream = async_stream::stream! { |
| 2658 | yield Ok::<_, Infallible>(mapped); |
| 2659 | }; |
| 2660 | let body = |
| 2661 | axum::body::to_bytes(Sse::new(stream).into_response().into_body(), usize::MAX).await?; |
| 2662 | let text = String::from_utf8_lossy(&body); |
| 2663 | assert!(text.contains("event: tool.started")); |
| 2664 | |
| 2665 | let tool_done = RuntimeEventRecord { |
| 2666 | schema_version: 1, |
| 2667 | seq: 3, |
| 2668 | timestamp: chrono::Utc::now(), |
| 2669 | thread_id: "thr_test".to_string(), |
| 2670 | turn_id: Some("turn_test".to_string()), |
| 2671 | item_id: Some("item_tool".to_string()), |
| 2672 | event: "item.completed".to_string(), |
| 2673 | payload: json!({ |
| 2674 | "item": { |
| 2675 | "id": "item_tool", |
| 2676 | "kind": "tool_call", |
| 2677 | "summary": "ok", |
| 2678 | "detail": "done" |
| 2679 | } |
| 2680 | }), |
| 2681 | }; |
| 2682 | let mapped = map_compat_stream_event(&tool_done).context("missing tool.completed event")?; |
| 2683 | let stream = async_stream::stream! { |
| 2684 | yield Ok::<_, Infallible>(mapped); |
| 2685 | }; |
| 2686 | let body = |
| 2687 | axum::body::to_bytes(Sse::new(stream).into_response().into_body(), usize::MAX).await?; |
| 2688 | let text = String::from_utf8_lossy(&body); |
| 2689 | assert!(text.contains("event: tool.completed")); |
| 2690 | assert!(text.contains("\"success\":true")); |
| 2691 | |
| 2692 | let unknown = RuntimeEventRecord { |
| 2693 | schema_version: 1, |
| 2694 | seq: 4, |
| 2695 | timestamp: chrono::Utc::now(), |
| 2696 | thread_id: "thr_test".to_string(), |
| 2697 | turn_id: Some("turn_test".to_string()), |
| 2698 | item_id: None, |
| 2699 | event: "item.delta".to_string(), |
| 2700 | payload: json!({ |
| 2701 | "kind": "context_compaction", |
| 2702 | "delta": "ignored", |
| 2703 | }), |
| 2704 | }; |
| 2705 | assert!(map_compat_stream_event(&unknown).is_none()); |
| 2706 | Ok(()) |
| 2707 | } |
| 2708 | |
| 2709 | #[tokio::test] |
| 2710 | async fn stream_endpoint_remains_backward_compatible() -> Result<()> { |
| 2711 | let Some((addr, runtime_threads, handle)) = spawn_test_server().await? else { |
| 2712 | return Ok(()); |
| 2713 | }; |
| 2714 | let client = reqwest::Client::new(); |
| 2715 | |
| 2716 | // Create a thread and install a mock engine so /v1/stream doesn't call the real API. |
| 2717 | let created: serde_json::Value = client |
| 2718 | .post(format!("http://{addr}/v1/threads")) |
| 2719 | .json(&json!({})) |
| 2720 | .send() |
| 2721 | .await? |
| 2722 | .error_for_status()? |
| 2723 | .json() |
| 2724 | .await?; |
| 2725 | let thread_id = created["id"] |
| 2726 | .as_str() |
| 2727 | .context("missing thread id")? |
| 2728 | .to_string(); |
| 2729 | |
| 2730 | let harness = crate::core::engine::mock_engine_handle(); |
| 2731 | runtime_threads |
| 2732 | .install_test_engine(&thread_id, harness.handle.clone()) |
| 2733 | .await?; |
| 2734 | let mut rx_op = harness.rx_op; |
| 2735 | let tx_event = harness.tx_event; |
| 2736 | tokio::spawn(async move { |
| 2737 | if !matches!(rx_op.recv().await, Some(Op::SendMessage { .. })) { |
| 2738 | return; |
| 2739 | } |
| 2740 | let _ = tx_event |
| 2741 | .send(EngineEvent::TurnStarted { |
| 2742 | turn_id: "mock_stream".to_string(), |
| 2743 | }) |
| 2744 | .await; |
| 2745 | let _ = tx_event |
| 2746 | .send(EngineEvent::MessageStarted { index: 0 }) |
| 2747 | .await; |
| 2748 | let _ = tx_event |
| 2749 | .send(EngineEvent::MessageDelta { |
| 2750 | index: 0, |
| 2751 | content: "streamed".to_string(), |
| 2752 | }) |
| 2753 | .await; |
| 2754 | let _ = tx_event |
| 2755 | .send(EngineEvent::MessageComplete { index: 0 }) |
| 2756 | .await; |
| 2757 | let _ = tx_event |
| 2758 | .send(EngineEvent::TurnComplete { |
| 2759 | usage: Usage { |
| 2760 | input_tokens: 4, |
| 2761 | output_tokens: 2, |
| 2762 | ..Usage::default() |
| 2763 | }, |
| 2764 | status: TurnOutcomeStatus::Completed, |
| 2765 | error: None, |
| 2766 | }) |
| 2767 | .await; |
| 2768 | }); |
| 2769 | |
| 2770 | // Start the turn and consume events via the SSE endpoint. |
| 2771 | let turn_start: serde_json::Value = client |
| 2772 | .post(format!("http://{addr}/v1/threads/{thread_id}/turns")) |
| 2773 | .json(&json!({ "prompt": "compatibility stream" })) |
| 2774 | .send() |
| 2775 | .await? |
| 2776 | .error_for_status()? |
| 2777 | .json() |
| 2778 | .await?; |
| 2779 | let turn_id = turn_start["turn"]["id"] |
| 2780 | .as_str() |
| 2781 | .context("missing turn id")? |
| 2782 | .to_string(); |
| 2783 | |
| 2784 | let _ = wait_for_terminal_turn_status( |
| 2785 | &client, |
| 2786 | addr, |
| 2787 | &thread_id, |
| 2788 | &turn_id, |
| 2789 | Duration::from_secs(2), |
| 2790 | ) |
| 2791 | .await?; |
| 2792 | |
| 2793 | // Verify that the persisted events include the expected turn lifecycle events. |
| 2794 | let events = runtime_threads.events_since(&thread_id, None)?; |
| 2795 | assert!( |
| 2796 | events.iter().any(|ev| ev.event == "turn.started"), |
| 2797 | "expected turn.started event" |
| 2798 | ); |
| 2799 | assert!( |
| 2800 | events.iter().any(|ev| ev.event == "turn.completed"), |
| 2801 | "expected turn.completed event" |
| 2802 | ); |
| 2803 | |
| 2804 | // Verify the SSE endpoint returns event-stream content type. |
| 2805 | let events_resp = client |
| 2806 | .get(format!( |
| 2807 | "http://{addr}/v1/threads/{thread_id}/events?since_seq=0" |
| 2808 | )) |
| 2809 | .send() |
| 2810 | .await? |
| 2811 | .error_for_status()?; |
| 2812 | let content_type = events_resp |
| 2813 | .headers() |
| 2814 | .get(reqwest::header::CONTENT_TYPE) |
| 2815 | .and_then(|v| v.to_str().ok()) |
| 2816 | .unwrap_or_default() |
| 2817 | .to_string(); |
| 2818 | assert!(content_type.starts_with("text/event-stream")); |
| 2819 | |
| 2820 | handle.abort(); |
| 2821 | Ok(()) |
| 2822 | } |
| 2823 | |
| 2824 | #[tokio::test] |
| 2825 | async fn session_get_returns_404_for_missing_id() -> Result<()> { |
| 2826 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 2827 | return Ok(()); |
| 2828 | }; |
| 2829 | let client = reqwest::Client::new(); |
| 2830 | |
| 2831 | let resp = client |
| 2832 | .get(format!("http://{addr}/v1/sessions/nonexistent_id")) |
| 2833 | .send() |
| 2834 | .await?; |
| 2835 | assert_eq!(resp.status(), StatusCode::NOT_FOUND); |
| 2836 | |
| 2837 | handle.abort(); |
| 2838 | Ok(()) |
| 2839 | } |
| 2840 | |
| 2841 | #[tokio::test] |
| 2842 | async fn session_endpoints_reject_invalid_id() -> Result<()> { |
| 2843 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 2844 | return Ok(()); |
| 2845 | }; |
| 2846 | let client = reqwest::Client::new(); |
| 2847 | |
| 2848 | let get_resp = client |
| 2849 | .get(format!("http://{addr}/v1/sessions/invalid%20id")) |
| 2850 | .send() |
| 2851 | .await?; |
| 2852 | assert_eq!(get_resp.status(), StatusCode::BAD_REQUEST); |
| 2853 | |
| 2854 | let resume_resp = client |
| 2855 | .post(format!( |
| 2856 | "http://{addr}/v1/sessions/invalid%20id/resume-thread" |
| 2857 | )) |
| 2858 | .json(&json!({})) |
| 2859 | .send() |
| 2860 | .await?; |
| 2861 | assert_eq!(resume_resp.status(), StatusCode::BAD_REQUEST); |
| 2862 | |
| 2863 | let delete_resp = client |
| 2864 | .delete(format!("http://{addr}/v1/sessions/invalid%20id")) |
| 2865 | .send() |
| 2866 | .await?; |
| 2867 | assert_eq!(delete_resp.status(), StatusCode::BAD_REQUEST); |
| 2868 | |
| 2869 | handle.abort(); |
| 2870 | Ok(()) |
| 2871 | } |
| 2872 | |
| 2873 | #[tokio::test] |
| 2874 | async fn session_resume_thread_returns_404_for_missing_session() -> Result<()> { |
| 2875 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 2876 | return Ok(()); |
| 2877 | }; |
| 2878 | let client = reqwest::Client::new(); |
| 2879 | |
| 2880 | let resp = client |
| 2881 | .post(format!( |
| 2882 | "http://{addr}/v1/sessions/nonexistent_session/resume-thread" |
| 2883 | )) |
| 2884 | .json(&json!({})) |
| 2885 | .send() |
| 2886 | .await?; |
| 2887 | assert_eq!(resp.status(), StatusCode::NOT_FOUND); |
| 2888 | |
| 2889 | handle.abort(); |
| 2890 | Ok(()) |
| 2891 | } |
| 2892 | |
| 2893 | #[tokio::test] |
| 2894 | async fn session_resume_thread_creates_thread_from_saved_session() -> Result<()> { |
| 2895 | let root = std::env::temp_dir().join(format!("deepseek-session-resume-{}", Uuid::new_v4())); |
| 2896 | let sessions_dir = root.join("sessions"); |
| 2897 | fs::create_dir_all(&sessions_dir)?; |
| 2898 | let session_id = "sess_test_resume"; |
| 2899 | let session = json!({ |
| 2900 | "schema_version": 1, |
| 2901 | "metadata": { |
| 2902 | "id": session_id, |
| 2903 | "title": "Test resume session", |
| 2904 | "created_at": "2025-01-01T00:00:00Z", |
| 2905 | "updated_at": "2025-01-01T00:10:00Z", |
| 2906 | "message_count": 2, |
| 2907 | "total_tokens": 100, |
| 2908 | "model": "deepseek-v4-pro", |
| 2909 | "workspace": "/tmp/test", |
| 2910 | "mode": "agent" |
| 2911 | }, |
| 2912 | "messages": [ |
| 2913 | { |
| 2914 | "role": "user", |
| 2915 | "content": [{ "type": "text", "text": "Hello, world!" }] |
| 2916 | }, |
| 2917 | { |
| 2918 | "role": "assistant", |
| 2919 | "content": [{ "type": "text", "text": "Hello! How can I help you?" }] |
| 2920 | } |
| 2921 | ], |
| 2922 | "system_prompt": null |
| 2923 | }); |
| 2924 | fs::write( |
| 2925 | sessions_dir.join(format!("{session_id}.json")), |
| 2926 | serde_json::to_string_pretty(&session)?, |
| 2927 | )?; |
| 2928 | |
| 2929 | let Some((addr, _runtime_threads, handle)) = |
| 2930 | spawn_test_server_with_root(root.clone(), sessions_dir.clone()).await? |
| 2931 | else { |
| 2932 | return Ok(()); |
| 2933 | }; |
| 2934 | let client = reqwest::Client::new(); |
| 2935 | |
| 2936 | let resp = client |
| 2937 | .post(format!( |
| 2938 | "http://{addr}/v1/sessions/{session_id}/resume-thread" |
| 2939 | )) |
| 2940 | .json(&json!({ "model": "deepseek-v4-pro" })) |
| 2941 | .send() |
| 2942 | .await?; |
| 2943 | assert_eq!(resp.status(), StatusCode::CREATED); |
| 2944 | let resumed: serde_json::Value = resp.json().await?; |
| 2945 | assert_eq!(resumed["session_id"], session_id); |
| 2946 | assert_eq!(resumed["message_count"], 2); |
| 2947 | |
| 2948 | let thread_id = resumed["thread_id"] |
| 2949 | .as_str() |
| 2950 | .context("missing resumed thread id")?; |
| 2951 | let detail: serde_json::Value = client |
| 2952 | .get(format!("http://{addr}/v1/threads/{thread_id}")) |
| 2953 | .send() |
| 2954 | .await? |
| 2955 | .error_for_status()? |
| 2956 | .json() |
| 2957 | .await?; |
| 2958 | assert_eq!(detail["thread"]["id"], thread_id); |
| 2959 | assert_eq!(detail["turns"].as_array().map_or(0, Vec::len), 1); |
| 2960 | assert_eq!(detail["items"].as_array().map_or(0, Vec::len), 2); |
| 2961 | |
| 2962 | handle.abort(); |
| 2963 | Ok(()) |
| 2964 | } |
| 2965 | |
| 2966 | #[tokio::test] |
| 2967 | async fn session_delete_returns_404_for_missing_id() -> Result<()> { |
| 2968 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 2969 | return Ok(()); |
| 2970 | }; |
| 2971 | let client = reqwest::Client::new(); |
| 2972 | let resp = client |
| 2973 | .delete(format!("http://{addr}/v1/sessions/nonexistent-id")) |
| 2974 | .send() |
| 2975 | .await?; |
| 2976 | assert_eq!(resp.status(), StatusCode::NOT_FOUND); |
| 2977 | handle.abort(); |
| 2978 | Ok(()) |
| 2979 | } |
| 2980 | |
| 2981 | /// #561 / whalescale#255 — extra CORS origins from `RuntimeApiOptions` |
| 2982 | /// are added on top of the built-in defaults and propagate through to the |
| 2983 | /// `Access-Control-Allow-Origin` response header for preflight requests. |
| 2984 | /// Built-in defaults must keep working unchanged. |
| 2985 | #[tokio::test] |
| 2986 | async fn cors_layer_appends_extra_origins_and_keeps_defaults() -> Result<()> { |
| 2987 | // The cors_layer fn is the layer factory — exercise it through a |
| 2988 | // Router with a single trivial route so we can issue OPTIONS preflights |
| 2989 | // and observe the response headers. |
| 2990 | let extra = vec!["http://localhost:5173".to_string()]; |
| 2991 | let layer = cors_layer(&extra); |
| 2992 | let router: Router = Router::new() |
| 2993 | .route("/probe", get(|| async { "ok" })) |
| 2994 | .layer(layer); |
| 2995 | |
| 2996 | let listener = match TcpListener::bind("127.0.0.1:0").await { |
| 2997 | Ok(listener) => listener, |
| 2998 | Err(err) if err.kind() == std::io::ErrorKind::PermissionDenied => return Ok(()), |
| 2999 | Err(err) => return Err(err.into()), |
| 3000 | }; |
| 3001 | let addr = listener.local_addr()?; |
| 3002 | let handle = tokio::spawn(async move { |
| 3003 | let _ = axum::serve(listener, router).await; |
| 3004 | }); |
| 3005 | |
| 3006 | let client = reqwest::Client::new(); |
| 3007 | |
| 3008 | // The user-supplied origin is allowed. |
| 3009 | let resp = client |
| 3010 | .request(reqwest::Method::OPTIONS, format!("http://{addr}/probe")) |
| 3011 | .header("Origin", "http://localhost:5173") |
| 3012 | .header("Access-Control-Request-Method", "GET") |
| 3013 | .send() |
| 3014 | .await?; |
| 3015 | assert_eq!( |
| 3016 | resp.headers() |
| 3017 | .get("access-control-allow-origin") |
| 3018 | .and_then(|v| v.to_str().ok()), |
| 3019 | Some("http://localhost:5173") |
| 3020 | ); |
| 3021 | |
| 3022 | // A built-in default origin still works. |
| 3023 | let resp = client |
| 3024 | .request(reqwest::Method::OPTIONS, format!("http://{addr}/probe")) |
| 3025 | .header("Origin", "http://localhost:1420") |
| 3026 | .header("Access-Control-Request-Method", "GET") |
| 3027 | .send() |
| 3028 | .await?; |
| 3029 | assert_eq!( |
| 3030 | resp.headers() |
| 3031 | .get("access-control-allow-origin") |
| 3032 | .and_then(|v| v.to_str().ok()), |
| 3033 | Some("http://localhost:1420") |
| 3034 | ); |
| 3035 | |
| 3036 | // An origin that's neither configured nor a default is rejected |
| 3037 | // (CorsLayer omits the Allow-Origin header on mismatch). |
| 3038 | let resp = client |
| 3039 | .request(reqwest::Method::OPTIONS, format!("http://{addr}/probe")) |
| 3040 | .header("Origin", "http://malicious.example") |
| 3041 | .header("Access-Control-Request-Method", "GET") |
| 3042 | .send() |
| 3043 | .await?; |
| 3044 | assert!( |
| 3045 | resp.headers().get("access-control-allow-origin").is_none(), |
| 3046 | "non-allowed origin must not be echoed back" |
| 3047 | ); |
| 3048 | |
| 3049 | handle.abort(); |
| 3050 | Ok(()) |
| 3051 | } |
| 3052 | |
| 3053 | /// #561 — invalid origins (non-ASCII, etc.) are skipped without aborting |
| 3054 | /// the layer build. |
| 3055 | #[test] |
| 3056 | fn cors_layer_skips_invalid_origins() { |
| 3057 | let extras = vec![ |
| 3058 | "http://valid.example".to_string(), |
| 3059 | // Embedded NUL char makes `HeaderValue::from_str` fail. |
| 3060 | "http://invalid.example\0".to_string(), |
| 3061 | " ".to_string(), // whitespace-only is dropped |
| 3062 | ]; |
| 3063 | // Should not panic. |
| 3064 | let _ = cors_layer(&extras); |
| 3065 | } |
| 3066 | |
| 3067 | /// #562 / whalescale#256 — `PATCH /v1/threads/{id}` accepts the new |
| 3068 | /// fields (allow_shell, trust_mode, auto_approve, model, mode, title, |
| 3069 | /// system_prompt). Each is independently optional; an empty string clears |
| 3070 | /// `title` / `system_prompt` back to None. |
| 3071 | #[tokio::test] |
| 3072 | async fn patch_thread_accepts_extended_field_set() -> Result<()> { |
| 3073 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 3074 | return Ok(()); |
| 3075 | }; |
| 3076 | let client = reqwest::Client::new(); |
| 3077 | |
| 3078 | let created: serde_json::Value = client |
| 3079 | .post(format!("http://{addr}/v1/threads")) |
| 3080 | .json(&json!({ |
| 3081 | "model": "deepseek-v4-flash", |
| 3082 | "mode": "agent" |
| 3083 | })) |
| 3084 | .send() |
| 3085 | .await? |
| 3086 | .error_for_status()? |
| 3087 | .json() |
| 3088 | .await?; |
| 3089 | let thread_id = created["id"] |
| 3090 | .as_str() |
| 3091 | .context("missing thread id")? |
| 3092 | .to_string(); |
| 3093 | |
| 3094 | // Patch every new field at once. |
| 3095 | let patched: serde_json::Value = client |
| 3096 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 3097 | .json(&json!({ |
| 3098 | "allow_shell": true, |
| 3099 | "trust_mode": true, |
| 3100 | "auto_approve": true, |
| 3101 | "model": "deepseek-v4-pro", |
| 3102 | "mode": "yolo", |
| 3103 | "title": "Whalescale UI test thread", |
| 3104 | "system_prompt": "You are a useful assistant." |
| 3105 | })) |
| 3106 | .send() |
| 3107 | .await? |
| 3108 | .error_for_status()? |
| 3109 | .json() |
| 3110 | .await?; |
| 3111 | |
| 3112 | assert_eq!(patched["allow_shell"], true); |
| 3113 | assert_eq!(patched["trust_mode"], true); |
| 3114 | assert_eq!(patched["auto_approve"], true); |
| 3115 | assert_eq!(patched["model"], "deepseek-v4-pro"); |
| 3116 | assert_eq!(patched["mode"], "yolo"); |
| 3117 | assert_eq!(patched["title"], "Whalescale UI test thread"); |
| 3118 | assert_eq!(patched["system_prompt"], "You are a useful assistant."); |
| 3119 | |
| 3120 | // Empty string clears title back to None. |
| 3121 | let cleared: serde_json::Value = client |
| 3122 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 3123 | .json(&json!({ "title": "" })) |
| 3124 | .send() |
| 3125 | .await? |
| 3126 | .error_for_status()? |
| 3127 | .json() |
| 3128 | .await?; |
| 3129 | assert!( |
| 3130 | cleared["title"].is_null() || !cleared.as_object().unwrap().contains_key("title"), |
| 3131 | "empty title must serialize as None: {cleared:?}" |
| 3132 | ); |
| 3133 | |
| 3134 | // Empty patch (no fields) is still rejected. |
| 3135 | let empty = client |
| 3136 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 3137 | .json(&json!({})) |
| 3138 | .send() |
| 3139 | .await?; |
| 3140 | assert_eq!(empty.status(), StatusCode::BAD_REQUEST); |
| 3141 | |
| 3142 | // Empty model is rejected (validation). |
| 3143 | let bad_model = client |
| 3144 | .patch(format!("http://{addr}/v1/threads/{thread_id}")) |
| 3145 | .json(&json!({ "model": " " })) |
| 3146 | .send() |
| 3147 | .await?; |
| 3148 | assert_eq!(bad_model.status(), StatusCode::BAD_REQUEST); |
| 3149 | |
| 3150 | handle.abort(); |
| 3151 | Ok(()) |
| 3152 | } |
| 3153 | |
| 3154 | /// #563 / whalescale#260 — `archived_only=true` returns archived-only |
| 3155 | /// (no active threads), distinct from `include_archived=true` which |
| 3156 | /// returns both. |
| 3157 | #[tokio::test] |
| 3158 | async fn list_threads_archived_only_filter_matches_only_archived() -> Result<()> { |
| 3159 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 3160 | return Ok(()); |
| 3161 | }; |
| 3162 | let client = reqwest::Client::new(); |
| 3163 | |
| 3164 | // Two threads — keep one active, archive the other. |
| 3165 | let active: serde_json::Value = client |
| 3166 | .post(format!("http://{addr}/v1/threads")) |
| 3167 | .json(&json!({})) |
| 3168 | .send() |
| 3169 | .await? |
| 3170 | .error_for_status()? |
| 3171 | .json() |
| 3172 | .await?; |
| 3173 | let active_id = active["id"].as_str().unwrap().to_string(); |
| 3174 | |
| 3175 | let archived: serde_json::Value = client |
| 3176 | .post(format!("http://{addr}/v1/threads")) |
| 3177 | .json(&json!({})) |
| 3178 | .send() |
| 3179 | .await? |
| 3180 | .error_for_status()? |
| 3181 | .json() |
| 3182 | .await?; |
| 3183 | let archived_id = archived["id"].as_str().unwrap().to_string(); |
| 3184 | |
| 3185 | client |
| 3186 | .patch(format!("http://{addr}/v1/threads/{archived_id}")) |
| 3187 | .json(&json!({ "archived": true })) |
| 3188 | .send() |
| 3189 | .await? |
| 3190 | .error_for_status()?; |
| 3191 | |
| 3192 | // Default (active only) → only the unarchived one. |
| 3193 | let active_list: serde_json::Value = client |
| 3194 | .get(format!("http://{addr}/v1/threads")) |
| 3195 | .send() |
| 3196 | .await? |
| 3197 | .error_for_status()? |
| 3198 | .json() |
| 3199 | .await?; |
| 3200 | let ids: Vec<&str> = active_list |
| 3201 | .as_array() |
| 3202 | .unwrap() |
| 3203 | .iter() |
| 3204 | .filter_map(|t| t["id"].as_str()) |
| 3205 | .collect(); |
| 3206 | assert!(ids.contains(&active_id.as_str())); |
| 3207 | assert!(!ids.contains(&archived_id.as_str())); |
| 3208 | |
| 3209 | // archived_only=true → only the archived one. |
| 3210 | let archived_list: serde_json::Value = client |
| 3211 | .get(format!("http://{addr}/v1/threads?archived_only=true")) |
| 3212 | .send() |
| 3213 | .await? |
| 3214 | .error_for_status()? |
| 3215 | .json() |
| 3216 | .await?; |
| 3217 | let ids: Vec<&str> = archived_list |
| 3218 | .as_array() |
| 3219 | .unwrap() |
| 3220 | .iter() |
| 3221 | .filter_map(|t| t["id"].as_str()) |
| 3222 | .collect(); |
| 3223 | assert_eq!(ids, vec![archived_id.as_str()]); |
| 3224 | |
| 3225 | // archived_only=true takes precedence over include_archived=true. |
| 3226 | let archived_list: serde_json::Value = client |
| 3227 | .get(format!( |
| 3228 | "http://{addr}/v1/threads?include_archived=true&archived_only=true" |
| 3229 | )) |
| 3230 | .send() |
| 3231 | .await? |
| 3232 | .error_for_status()? |
| 3233 | .json() |
| 3234 | .await?; |
| 3235 | let ids: Vec<&str> = archived_list |
| 3236 | .as_array() |
| 3237 | .unwrap() |
| 3238 | .iter() |
| 3239 | .filter_map(|t| t["id"].as_str()) |
| 3240 | .collect(); |
| 3241 | assert_eq!(ids, vec![archived_id.as_str()]); |
| 3242 | |
| 3243 | // Same filter works on the summary endpoint. |
| 3244 | let summary: serde_json::Value = client |
| 3245 | .get(format!( |
| 3246 | "http://{addr}/v1/threads/summary?archived_only=true&limit=10" |
| 3247 | )) |
| 3248 | .send() |
| 3249 | .await? |
| 3250 | .error_for_status()? |
| 3251 | .json() |
| 3252 | .await?; |
| 3253 | let summary_ids: Vec<&str> = summary |
| 3254 | .as_array() |
| 3255 | .unwrap() |
| 3256 | .iter() |
| 3257 | .filter_map(|t| t["id"].as_str()) |
| 3258 | .collect(); |
| 3259 | assert_eq!(summary_ids, vec![archived_id.as_str()]); |
| 3260 | |
| 3261 | handle.abort(); |
| 3262 | Ok(()) |
| 3263 | } |
| 3264 | |
| 3265 | /// #564 / whalescale#261 — `GET /v1/usage` aggregates per-turn token + |
| 3266 | /// cost data. With no threads the response is well-formed and totals are |
| 3267 | /// zero with empty buckets (never a 404). |
| 3268 | #[tokio::test] |
| 3269 | async fn usage_endpoint_returns_empty_aggregation_for_fresh_store() -> Result<()> { |
| 3270 | let Some((addr, _runtime_threads, handle)) = spawn_test_server().await? else { |
| 3271 | return Ok(()); |
| 3272 | }; |
| 3273 | let client = reqwest::Client::new(); |
| 3274 | |
| 3275 | let body: serde_json::Value = client |
| 3276 | .get(format!("http://{addr}/v1/usage")) |
| 3277 | .send() |
| 3278 | .await? |
| 3279 | .error_for_status()? |
| 3280 | .json() |
| 3281 | .await?; |
| 3282 | assert_eq!(body["group_by"], "day"); |
| 3283 | assert_eq!(body["totals"]["input_tokens"], 0); |
| 3284 | assert_eq!(body["totals"]["output_tokens"], 0); |
| 3285 | assert_eq!(body["totals"]["turns"], 0); |
| 3286 | assert!( |
| 3287 | body["buckets"].as_array().unwrap().is_empty(), |
| 3288 | "buckets must be empty when no turns exist: {body}" |
| 3289 | ); |
| 3290 | |
| 3291 | // group_by query options are validated. |
| 3292 | let bad_group = client |
| 3293 | .get(format!("http://{addr}/v1/usage?group_by=galaxy")) |
| 3294 | .send() |
| 3295 | .await?; |
| 3296 | assert_eq!(bad_group.status(), StatusCode::BAD_REQUEST); |
| 3297 | |
| 3298 | // Each accepted group_by value succeeds. |
| 3299 | for gb in ["day", "model", "provider", "thread"] { |
| 3300 | let resp = client |
| 3301 | .get(format!("http://{addr}/v1/usage?group_by={gb}")) |
| 3302 | .send() |
| 3303 | .await?; |
| 3304 | assert!(resp.status().is_success(), "group_by={gb} failed: {resp:?}"); |
| 3305 | } |
| 3306 | |
| 3307 | // Bad ISO-8601 timestamp rejected. |
| 3308 | let bad_since = client |
| 3309 | .get(format!("http://{addr}/v1/usage?since=not-a-date")) |
| 3310 | .send() |
| 3311 | .await?; |
| 3312 | assert_eq!(bad_since.status(), StatusCode::BAD_REQUEST); |
| 3313 | |
| 3314 | // since > until rejected. |
| 3315 | let inverted = client |
| 3316 | .get(format!( |
| 3317 | "http://{addr}/v1/usage?since=2030-01-02T00:00:00Z&until=2030-01-01T00:00:00Z" |
| 3318 | )) |
| 3319 | .send() |
| 3320 | .await?; |
| 3321 | assert_eq!(inverted.status(), StatusCode::BAD_REQUEST); |
| 3322 | |
| 3323 | handle.abort(); |
| 3324 | Ok(()) |
| 3325 | } |
| 3326 | } |
| 3327 |