1use super::{permission, result_collector, update_mapper};
27use async_trait::async_trait;
28use luft_core::contract::backend::{
29 AgentBackend, AgentCapabilities, AgentResult, AgentTask, BackendError, RunContext, ToolPolicy,
30};
31use luft_core::contract::event::EventSender;
32#[cfg(feature = "unstable_end_turn_token_usage")]
33use luft_core::contract::ids::TokenUsage;
34use luft_core::contract::ids::{AgentId, RunId};
35use std::path::PathBuf;
36use std::process::Stdio;
37use std::sync::{Arc, Mutex};
38use std::time::Duration;
39use tokio_util::compat::{Compat, TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt};
40
41type AcpTransport =
43 ByteStreams<Compat<tokio::process::ChildStdin>, Compat<tokio::process::ChildStdout>>;
44
45use agent_client_protocol::schema::{
46 ContentBlock, InitializeRequest, McpServer, McpServerStdio, NewSessionRequest,
47 NewSessionResponse, PromptRequest, ProtocolVersion, RequestPermissionOutcome,
48 RequestPermissionRequest, RequestPermissionResponse, SelectedPermissionOutcome,
49 SessionConfigKind, SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId,
50 SessionNotification, SetSessionConfigOptionRequest, StopReason, TextContent,
51};
52use agent_client_protocol::{Agent, ByteStreams, Client, ConnectionTo, Responder};
53
54const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
57
58const POST_SUBMISSION_IDLE: Duration = Duration::from_secs(5);
66
67const STOP_REASON_END_TURN: &str = "EndTurn";
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76enum WatchdogOutcome {
77 PreIdleTimeout,
79 PostSubmissionTimeout,
83 ChannelClosed,
86}
87
88struct SessionState {
91 acc: Arc<update_mapper::Accumulator>,
92 stop_holder: Arc<Mutex<Option<String>>>,
93 events: EventSender,
94 activity_tx: tokio::sync::mpsc::UnboundedSender<()>,
95 activity_rx: tokio::sync::mpsc::UnboundedReceiver<()>,
96 submit_signal: Arc<tokio::sync::Notify>,
97 run_id: RunId,
98 agent_id: AgentId,
99 emit_raw: bool,
100 policy: Option<ToolPolicy>,
101 prompt: String,
102 cwd: PathBuf,
103}
104
105#[derive(Debug, Clone)]
107pub struct AcpConfig {
108 pub id: &'static str,
110 pub binary: PathBuf,
112 pub acp_args: Vec<String>,
114 pub log_level: Option<String>,
116 pub connect_timeout: Duration,
118 pub emit_raw_events: bool,
122 pub env_passthrough: Vec<String>,
141 pub model: Option<String>,
145}
146
147impl AcpConfig {
148 pub const DEFAULT_ENV_PASSTHROUGH: &'static [&'static str] = &[
153 "PATH",
155 "SYSTEMROOT",
156 "WINDIR",
157 "COMSPEC",
158 "PATHEXT",
159 "USERPROFILE",
161 "HOME",
162 "USER",
163 "USERNAME",
164 "LOGNAME",
165 "TMPDIR",
167 "TMP",
168 "TEMP",
169 "LANG",
171 "LC_ALL",
172 "LC_CTYPE",
173 "SHELL",
175 "APPDATA",
177 "LOCALAPPDATA",
178 "ProgramFiles",
179 "ProgramFiles(x86)",
180 ];
181}
182
183impl Default for AcpConfig {
184 fn default() -> Self {
185 Self {
186 id: "opencode",
187 binary: PathBuf::from("opencode"),
188 acp_args: vec!["acp".to_string()],
189 log_level: None,
190 connect_timeout: Duration::from_secs(10),
191 emit_raw_events: true,
192 env_passthrough: Self::DEFAULT_ENV_PASSTHROUGH
193 .iter()
194 .map(|s| s.to_string())
195 .collect(),
196 model: None,
197 }
198 }
199}
200
201pub struct AcpAdapter {
203 config: AcpConfig,
204}
205
206impl AcpAdapter {
207 pub fn new(config: AcpConfig) -> Self {
208 Self { config }
209 }
210
211 pub fn default_opencode() -> Self {
213 Self::new(AcpConfig::default())
214 }
215
216 pub fn config(&self) -> &AcpConfig {
220 &self.config
221 }
222}
223
224#[async_trait]
225impl AgentBackend for AcpAdapter {
226 fn id(&self) -> &'static str {
227 self.config.id
228 }
229
230 fn capabilities(&self) -> AgentCapabilities {
231 AgentCapabilities {
232 streaming: true,
233 mcp_injection: true,
234 structured_output: true,
235 models: vec![],
236 }
237 }
238
239 fn as_any(&self) -> &dyn std::any::Any {
240 self
241 }
242
243 async fn run(&self, task: AgentTask, ctx: RunContext) -> Result<AgentResult, BackendError> {
244 let config = self.config.clone();
245 let cancel = ctx.cancel.clone();
246 let events = ctx.events.clone();
247 let run_id = ctx.run_id;
248
249 let handle = tokio::task::spawn_blocking(move || {
252 let rt = tokio::runtime::Builder::new_current_thread()
253 .enable_all()
254 .build()
255 .map_err(|e| BackendError::Execution(format!("acp runtime: {e}")))?;
256 let local = tokio::task::LocalSet::new();
257 local.block_on(&rt, run_acp_session(config, task, run_id, cancel, events))
258 });
259
260 handle
261 .await
262 .map_err(|e| BackendError::Execution(format!("acp task join: {e}")))?
263 }
264}
265
266#[tracing::instrument(
268 name = "backend",
269 skip_all,
270 fields(run_id = %run_id, agent_id = %task.agent_id, backend = "opencode")
271)]
272async fn run_acp_session(
273 config: AcpConfig,
274 task: AgentTask,
275 run_id: RunId,
276 cancel: tokio_util::sync::CancellationToken,
277 events: EventSender,
278) -> Result<AgentResult, BackendError> {
279 let (mut child, transport) = spawn_agent(&config)?;
281
282 let (activity_tx, activity_rx) = tokio::sync::mpsc::unbounded_channel::<()>();
288 let submit_signal = Arc::new(tokio::sync::Notify::new());
289 let mut state = SessionState {
290 acc: Arc::new(update_mapper::Accumulator::new()),
291 stop_holder: Arc::new(Mutex::new(None)),
292 events: events.clone(),
293 activity_tx,
294 activity_rx,
295 submit_signal: submit_signal.clone(),
296 run_id,
297 agent_id: task.agent_id,
298 emit_raw: config.emit_raw_events,
299 policy: task.allowlist.clone(),
300 prompt: task.prompt.clone(),
301 cwd: std::fs::canonicalize(&task.workdir).unwrap_or_else(|_| task.workdir.clone()),
302 };
303
304 let schema_guard = prepare_schema_mcp(task.output_schema.as_ref())?;
308 let schema_file_path = schema_guard
309 .as_ref()
310 .map(|g| g.0.path().to_string_lossy().into_owned());
311
312 let conn_fut = drive_connection(&state, transport, schema_file_path, config.model.clone());
317 let idle_timeout = task.timeout.unwrap_or(DEFAULT_IDLE_TIMEOUT);
318
319 let outcome = tokio::select! {
324 r = conn_fut => r,
325 _ = cancel.cancelled() => {
326 tracing::debug!("ACP session cancelled");
327 let _ = child.start_kill();
328 return Err(BackendError::Cancelled);
329 }
330 res = idle_watchdog(
331 idle_timeout,
332 POST_SUBMISSION_IDLE,
333 &mut state.activity_rx,
334 state.submit_signal.clone(),
335 ) => {
336 handle_watchdog_outcome(res, &mut child, &state.stop_holder, idle_timeout)?;
337 Ok(())
338 }
339 };
340 let _ = child.start_kill();
341
342 outcome.map_err(classify_protocol_error)?;
343
344 Ok(collect_session_result(&task, &state))
346}
347
348fn spawn_agent(config: &AcpConfig) -> Result<(tokio::process::Child, AcpTransport), BackendError> {
358 let mut cmd = tokio::process::Command::new(&config.binary);
359 cmd.args(&config.acp_args);
360 if let Some(level) = &config.log_level {
361 cmd.arg("--log-level").arg(level);
362 }
363 cmd.stdin(Stdio::piped())
364 .stdout(Stdio::piped())
365 .stderr(Stdio::null());
366
367 cmd.env_clear();
368 for name in &config.env_passthrough {
369 if let Ok(value) = std::env::var(name) {
370 cmd.env(name, value);
371 }
372 }
373
374 let mut child = cmd.spawn().map_err(|e| {
375 tracing::error!(binary = %config.binary.display(), error = %e, "failed to spawn ACP backend");
376 BackendError::Spawn(format!("failed to spawn {}: {e}", config.binary.display()))
377 })?;
378 let stdin = child
379 .stdin
380 .take()
381 .ok_or_else(|| BackendError::Spawn("no child stdin".into()))?;
382 let stdout = child
383 .stdout
384 .take()
385 .ok_or_else(|| BackendError::Spawn("no child stdout".into()))?;
386 let transport = ByteStreams::new(stdin.compat_write(), stdout.compat());
387 Ok((child, transport))
388}
389
390fn prepare_schema_mcp(
397 schema: Option<&serde_json::Value>,
398) -> Result<Option<SchemaFileGuard>, BackendError> {
399 let Some(schema) = schema else {
400 return Ok(None);
401 };
402 let schema_json = serde_json::to_string(schema)
403 .map_err(|e| BackendError::Execution(format!("schema serialize: {e}")))?;
404 let schema_file = tempfile::NamedTempFile::new()
405 .map_err(|e| BackendError::Execution(format!("schema temp file: {e}")))?;
406 std::fs::write(&schema_file, &schema_json)
407 .map_err(|e| BackendError::Execution(format!("schema temp write: {e}")))?;
408 let path = schema_file.path().to_string_lossy().into_owned();
409 tracing::debug!(schema_file = %path, "prepared MCP structured-output server");
410 Ok(Some(SchemaFileGuard(schema_file)))
411}
412
413struct SchemaFileGuard(tempfile::NamedTempFile);
414
415fn drive_connection(
420 state: &SessionState,
421 transport: AcpTransport,
422 schema_file_path: Option<String>,
423 model: Option<String>,
424) -> impl std::future::Future<Output = Result<(), agent_client_protocol::Error>> {
425 let acc = state.acc.clone();
426 let events = state.events.clone();
427 let stop_holder = state.stop_holder.clone();
428 let activity_tx = state.activity_tx.clone();
429 let submit_signal = state.submit_signal.clone();
430 let run_id = state.run_id;
431 let agent_id = state.agent_id;
432 let emit_raw = state.emit_raw;
433 let policy = state.policy.clone();
434
435 let acc_for_prompt = acc.clone();
436 let stop_holder_for_prompt = stop_holder.clone();
437 let cwd = state.cwd.clone();
438 let prompt = state.prompt.clone();
439
440 async move {
441 Client
442 .builder()
443 .name("luft")
444 .on_receive_notification(
445 {
446 let acc = acc.clone();
447 let events = events.clone();
448 let activity_tx = activity_tx.clone();
449 let submit_signal = submit_signal.clone();
450 move |n: SessionNotification, _cx: ConnectionTo<Agent>| {
451 let acc = acc.clone();
452 let events = events.clone();
453 let activity_tx = activity_tx.clone();
454 let submit_signal = submit_signal.clone();
455 async move {
456 handle_session_update(
457 n,
458 &acc,
459 &events,
460 &activity_tx,
461 &submit_signal,
462 run_id,
463 agent_id,
464 emit_raw,
465 );
466 Ok(())
467 }
468 }
469 },
470 agent_client_protocol::on_receive_notification!(),
471 )
472 .on_receive_request(
473 {
474 let policy = policy.clone();
475 move |req: RequestPermissionRequest,
476 responder: Responder<RequestPermissionResponse>,
477 _conn: ConnectionTo<Agent>| {
478 let policy = policy.clone();
479 async move { decide_permission(req, responder, policy).await }
480 }
481 },
482 agent_client_protocol::on_receive_request!(),
483 )
484 .connect_with(transport, move |conn: ConnectionTo<Agent>| {
485 let acc_for_prompt = acc_for_prompt.clone();
486 let stop_holder_for_prompt = stop_holder_for_prompt.clone();
487 let model = model.clone();
488 async move {
489 run_handshake_and_prompt(
490 &conn,
491 &cwd,
492 schema_file_path.as_deref(),
493 model.as_deref(),
494 &prompt,
495 &acc_for_prompt,
496 &stop_holder_for_prompt,
497 )
498 .await
499 }
500 })
501 .await
502 }
503}
504
505#[allow(clippy::too_many_arguments)]
509fn handle_session_update(
510 n: SessionNotification,
511 acc: &Arc<update_mapper::Accumulator>,
512 events: &EventSender,
513 activity_tx: &tokio::sync::mpsc::UnboundedSender<()>,
514 submit_signal: &Arc<tokio::sync::Notify>,
515 run_id: RunId,
516 agent_id: AgentId,
517 emit_raw: bool,
518) {
519 let _ = activity_tx.send(());
520 let kind = serde_json::to_value(&n.update)
521 .ok()
522 .and_then(|v| {
523 v.get("sessionUpdate")
524 .and_then(|v| v.as_str())
525 .map(String::from)
526 })
527 .unwrap_or_else(|| "unknown".to_string());
528 tracing::debug!(%kind, "ACP session/update");
529
530 let was_submitted = acc.structured_output.lock().unwrap().is_some();
534 update_mapper::handle_update(&n.update, run_id, agent_id, acc, events, emit_raw);
535 if !was_submitted && acc.structured_output.lock().unwrap().is_some() {
536 submit_signal.notify_one();
537 tracing::debug!(
538 "ACP structured_output captured; watchdog switching to post-submission mode"
539 );
540 }
541}
542
543async fn decide_permission(
547 req: RequestPermissionRequest,
548 responder: Responder<RequestPermissionResponse>,
549 policy: Option<ToolPolicy>,
550) -> Result<(), agent_client_protocol::Error> {
551 let inputs = permission::extract_inputs(&req);
552 let approve = matches!(
553 permission::decide(policy.as_ref(), &inputs),
554 permission::Decision::Approve
555 );
556 tracing::debug!(
557 approve,
558 options = req.options.len(),
559 "ACP permission request"
560 );
561 let outcome = match (approve, req.options.first()) {
562 (true, Some(opt)) => RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(
563 opt.option_id.clone(),
564 )),
565 _ => RequestPermissionOutcome::Cancelled,
566 };
567 responder.respond(RequestPermissionResponse::new(outcome))
568}
569
570async fn run_handshake_and_prompt(
577 conn: &ConnectionTo<Agent>,
578 cwd: &std::path::Path,
579 schema_file_path: Option<&str>,
580 model: Option<&str>,
581 prompt: &str,
582 acc: &Arc<update_mapper::Accumulator>,
583 stop_holder: &Arc<Mutex<Option<String>>>,
584) -> Result<(), agent_client_protocol::Error> {
585 tracing::debug!("ACP handshake: initialize");
586 conn.send_request(InitializeRequest::new(ProtocolVersion::V1))
587 .block_task()
588 .await?;
589
590 tracing::debug!("ACP handshake: session/new");
591 let ns = session_new(conn, cwd.to_path_buf(), schema_file_path).await?;
592
593 if let Some(model_name) = model {
594 validate_and_set_model(conn, &ns, model_name).await?;
595 }
596
597 tracing::debug!("ACP handshake: session/prompt");
598 let pr = send_prompt(conn, ns.session_id, prompt.to_string()).await?;
599 record_prompt_result(&pr, stop_holder, acc);
600 Ok(())
601}
602
603async fn session_new(
606 conn: &ConnectionTo<Agent>,
607 cwd: PathBuf,
608 schema_file_path: Option<&str>,
609) -> Result<NewSessionResponse, agent_client_protocol::Error> {
610 let req = NewSessionRequest::new(cwd);
611 let req = match schema_file_path {
612 Some(sf) => {
613 let luft_bin =
614 std::env::current_exe().unwrap_or_else(|_| std::path::PathBuf::from("luft"));
615 let mcp = McpServerStdio::new("luft-structured-output", luft_bin).args(vec![
616 "mcp-structured-output".to_string(),
617 "--schema-file".to_string(),
618 sf.to_string(),
619 ]);
620 req.mcp_servers(vec![McpServer::Stdio(mcp)])
621 }
622 None => req,
623 };
624 conn.send_request(req).block_task().await
625}
626
627async fn validate_and_set_model(
632 conn: &ConnectionTo<Agent>,
633 ns: &NewSessionResponse,
634 model_name: &str,
635) -> Result<(), agent_client_protocol::Error> {
636 let config_options = match ns.config_options.as_ref() {
637 Some(opts) => opts,
638 None => {
639 tracing::debug!("ACP: agent does not advertise config_options");
640 return Ok(());
641 }
642 };
643 let model_option = match config_options
644 .iter()
645 .find(|opt| opt.category.as_ref() == Some(&SessionConfigOptionCategory::Model))
646 {
647 Some(o) => o,
648 None => {
649 tracing::debug!("ACP: agent does not support model selection");
650 return Ok(());
651 }
652 };
653 let select = match &model_option.kind {
654 SessionConfigKind::Select(s) => s,
655 _ => {
656 tracing::debug!("ACP: model option is not a Select kind");
657 return Ok(());
658 }
659 };
660 let valid = match &select.options {
661 SessionConfigSelectOptions::Ungrouped(opts) => {
662 opts.iter().any(|o| o.value.0.as_ref() == model_name)
663 }
664 SessionConfigSelectOptions::Grouped(groups) => groups
665 .iter()
666 .any(|g| g.options.iter().any(|o| o.value.0.as_ref() == model_name)),
667 _ => false,
668 };
669 if valid {
670 tracing::debug!(model = %model_name, "ACP: setting session model");
671 let req = SetSessionConfigOptionRequest::new(
672 ns.session_id.clone(),
673 model_option.id.clone(),
674 model_name.to_string(),
675 );
676 conn.send_request(req).block_task().await?;
677 } else {
678 tracing::warn!(
679 model = %model_name,
680 "ACP: requested model not available, using agent default"
681 );
682 }
683 Ok(())
684}
685
686async fn send_prompt(
690 conn: &ConnectionTo<Agent>,
691 session_id: SessionId,
692 prompt: String,
693) -> Result<agent_client_protocol::schema::PromptResponse, agent_client_protocol::Error> {
694 conn.send_request(PromptRequest::new(
695 session_id,
696 vec![ContentBlock::Text(TextContent::new(prompt))],
697 ))
698 .block_task()
699 .await
700}
701
702fn record_prompt_result(
706 pr: &agent_client_protocol::schema::PromptResponse,
707 stop_holder: &Arc<Mutex<Option<String>>>,
708 #[cfg_attr(
709 not(feature = "unstable_end_turn_token_usage"),
710 allow(unused_variables)
711 )]
712 acc: &Arc<update_mapper::Accumulator>,
713) {
714 tracing::debug!(stop_reason = ?pr.stop_reason, "ACP prompt complete");
715 *stop_holder.lock().unwrap() = Some(stop_reason_as_str(&pr.stop_reason));
716 #[cfg(feature = "unstable_end_turn_token_usage")]
717 {
718 if let Some(u) = pr.usage.as_ref() {
719 tracing::debug!(
720 input = u.input_tokens,
721 output = u.output_tokens,
722 total = u.total_tokens,
723 "ACP prompt usage"
724 );
725 *acc.tokens.lock().unwrap() = TokenUsage {
726 input: u.input_tokens,
727 output: u.output_tokens,
728 cache_read: u.cached_read_tokens.unwrap_or(0),
729 cache_write: u.cached_write_tokens.unwrap_or(0),
730 };
731 }
732 }
733}
734
735fn stop_reason_as_str(r: &StopReason) -> String {
741 match r {
742 StopReason::EndTurn => STOP_REASON_END_TURN.to_string(),
743 StopReason::MaxTokens => "MaxTokens".to_string(),
744 StopReason::MaxTurnRequests => "MaxTurnRequests".to_string(),
745 StopReason::Refusal => "Refusal".to_string(),
746 StopReason::Cancelled => "Cancelled".to_string(),
747 #[allow(unreachable_patterns)]
748 other => format!("{other:?}"),
749 }
750}
751
752fn handle_watchdog_outcome(
758 res: WatchdogOutcome,
759 child: &mut tokio::process::Child,
760 stop_holder: &Arc<Mutex<Option<String>>>,
761 idle_timeout: Duration,
762) -> Result<(), BackendError> {
763 let _ = child.start_kill();
764 match res {
765 WatchdogOutcome::PreIdleTimeout => {
766 tracing::warn!(
767 idle_timeout_ms = idle_timeout.as_millis() as u64,
768 "ACP session idle timeout (no protocol activity)"
769 );
770 Err(BackendError::Timeout)
771 }
772 WatchdogOutcome::ChannelClosed => {
773 tracing::debug!("ACP activity channel closed");
774 Err(BackendError::Timeout)
775 }
776 WatchdogOutcome::PostSubmissionTimeout => {
777 tracing::info!(
784 post_idle_ms = POST_SUBMISSION_IDLE.as_millis() as u64,
785 "ACP post-submission timeout; treating structured_output as result"
786 );
787 let mut guard = stop_holder.lock().unwrap();
792 if guard.is_none() {
793 *guard = Some(STOP_REASON_END_TURN.to_string());
794 }
795 Ok(())
796 }
797 }
798}
799
800fn classify_protocol_error(e: agent_client_protocol::Error) -> BackendError {
810 let s = e.to_string();
811 if is_connection_closed(&s) {
812 tracing::warn!("ACP connection closed");
813 BackendError::Protocol("connection closed".into())
814 } else {
815 tracing::error!(error = %s, "ACP protocol error");
816 BackendError::Protocol(s)
817 }
818}
819
820fn is_connection_closed(s: &str) -> bool {
821 s.contains("receiver dropped")
825 || s.contains("broken pipe")
826 || s.contains("unexpected eof")
827 || s.contains("connection closed")
828}
829
830fn collect_session_result(task: &AgentTask, state: &SessionState) -> AgentResult {
832 let stop = state.stop_holder.lock().unwrap().take().unwrap_or_default();
833 let message = std::mem::take(&mut *state.acc.message.lock().unwrap());
834 let tokens = *state.acc.tokens.lock().unwrap();
835 let structured = state.acc.structured_output.lock().unwrap().take();
836 result_collector::collect(task, &stop, message, tokens, structured)
837}
838
839async fn idle_watchdog(
862 pre_idle: Duration,
863 post_idle: Duration,
864 activity_rx: &mut tokio::sync::mpsc::UnboundedReceiver<()>,
865 submit_signal: Arc<tokio::sync::Notify>,
866) -> WatchdogOutcome {
867 let mut submitted = false;
868 loop {
869 if submitted {
870 while activity_rx.try_recv().is_ok() {}
879 tokio::time::sleep(post_idle).await;
880 return WatchdogOutcome::PostSubmissionTimeout;
881 }
882 tokio::select! {
883 biased;
884 _ = submit_signal.notified() => {
885 submitted = true;
886 tracing::debug!(
887 post_idle_ms = post_idle.as_millis() as u64,
888 "ACP watchdog entered post-submission mode"
889 );
890 }
891 msg = activity_rx.recv() => match msg {
892 Some(()) => { while activity_rx.try_recv().is_ok() {} }
893 None => return WatchdogOutcome::ChannelClosed,
894 },
895 _ = tokio::time::sleep(pre_idle) => {
896 return WatchdogOutcome::PreIdleTimeout;
897 }
898 }
899 }
900}
901
902#[cfg(test)]
903mod tests {
904 use super::*;
905
906 #[tokio::test]
913 async fn idle_watchdog_fires_after_idle_period() {
914 let (_atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
915 let submit = Arc::new(tokio::sync::Notify::new());
916 let r = tokio::time::timeout(
917 Duration::from_millis(500),
918 idle_watchdog(
919 Duration::from_millis(50),
920 Duration::from_millis(50),
921 &mut arx,
922 submit,
923 ),
924 )
925 .await;
926 let outcome = r.expect("should fire after idle period");
927 assert_eq!(outcome, WatchdogOutcome::PreIdleTimeout);
928 }
929
930 #[tokio::test]
931 async fn idle_watchdog_does_not_fire_with_activity() {
932 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
933 let submit = Arc::new(tokio::sync::Notify::new());
934 tokio::spawn(async move {
935 for _ in 0..5 {
936 tokio::time::sleep(Duration::from_millis(20)).await;
937 let _ = atx.send(());
938 }
939 });
940 let r = tokio::time::timeout(
941 Duration::from_millis(80),
942 idle_watchdog(
943 Duration::from_millis(50),
944 Duration::from_millis(50),
945 &mut arx,
946 submit,
947 ),
948 )
949 .await;
950 assert!(
951 r.is_err(),
952 "should not fire while activity is within idle window"
953 );
954 }
955
956 #[tokio::test]
957 async fn idle_watchdog_fires_after_activity_stops() {
958 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
959 let submit = Arc::new(tokio::sync::Notify::new());
960 let _ = atx.send(());
961 drop(atx);
962 let r = tokio::time::timeout(
963 Duration::from_millis(30),
964 idle_watchdog(
965 Duration::from_millis(80),
966 Duration::from_millis(80),
967 &mut arx,
968 submit,
969 ),
970 )
971 .await;
972 let outcome = r.expect("should return immediately when channel closes");
973 assert_eq!(outcome, WatchdogOutcome::ChannelClosed);
974 }
975
976 #[tokio::test]
988 async fn idle_watchdog_enters_post_mode_after_submit_signal() {
989 let (_atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
990 let submit = Arc::new(tokio::sync::Notify::new());
991 let submit_h = submit.clone();
992 tokio::spawn(async move {
993 tokio::time::sleep(Duration::from_millis(20)).await;
994 submit_h.notify_one();
995 });
996 let r = tokio::time::timeout(
997 Duration::from_millis(500),
998 idle_watchdog(
999 Duration::from_secs(60), Duration::from_millis(50), &mut arx,
1002 submit,
1003 ),
1004 )
1005 .await;
1006 let outcome = r.expect("watchdog should return after post_idle");
1007 assert_eq!(outcome, WatchdogOutcome::PostSubmissionTimeout);
1008 }
1009
1010 #[tokio::test]
1011 async fn idle_watchdog_post_mode_is_not_reset_by_activity() {
1012 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
1016 let submit = Arc::new(tokio::sync::Notify::new());
1017 submit.notify_one();
1018 tokio::spawn(async move {
1019 for _ in 0..20 {
1020 tokio::time::sleep(Duration::from_millis(20)).await;
1021 let _ = atx.send(());
1022 }
1023 });
1024 let start = std::time::Instant::now();
1025 let r = tokio::time::timeout(
1026 Duration::from_millis(500),
1027 idle_watchdog(
1028 Duration::from_secs(60),
1029 Duration::from_millis(80),
1030 &mut arx,
1031 submit,
1032 ),
1033 )
1034 .await;
1035 let outcome = r.expect("watchdog should return after post_idle");
1036 let elapsed = start.elapsed();
1037 assert_eq!(outcome, WatchdogOutcome::PostSubmissionTimeout);
1038 assert!(
1042 elapsed < Duration::from_millis(300),
1043 "post-mode timer was reset by activity: elapsed={elapsed:?}"
1044 );
1045 }
1046
1047 #[tokio::test]
1048 async fn idle_watchdog_pre_mode_resets_on_activity() {
1049 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
1051 let submit = Arc::new(tokio::sync::Notify::new());
1052 tokio::spawn(async move {
1053 for _ in 0..10 {
1054 tokio::time::sleep(Duration::from_millis(30)).await;
1055 let _ = atx.send(());
1056 }
1057 });
1058 let r = tokio::time::timeout(
1059 Duration::from_millis(200),
1060 idle_watchdog(
1061 Duration::from_millis(60),
1062 Duration::from_millis(60),
1063 &mut arx,
1064 submit,
1065 ),
1066 )
1067 .await;
1068 assert!(
1069 r.is_err(),
1070 "pre-mode should not fire while activity keeps resetting timer"
1071 );
1072 }
1073
1074 #[test]
1082 fn stop_reason_as_str_end_turn_matches_constant() {
1083 assert_eq!(
1084 stop_reason_as_str(&StopReason::EndTurn),
1085 STOP_REASON_END_TURN
1086 );
1087 assert_eq!(stop_reason_as_str(&StopReason::EndTurn), "EndTurn");
1088 }
1089
1090 #[test]
1091 fn stop_reason_as_str_cancelled_contains_cancel() {
1092 assert_eq!(stop_reason_as_str(&StopReason::Cancelled), "Cancelled");
1093 }
1094
1095 #[test]
1096 fn stop_reason_as_str_other_variants_stable() {
1097 assert_eq!(stop_reason_as_str(&StopReason::MaxTokens), "MaxTokens");
1098 assert_eq!(
1099 stop_reason_as_str(&StopReason::MaxTurnRequests),
1100 "MaxTurnRequests"
1101 );
1102 assert_eq!(stop_reason_as_str(&StopReason::Refusal), "Refusal");
1103 }
1104
1105 #[tokio::test]
1111 async fn handle_watchdog_post_submission_synthesizes_end_turn() {
1112 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1113 let mut child = match tokio::process::Command::new("cmd")
1116 .arg("/C")
1117 .arg("exit 0")
1118 .stdin(Stdio::null())
1119 .stdout(Stdio::null())
1120 .stderr(Stdio::null())
1121 .kill_on_drop(true)
1122 .spawn()
1123 {
1124 Ok(c) => c,
1125 Err(_) => return, };
1127 let r = handle_watchdog_outcome(
1128 WatchdogOutcome::PostSubmissionTimeout,
1129 &mut child,
1130 &stop,
1131 Duration::from_secs(300),
1132 );
1133 assert!(
1134 r.is_ok(),
1135 "post-submission outcome should fall through to collect"
1136 );
1137 assert_eq!(stop.lock().unwrap().as_deref(), Some("EndTurn"));
1138 }
1139
1140 #[tokio::test]
1141 async fn handle_watchdog_post_submission_preserves_existing_stop() {
1142 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(Some("Cancelled".into())));
1143 let mut child = match tokio::process::Command::new("cmd")
1144 .arg("/C")
1145 .arg("exit 0")
1146 .stdin(Stdio::null())
1147 .stdout(Stdio::null())
1148 .stderr(Stdio::null())
1149 .kill_on_drop(true)
1150 .spawn()
1151 {
1152 Ok(c) => c,
1153 Err(_) => return,
1154 };
1155 let r = handle_watchdog_outcome(
1156 WatchdogOutcome::PostSubmissionTimeout,
1157 &mut child,
1158 &stop,
1159 Duration::from_secs(300),
1160 );
1161 assert!(r.is_ok());
1162 assert_eq!(stop.lock().unwrap().as_deref(), Some("Cancelled"));
1164 }
1165
1166 #[tokio::test]
1167 async fn handle_watchdog_pre_idle_returns_timeout_error() {
1168 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1169 let mut child = match tokio::process::Command::new("cmd")
1170 .arg("/C")
1171 .arg("exit 0")
1172 .stdin(Stdio::null())
1173 .stdout(Stdio::null())
1174 .stderr(Stdio::null())
1175 .kill_on_drop(true)
1176 .spawn()
1177 {
1178 Ok(c) => c,
1179 Err(_) => return,
1180 };
1181 let r = handle_watchdog_outcome(
1182 WatchdogOutcome::PreIdleTimeout,
1183 &mut child,
1184 &stop,
1185 Duration::from_secs(1),
1186 );
1187 assert!(matches!(r, Err(BackendError::Timeout)));
1188 }
1189
1190 #[tokio::test]
1191 async fn handle_watchdog_channel_closed_returns_timeout_error() {
1192 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1193 let mut child = match tokio::process::Command::new("cmd")
1194 .arg("/C")
1195 .arg("exit 0")
1196 .stdin(Stdio::null())
1197 .stdout(Stdio::null())
1198 .stderr(Stdio::null())
1199 .kill_on_drop(true)
1200 .spawn()
1201 {
1202 Ok(c) => c,
1203 Err(_) => return,
1204 };
1205 let r = handle_watchdog_outcome(
1206 WatchdogOutcome::ChannelClosed,
1207 &mut child,
1208 &stop,
1209 Duration::from_secs(1),
1210 );
1211 assert!(matches!(r, Err(BackendError::Timeout)));
1212 }
1213
1214 #[test]
1217 fn is_connection_closed_matches_documented_substrings() {
1218 assert!(is_connection_closed("receiver dropped"));
1219 assert!(is_connection_closed("broken pipe"));
1220 assert!(is_connection_closed("unexpected eof"));
1221 assert!(is_connection_closed("connection closed"));
1222 assert!(is_connection_closed(
1223 "io error: broken pipe writing to stdin"
1224 ));
1225 assert!(!is_connection_closed("unknown protocol method"));
1226 assert!(!is_connection_closed(""));
1227 }
1228}