1use std::sync::Arc;
4use std::time::Duration;
5
6use serde_json::{json, Map, Value};
7use tokio::sync::Mutex;
8
9use crate::client::{unwrap_next_frame, Client, SendInputOptions};
10use crate::errors::{Error, Result};
11use crate::session::{bootstrap_loop_session, connect_with_retries, BootstrapOptions};
12use crate::stream_terminal::{is_turn_end_custom_data, is_turn_progress_chunk, STREAM_END};
13use crate::turn_boundary::{
14 frame_turn_id, is_idle_terminal_allowed, is_turn_terminal_allowed, parse_turn_generation,
15 turn_ids_match,
16};
17
18use super::chunk_filter::should_drop_stream_chunk_early;
19use super::observability::TurnEventStats;
20
21pub const DEFAULT_POST_IDLE_DRAIN: Duration = Duration::from_millis(500);
23
24pub type EarlyDropFn = Arc<dyn Fn(&[Value], &str, &Value) -> bool + Send + Sync>;
26
27#[derive(Clone)]
29pub struct DaemonSessionOptions {
30 pub workspace: Option<String>,
32 pub stream_delivery: String,
34 pub post_idle_drain: Duration,
36 pub early_drop_fn: Option<EarlyDropFn>,
38}
39
40impl Default for DaemonSessionOptions {
41 fn default() -> Self {
42 Self {
43 workspace: None,
44 stream_delivery: "adaptive".into(),
45 post_idle_drain: DEFAULT_POST_IDLE_DRAIN,
46 early_drop_fn: None,
47 }
48 }
49}
50
51impl std::fmt::Debug for DaemonSessionOptions {
52 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
53 f.debug_struct("DaemonSessionOptions")
54 .field("workspace", &self.workspace)
55 .field("stream_delivery", &self.stream_delivery)
56 .field("post_idle_drain", &self.post_idle_drain)
57 .field(
58 "early_drop_fn",
59 &self.early_drop_fn.as_ref().map(|_| "<fn>"),
60 )
61 .finish()
62 }
63}
64
65#[derive(Debug, Clone, Default)]
67pub struct SendTurnOptions {
68 pub preferred_subagent: Option<String>,
70 pub intake_scope: Option<String>,
72 pub model: Option<String>,
74 pub model_params: Option<Value>,
76 pub attachments: Option<Value>,
78 pub clarification_mode: Option<String>,
80 pub clarification_answer: bool,
82 pub intent_hint: Option<String>,
84}
85
86#[derive(Debug, Clone)]
88pub struct TurnChunk {
89 pub namespace: Value,
91 pub mode: String,
93 pub data: Value,
95}
96
97pub struct DaemonSession {
99 opts: DaemonSessionOptions,
100 client: Client,
101 rpc_client: Client,
102 rpc_connected: Mutex<bool>,
103 loop_id: Mutex<String>,
104 read_lock: Mutex<()>,
105 early_drop_fn: EarlyDropFn,
106 pub turn_event_stats: Mutex<TurnEventStats>,
108 pub last_turn_end_state: Mutex<String>,
110 pub last_turn_cancel_seen: Mutex<bool>,
112 pub last_turn_error_message: Mutex<String>,
114}
115
116impl DaemonSession {
117 pub fn new(ws_url: impl Into<String>, opts: Option<DaemonSessionOptions>) -> Self {
119 let ws_url = ws_url.into();
120 let opts = opts.unwrap_or_default();
121 let early_drop_fn = opts.early_drop_fn.clone().unwrap_or_else(|| {
122 Arc::new(|ns: &[Value], mode: &str, data: &Value| {
123 should_drop_stream_chunk_early(ns, mode, data)
124 })
125 });
126 Self {
127 client: Client::new(&ws_url),
128 rpc_client: Client::new(&ws_url),
129 rpc_connected: Mutex::new(false),
130 loop_id: Mutex::new(String::new()),
131 read_lock: Mutex::new(()),
132 early_drop_fn,
133 turn_event_stats: Mutex::new(TurnEventStats::new()),
134 last_turn_end_state: Mutex::new(String::new()),
135 last_turn_cancel_seen: Mutex::new(false),
136 last_turn_error_message: Mutex::new(String::new()),
137 opts,
138 }
139 }
140
141 pub fn stream_client(&self) -> &Client {
143 &self.client
144 }
145
146 pub fn rpc_client(&self) -> &Client {
148 &self.rpc_client
149 }
150
151 pub async fn loop_id(&self) -> String {
153 self.loop_id.lock().await.clone()
154 }
155
156 pub async fn connect(&self, resume_loop_id: Option<&str>) -> Result<Map<String, Value>> {
158 connect_with_retries(&self.client, 40, Duration::from_millis(250)).await?;
159 self.bootstrap_loop(resume_loop_id).await
160 }
161
162 async fn bootstrap_loop(&self, resume_loop_id: Option<&str>) -> Result<Map<String, Value>> {
163 let mut boot = BootstrapOptions::new();
164 boot.resume_loop_id = resume_loop_id.map(|s| s.to_string());
165 boot.workspace = self.opts.workspace.clone();
166 boot.stream_delivery = self.opts.stream_delivery.clone();
167 let ready = bootstrap_loop_session(&self.client, boot, None).await?;
168 if let Some(lid) = ready.get("loop_id").and_then(|v| v.as_str()) {
169 *self.loop_id.lock().await = lid.to_string();
170 }
171 Ok(ready)
172 }
173
174 pub async fn new_loop(&self) -> Result<Map<String, Value>> {
176 self.bootstrap_loop(None).await
177 }
178
179 pub async fn switch_loop(&self, loop_id: &str) -> Result<Map<String, Value>> {
181 self.bootstrap_loop(Some(loop_id)).await
182 }
183
184 pub async fn ensure_connected(&self) -> Result<()> {
186 if self.client.is_connection_alive() {
187 return Ok(());
188 }
189 connect_with_retries(&self.client, 40, Duration::from_millis(250)).await?;
190 let lid = self.loop_id().await;
191 if lid.is_empty() {
192 self.bootstrap_loop(None).await?;
193 return Ok(());
194 }
195 match self.client.reattach_and_probe(&lid).await {
196 Ok(()) => Ok(()),
197 Err(Error::StaleLoop(_)) | Err(_) => {
198 let _ = self.rpc_client.close().await;
200 *self.rpc_connected.lock().await = false;
201 self.bootstrap_loop(None).await?;
202 Ok(())
203 }
204 }
205 }
206
207 pub async fn close(&self) -> Result<()> {
209 let _ = self.client.close().await;
210 let _ = self.rpc_client.close().await;
211 *self.rpc_connected.lock().await = false;
212 Ok(())
213 }
214
215 pub async fn detach(&self) -> Result<()> {
217 self.client.notify("disconnect", Map::new()).await
218 }
219
220 pub async fn send_turn(&self, text: &str, opts: Option<SendTurnOptions>) -> Result<()> {
222 let loop_id = self.loop_id().await;
223 if loop_id.is_empty() {
224 return Err(Error::msg("no active loop session"));
225 }
226 let opts = opts.unwrap_or_default();
227 let input = SendInputOptions {
228 loop_id: Some(loop_id),
229 preferred_subagent: opts.preferred_subagent,
230 intake_scope: opts.intake_scope,
231 model: opts.model,
232 model_params: opts.model_params,
233 attachments: opts.attachments,
234 clarification_mode: opts.clarification_mode,
235 clarification_answer: opts.clarification_answer,
236 intent_hint: opts.intent_hint,
237 ..Default::default()
238 };
239 self.client.send_input(text, input).await
240 }
241
242 pub async fn cancel_active_turn(&self) -> Result<()> {
244 let mut params = Map::new();
245 params.insert("cmd".into(), json!("/cancel"));
246 self.client.notify("slash_command", params).await
247 }
248
249 async fn ensure_rpc_connected(&self) -> Result<()> {
250 let mut flag = self.rpc_connected.lock().await;
251 if *flag && self.rpc_client.is_connected() {
252 return Ok(());
253 }
254 connect_with_retries(&self.rpc_client, 5, Duration::from_millis(250)).await?;
255 *flag = true;
256 Ok(())
257 }
258
259 pub async fn list_loops(&self, limit: u32) -> Result<Map<String, Value>> {
261 self.ensure_rpc_connected().await?;
262 let lim = if limit == 0 { 20 } else { limit };
263 self.rpc_client.loop_list(lim).await
264 }
265
266 pub async fn fetch_loop_history(&self, loop_id: &str) -> Result<Map<String, Value>> {
268 self.ensure_rpc_connected().await?;
269 self.rpc_client.loop_history_fetch(loop_id).await
270 }
271
272 pub async fn iter_turn_chunks(&self, max_wait: Option<Duration>) -> Result<Vec<TurnChunk>> {
276 let _guard = self.read_lock.lock().await;
277 self.iter_turn_chunks_locked(max_wait).await
278 }
279
280 async fn iter_turn_chunks_locked(&self, max_wait: Option<Duration>) -> Result<Vec<TurnChunk>> {
281 *self.last_turn_end_state.lock().await = String::new();
282 *self.last_turn_error_message.lock().await = String::new();
283 *self.last_turn_cancel_seen.lock().await = false;
284 *self.turn_event_stats.lock().await = TurnEventStats::new();
285
286 let mut out = Vec::new();
287 let mut query_started = false;
288 let mut expected_loop_id = self.loop_id().await;
289 let mut expected_turn_id: Option<String> = None;
290 let mut turn_progress_seen = false;
291 let mut cancel_seen = false;
292 let absolute_deadline = max_wait.map(|d| tokio::time::Instant::now() + d);
293
294 let _ = self.client.peel_stale_pending_control_events().await;
295
296 loop {
297 if let Some(deadline) = absolute_deadline {
298 if tokio::time::Instant::now() > deadline {
299 let err = format!(
300 "turn timed out after {:?} (loop={})",
301 max_wait.unwrap_or_default(),
302 expected_loop_id
303 );
304 *self.last_turn_error_message.lock().await = err.clone();
305 return Err(Error::msg(err));
306 }
307 }
308
309 let ev = self
310 .client
311 .read_event_with_timeout(Duration::from_millis(250))
312 .await?;
313 let Some(ev) = ev else {
314 if query_started && !self.client.is_connection_alive() {
315 *self.last_turn_end_state.lock().await = "connection_lost".into();
316 return Err(Error::msg("daemon connection lost"));
317 }
318 continue;
320 };
321
322 let mut frame = ev;
323 let mut event_type = frame
324 .get("type")
325 .and_then(|v| v.as_str())
326 .unwrap_or("")
327 .to_string();
328 if event_type == "next" {
329 frame = unwrap_next_frame(&frame);
330 event_type = frame
331 .get("type")
332 .and_then(|v| v.as_str())
333 .unwrap_or("")
334 .to_string();
335 }
336
337 let event_loop_id = frame
338 .get("loop_id")
339 .and_then(|v| v.as_str())
340 .unwrap_or("")
341 .to_string();
342 if !expected_loop_id.is_empty()
343 && !event_loop_id.is_empty()
344 && event_loop_id != expected_loop_id
345 {
346 continue;
347 }
348
349 let ev_turn = frame_turn_id(Some(&frame));
350 let status_state = if event_type == "status" {
351 frame.get("state").and_then(|v| v.as_str()).unwrap_or("")
352 } else {
353 ""
354 };
355 let is_running_status = status_state == "running";
356 let is_terminal_status = status_state == "idle" || status_state == "stopped";
357 if let Some(ref expected) = expected_turn_id {
358 if (event_type == "event" || event_type == "status") && !is_running_status {
359 if is_terminal_status {
360 if let Some(ref tid) = ev_turn {
361 if !turn_ids_match(Some(expected.as_str()), Some(tid.as_str())) {
362 continue;
363 }
364 }
365 } else if !turn_ids_match(Some(expected.as_str()), ev_turn.as_deref()) {
366 continue;
367 }
368 }
369 }
370
371 if event_type == "error" {
372 let msg = frame
373 .get("error")
374 .and_then(|e| e.get("message"))
375 .and_then(|m| m.as_str())
376 .or_else(|| frame.get("message").and_then(|m| m.as_str()))
377 .unwrap_or("daemon error")
378 .to_string();
379 *self.last_turn_error_message.lock().await = msg.clone();
380 return Err(Error::msg(msg));
381 }
382
383 if event_type == "status" {
384 if let Some(lid) = frame.get("loop_id").and_then(|v| v.as_str()) {
385 if !lid.is_empty() {
386 *self.loop_id.lock().await = lid.to_string();
387 expected_loop_id = lid.to_string();
388 }
389 }
390 match status_state {
391 "running" => {
392 query_started = true;
393 if let Some(status_turn) = frame_turn_id(Some(&frame)) {
394 let new_gen = parse_turn_generation(Some(&status_turn));
395 let old_gen = parse_turn_generation(expected_turn_id.as_deref());
396 if expected_turn_id.is_none()
397 || (new_gen.is_some()
398 && (old_gen.is_none() || new_gen.unwrap() >= old_gen.unwrap()))
399 {
400 if expected_turn_id.as_ref().is_some_and(|e| e != &status_turn) {
401 turn_progress_seen = false;
402 }
403 expected_turn_id = Some(status_turn);
404 }
405 }
406 }
407 "stopped" if query_started => {
408 let stop_turn = frame_turn_id(Some(&frame));
409 if expected_turn_id.is_some()
410 && !turn_ids_match(expected_turn_id.as_deref(), stop_turn.as_deref())
411 {
412 continue;
413 }
414 *self.last_turn_end_state.lock().await = status_state.into();
415 self.drain_after_idle(&expected_loop_id, &mut out).await;
416 return Ok(out);
417 }
418 "idle" if query_started => {
419 let idle_turn = frame_turn_id(Some(&frame));
420 if !is_idle_terminal_allowed(
421 expected_turn_id.as_deref(),
422 idle_turn.as_deref(),
423 query_started,
424 turn_progress_seen,
425 cancel_seen,
426 ) {
427 continue;
428 }
429 *self.last_turn_end_state.lock().await = status_state.into();
430 self.drain_after_idle(&expected_loop_id, &mut out).await;
431 return Ok(out);
432 }
433 _ => {}
434 }
435 continue;
436 }
437
438 if event_type == "command_response" {
439 let content = frame.get("content").and_then(|v| v.as_str()).unwrap_or("");
440 if content.contains("Cancellation requested") {
441 cancel_seen = true;
442 *self.last_turn_cancel_seen.lock().await = true;
443 }
444 continue;
445 }
446
447 if event_type != "event" {
448 continue;
449 }
450
451 let data = frame.get("data").cloned().unwrap_or(Value::Null);
452 let namespace = frame
453 .get("namespace")
454 .cloned()
455 .unwrap_or(Value::Array(vec![]));
456 let mode = frame
457 .get("mode")
458 .and_then(|v| v.as_str())
459 .unwrap_or("")
460 .to_string();
461
462 let ns_slice: Vec<Value> = match &namespace {
463 Value::Array(a) => a.clone(),
464 _ => vec![],
465 };
466 if (self.early_drop_fn)(&ns_slice, &mode, &data) {
467 self.turn_event_stats.lock().await.filtered_early += 1;
468 continue;
469 }
470
471 if mode == "custom" && is_turn_end_custom_data(&data) {
472 let data_turn = frame_turn_id(Some(&data)).or_else(|| ev_turn.clone());
473 if !is_turn_terminal_allowed(
474 expected_turn_id.as_deref(),
475 data_turn.as_deref(),
476 query_started,
477 turn_progress_seen,
478 ) {
479 continue;
480 }
481 }
482
483 if is_turn_progress_chunk(&mode, &data) {
484 turn_progress_seen = true;
485 }
486
487 out.push(TurnChunk {
488 namespace,
489 mode: mode.clone(),
490 data: data.clone(),
491 });
492
493 if mode == "custom" && is_turn_end_custom_data(&data) {
494 let custom_type = data.get("type").and_then(|v| v.as_str()).unwrap_or("");
495 *self.last_turn_end_state.lock().await = if custom_type == STREAM_END {
496 "stream_end".into()
497 } else {
498 "completed".into()
499 };
500 self.drain_after_idle(&expected_loop_id, &mut out).await;
501 return Ok(out);
502 }
503 }
504 }
505
506 async fn drain_after_idle(&self, expected_loop_id: &str, out: &mut Vec<TurnChunk>) {
507 let deadline = tokio::time::Instant::now() + self.opts.post_idle_drain;
508 while tokio::time::Instant::now() < deadline {
509 let ev = match self
510 .client
511 .read_event_with_timeout(Duration::from_millis(250))
512 .await
513 {
514 Ok(Some(e)) => e,
515 _ => return,
516 };
517 let mut frame = ev;
518 let mut event_type = frame
519 .get("type")
520 .and_then(|v| v.as_str())
521 .unwrap_or("")
522 .to_string();
523 if event_type == "next" {
524 frame = unwrap_next_frame(&frame);
525 event_type = frame
526 .get("type")
527 .and_then(|v| v.as_str())
528 .unwrap_or("")
529 .to_string();
530 }
531 let event_loop_id = frame.get("loop_id").and_then(|v| v.as_str()).unwrap_or("");
532 if !expected_loop_id.is_empty()
533 && !event_loop_id.is_empty()
534 && event_loop_id != expected_loop_id
535 {
536 continue;
537 }
538 if event_type == "error" {
539 return;
540 }
541 if event_type == "status" {
542 if let Some(lid) = frame.get("loop_id").and_then(|v| v.as_str()) {
543 if !lid.is_empty() {
544 *self.loop_id.lock().await = lid.to_string();
545 }
546 }
547 continue;
548 }
549 if event_type != "event" {
550 continue;
551 }
552 let data = frame.get("data").cloned().unwrap_or(Value::Null);
553 let namespace = frame
554 .get("namespace")
555 .cloned()
556 .unwrap_or(Value::Array(vec![]));
557 let mode = frame
558 .get("mode")
559 .and_then(|v| v.as_str())
560 .unwrap_or("")
561 .to_string();
562 let ns_slice: Vec<Value> = match &namespace {
563 Value::Array(a) => a.clone(),
564 _ => vec![],
565 };
566 if (self.early_drop_fn)(&ns_slice, &mode, &data) {
567 self.turn_event_stats.lock().await.filtered_early += 1;
568 continue;
569 }
570 self.turn_event_stats.lock().await.post_idle_drained += 1;
571 out.push(TurnChunk {
572 namespace,
573 mode,
574 data,
575 });
576 }
577 }
578}