1use super::{permission, result_collector, update_mapper};
27use luft_core::contract::backend::{
28 AgentBackend, AgentCapabilities, AgentResult, AgentTask, BackendError, RunContext, ToolPolicy,
29};
30use luft_core::contract::event::EventSender;
31#[cfg(feature = "unstable_end_turn_token_usage")]
32use luft_core::contract::ids::TokenUsage;
33use luft_core::contract::ids::{AgentId, RunId};
34use async_trait::async_trait;
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 ];
176}
177
178impl Default for AcpConfig {
179 fn default() -> Self {
180 Self {
181 id: "opencode",
182 binary: PathBuf::from("opencode"),
183 acp_args: vec!["acp".to_string()],
184 log_level: None,
185 connect_timeout: Duration::from_secs(10),
186 emit_raw_events: true,
187 env_passthrough: Self::DEFAULT_ENV_PASSTHROUGH
188 .iter()
189 .map(|s| s.to_string())
190 .collect(),
191 model: None,
192 }
193 }
194}
195
196pub struct AcpAdapter {
198 config: AcpConfig,
199}
200
201impl AcpAdapter {
202 pub fn new(config: AcpConfig) -> Self {
203 Self { config }
204 }
205
206 pub fn default_opencode() -> Self {
208 Self::new(AcpConfig::default())
209 }
210
211 pub fn config(&self) -> &AcpConfig {
215 &self.config
216 }
217}
218
219#[async_trait]
220impl AgentBackend for AcpAdapter {
221 fn id(&self) -> &'static str {
222 self.config.id
223 }
224
225 fn capabilities(&self) -> AgentCapabilities {
226 AgentCapabilities {
227 streaming: true,
228 mcp_injection: true,
229 structured_output: true,
230 models: vec![],
231 }
232 }
233
234 fn as_any(&self) -> &dyn std::any::Any {
235 self
236 }
237
238 async fn run(&self, task: AgentTask, ctx: RunContext) -> Result<AgentResult, BackendError> {
239 let config = self.config.clone();
240 let cancel = ctx.cancel.clone();
241 let events = ctx.events.clone();
242 let run_id = ctx.run_id;
243
244 let handle = tokio::task::spawn_blocking(move || {
247 let rt = tokio::runtime::Builder::new_current_thread()
248 .enable_all()
249 .build()
250 .map_err(|e| BackendError::Execution(format!("acp runtime: {e}")))?;
251 let local = tokio::task::LocalSet::new();
252 local.block_on(&rt, run_acp_session(config, task, run_id, cancel, events))
253 });
254
255 handle
256 .await
257 .map_err(|e| BackendError::Execution(format!("acp task join: {e}")))?
258 }
259}
260
261#[tracing::instrument(
263 name = "backend",
264 skip_all,
265 fields(run_id = %run_id, agent_id = %task.agent_id, backend = "opencode")
266)]
267async fn run_acp_session(
268 config: AcpConfig,
269 task: AgentTask,
270 run_id: RunId,
271 cancel: tokio_util::sync::CancellationToken,
272 events: EventSender,
273) -> Result<AgentResult, BackendError> {
274 let (mut child, transport) = spawn_agent(&config)?;
276
277 let (activity_tx, activity_rx) = tokio::sync::mpsc::unbounded_channel::<()>();
283 let submit_signal = Arc::new(tokio::sync::Notify::new());
284 let mut state = SessionState {
285 acc: Arc::new(update_mapper::Accumulator::new()),
286 stop_holder: Arc::new(Mutex::new(None)),
287 events: events.clone(),
288 activity_tx,
289 activity_rx,
290 submit_signal: submit_signal.clone(),
291 run_id,
292 agent_id: task.agent_id,
293 emit_raw: config.emit_raw_events,
294 policy: task.allowlist.clone(),
295 prompt: task.prompt.clone(),
296 cwd: std::fs::canonicalize(&task.workdir).unwrap_or_else(|_| task.workdir.clone()),
297 };
298
299 let schema_guard = prepare_schema_mcp(task.output_schema.as_ref())?;
303 let schema_file_path = schema_guard
304 .as_ref()
305 .map(|g| g.0.path().to_string_lossy().into_owned());
306
307 let conn_fut = drive_connection(&state, transport, schema_file_path, config.model.clone());
312 let idle_timeout = task.timeout.unwrap_or(DEFAULT_IDLE_TIMEOUT);
313
314 let outcome = tokio::select! {
319 r = conn_fut => r,
320 _ = cancel.cancelled() => {
321 tracing::debug!("ACP session cancelled");
322 let _ = child.start_kill();
323 return Err(BackendError::Cancelled);
324 }
325 res = idle_watchdog(
326 idle_timeout,
327 POST_SUBMISSION_IDLE,
328 &mut state.activity_rx,
329 state.submit_signal.clone(),
330 ) => {
331 handle_watchdog_outcome(res, &mut child, &state.stop_holder, idle_timeout)?;
332 Ok(())
333 }
334 };
335 let _ = child.start_kill();
336
337 outcome.map_err(classify_protocol_error)?;
338
339 Ok(collect_session_result(&task, &state))
341}
342
343fn spawn_agent(config: &AcpConfig) -> Result<(tokio::process::Child, AcpTransport), BackendError> {
353 let mut cmd = tokio::process::Command::new(&config.binary);
354 cmd.args(&config.acp_args);
355 if let Some(level) = &config.log_level {
356 cmd.arg("--log-level").arg(level);
357 }
358 cmd.stdin(Stdio::piped())
359 .stdout(Stdio::piped())
360 .stderr(Stdio::null());
361
362 cmd.env_clear();
363 for name in &config.env_passthrough {
364 if let Ok(value) = std::env::var(name) {
365 cmd.env(name, value);
366 }
367 }
368
369 let mut child = cmd.spawn().map_err(|e| {
370 tracing::error!(binary = %config.binary.display(), error = %e, "failed to spawn ACP backend");
371 BackendError::Spawn(format!("failed to spawn {}: {e}", config.binary.display()))
372 })?;
373 let stdin = child
374 .stdin
375 .take()
376 .ok_or_else(|| BackendError::Spawn("no child stdin".into()))?;
377 let stdout = child
378 .stdout
379 .take()
380 .ok_or_else(|| BackendError::Spawn("no child stdout".into()))?;
381 let transport = ByteStreams::new(stdin.compat_write(), stdout.compat());
382 Ok((child, transport))
383}
384
385fn prepare_schema_mcp(
392 schema: Option<&serde_json::Value>,
393) -> Result<Option<SchemaFileGuard>, BackendError> {
394 let Some(schema) = schema else {
395 return Ok(None);
396 };
397 let schema_json = serde_json::to_string(schema)
398 .map_err(|e| BackendError::Execution(format!("schema serialize: {e}")))?;
399 let schema_file = tempfile::NamedTempFile::new()
400 .map_err(|e| BackendError::Execution(format!("schema temp file: {e}")))?;
401 std::fs::write(&schema_file, &schema_json)
402 .map_err(|e| BackendError::Execution(format!("schema temp write: {e}")))?;
403 let path = schema_file.path().to_string_lossy().into_owned();
404 tracing::debug!(schema_file = %path, "prepared MCP structured-output server");
405 Ok(Some(SchemaFileGuard(schema_file)))
406}
407
408struct SchemaFileGuard(tempfile::NamedTempFile);
409
410fn drive_connection(
415 state: &SessionState,
416 transport: AcpTransport,
417 schema_file_path: Option<String>,
418 model: Option<String>,
419) -> impl std::future::Future<Output = Result<(), agent_client_protocol::Error>> {
420 let acc = state.acc.clone();
421 let events = state.events.clone();
422 let stop_holder = state.stop_holder.clone();
423 let activity_tx = state.activity_tx.clone();
424 let submit_signal = state.submit_signal.clone();
425 let run_id = state.run_id;
426 let agent_id = state.agent_id;
427 let emit_raw = state.emit_raw;
428 let policy = state.policy.clone();
429
430 let acc_for_prompt = acc.clone();
431 let stop_holder_for_prompt = stop_holder.clone();
432 let cwd = state.cwd.clone();
433 let prompt = state.prompt.clone();
434
435 async move {
436 Client
437 .builder()
438 .name("luft")
439 .on_receive_notification(
440 {
441 let acc = acc.clone();
442 let events = events.clone();
443 let activity_tx = activity_tx.clone();
444 let submit_signal = submit_signal.clone();
445 move |n: SessionNotification, _cx: ConnectionTo<Agent>| {
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 async move {
451 handle_session_update(
452 n,
453 &acc,
454 &events,
455 &activity_tx,
456 &submit_signal,
457 run_id,
458 agent_id,
459 emit_raw,
460 );
461 Ok(())
462 }
463 }
464 },
465 agent_client_protocol::on_receive_notification!(),
466 )
467 .on_receive_request(
468 {
469 let policy = policy.clone();
470 move |req: RequestPermissionRequest,
471 responder: Responder<RequestPermissionResponse>,
472 _conn: ConnectionTo<Agent>| {
473 let policy = policy.clone();
474 async move { decide_permission(req, responder, policy).await }
475 }
476 },
477 agent_client_protocol::on_receive_request!(),
478 )
479 .connect_with(transport, move |conn: ConnectionTo<Agent>| {
480 let acc_for_prompt = acc_for_prompt.clone();
481 let stop_holder_for_prompt = stop_holder_for_prompt.clone();
482 let model = model.clone();
483 async move {
484 run_handshake_and_prompt(
485 &conn,
486 &cwd,
487 schema_file_path.as_deref(),
488 model.as_deref(),
489 &prompt,
490 &acc_for_prompt,
491 &stop_holder_for_prompt,
492 )
493 .await
494 }
495 })
496 .await
497 }
498}
499
500#[allow(clippy::too_many_arguments)]
504fn handle_session_update(
505 n: SessionNotification,
506 acc: &Arc<update_mapper::Accumulator>,
507 events: &EventSender,
508 activity_tx: &tokio::sync::mpsc::UnboundedSender<()>,
509 submit_signal: &Arc<tokio::sync::Notify>,
510 run_id: RunId,
511 agent_id: AgentId,
512 emit_raw: bool,
513) {
514 let _ = activity_tx.send(());
515 let kind = serde_json::to_value(&n.update)
516 .ok()
517 .and_then(|v| {
518 v.get("sessionUpdate")
519 .and_then(|v| v.as_str())
520 .map(String::from)
521 })
522 .unwrap_or_else(|| "unknown".to_string());
523 tracing::debug!(%kind, "ACP session/update");
524
525 let was_submitted = acc.structured_output.lock().unwrap().is_some();
529 update_mapper::handle_update(&n.update, run_id, agent_id, acc, events, emit_raw);
530 if !was_submitted && acc.structured_output.lock().unwrap().is_some() {
531 submit_signal.notify_one();
532 tracing::debug!(
533 "ACP structured_output captured; watchdog switching to post-submission mode"
534 );
535 }
536}
537
538async fn decide_permission(
542 req: RequestPermissionRequest,
543 responder: Responder<RequestPermissionResponse>,
544 policy: Option<ToolPolicy>,
545) -> Result<(), agent_client_protocol::Error> {
546 let inputs = permission::extract_inputs(&req);
547 let approve = matches!(
548 permission::decide(policy.as_ref(), &inputs),
549 permission::Decision::Approve
550 );
551 tracing::debug!(
552 approve,
553 options = req.options.len(),
554 "ACP permission request"
555 );
556 let outcome = match (approve, req.options.first()) {
557 (true, Some(opt)) => RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(
558 opt.option_id.clone(),
559 )),
560 _ => RequestPermissionOutcome::Cancelled,
561 };
562 responder.respond(RequestPermissionResponse::new(outcome))
563}
564
565async fn run_handshake_and_prompt(
572 conn: &ConnectionTo<Agent>,
573 cwd: &std::path::Path,
574 schema_file_path: Option<&str>,
575 model: Option<&str>,
576 prompt: &str,
577 acc: &Arc<update_mapper::Accumulator>,
578 stop_holder: &Arc<Mutex<Option<String>>>,
579) -> Result<(), agent_client_protocol::Error> {
580 tracing::debug!("ACP handshake: initialize");
581 conn.send_request(InitializeRequest::new(ProtocolVersion::V1))
582 .block_task()
583 .await?;
584
585 tracing::debug!("ACP handshake: session/new");
586 let ns = session_new(conn, cwd.to_path_buf(), schema_file_path).await?;
587
588 if let Some(model_name) = model {
589 validate_and_set_model(conn, &ns, model_name).await?;
590 }
591
592 tracing::debug!("ACP handshake: session/prompt");
593 let pr = send_prompt(conn, ns.session_id, prompt.to_string()).await?;
594 record_prompt_result(&pr, stop_holder, acc);
595 Ok(())
596}
597
598async fn session_new(
601 conn: &ConnectionTo<Agent>,
602 cwd: PathBuf,
603 schema_file_path: Option<&str>,
604) -> Result<NewSessionResponse, agent_client_protocol::Error> {
605 let req = NewSessionRequest::new(cwd);
606 let req = match schema_file_path {
607 Some(sf) => {
608 let luft_bin =
609 std::env::current_exe().unwrap_or_else(|_| std::path::PathBuf::from("luft"));
610 let mcp = McpServerStdio::new("luft-structured-output", luft_bin).args(vec![
611 "mcp-structured-output".to_string(),
612 "--schema-file".to_string(),
613 sf.to_string(),
614 ]);
615 req.mcp_servers(vec![McpServer::Stdio(mcp)])
616 }
617 None => req,
618 };
619 conn.send_request(req).block_task().await
620}
621
622async fn validate_and_set_model(
627 conn: &ConnectionTo<Agent>,
628 ns: &NewSessionResponse,
629 model_name: &str,
630) -> Result<(), agent_client_protocol::Error> {
631 let config_options = match ns.config_options.as_ref() {
632 Some(opts) => opts,
633 None => {
634 tracing::debug!("ACP: agent does not advertise config_options");
635 return Ok(());
636 }
637 };
638 let model_option = match config_options
639 .iter()
640 .find(|opt| opt.category.as_ref() == Some(&SessionConfigOptionCategory::Model))
641 {
642 Some(o) => o,
643 None => {
644 tracing::debug!("ACP: agent does not support model selection");
645 return Ok(());
646 }
647 };
648 let select = match &model_option.kind {
649 SessionConfigKind::Select(s) => s,
650 _ => {
651 tracing::debug!("ACP: model option is not a Select kind");
652 return Ok(());
653 }
654 };
655 let valid = match &select.options {
656 SessionConfigSelectOptions::Ungrouped(opts) => {
657 opts.iter().any(|o| o.value.0.as_ref() == model_name)
658 }
659 SessionConfigSelectOptions::Grouped(groups) => groups
660 .iter()
661 .any(|g| g.options.iter().any(|o| o.value.0.as_ref() == model_name)),
662 _ => false,
663 };
664 if valid {
665 tracing::debug!(model = %model_name, "ACP: setting session model");
666 let req = SetSessionConfigOptionRequest::new(
667 ns.session_id.clone(),
668 model_option.id.clone(),
669 model_name.to_string(),
670 );
671 conn.send_request(req).block_task().await?;
672 } else {
673 tracing::warn!(
674 model = %model_name,
675 "ACP: requested model not available, using agent default"
676 );
677 }
678 Ok(())
679}
680
681async fn send_prompt(
685 conn: &ConnectionTo<Agent>,
686 session_id: SessionId,
687 prompt: String,
688) -> Result<agent_client_protocol::schema::PromptResponse, agent_client_protocol::Error> {
689 conn.send_request(PromptRequest::new(
690 session_id,
691 vec![ContentBlock::Text(TextContent::new(prompt))],
692 ))
693 .block_task()
694 .await
695}
696
697fn record_prompt_result(
701 pr: &agent_client_protocol::schema::PromptResponse,
702 stop_holder: &Arc<Mutex<Option<String>>>,
703 #[cfg_attr(
704 not(feature = "unstable_end_turn_token_usage"),
705 allow(unused_variables)
706 )]
707 acc: &Arc<update_mapper::Accumulator>,
708) {
709 tracing::debug!(stop_reason = ?pr.stop_reason, "ACP prompt complete");
710 *stop_holder.lock().unwrap() = Some(stop_reason_as_str(&pr.stop_reason));
711 #[cfg(feature = "unstable_end_turn_token_usage")]
712 {
713 if let Some(u) = pr.usage.as_ref() {
714 tracing::debug!(
715 input = u.input_tokens,
716 output = u.output_tokens,
717 total = u.total_tokens,
718 "ACP prompt usage"
719 );
720 *acc.tokens.lock().unwrap() = TokenUsage {
721 input: u.input_tokens,
722 output: u.output_tokens,
723 cache_read: u.cached_read_tokens.unwrap_or(0),
724 cache_write: u.cached_write_tokens.unwrap_or(0),
725 };
726 }
727 }
728}
729
730fn stop_reason_as_str(r: &StopReason) -> String {
736 match r {
737 StopReason::EndTurn => STOP_REASON_END_TURN.to_string(),
738 StopReason::MaxTokens => "MaxTokens".to_string(),
739 StopReason::MaxTurnRequests => "MaxTurnRequests".to_string(),
740 StopReason::Refusal => "Refusal".to_string(),
741 StopReason::Cancelled => "Cancelled".to_string(),
742 #[allow(unreachable_patterns)]
743 other => format!("{other:?}"),
744 }
745}
746
747fn handle_watchdog_outcome(
753 res: WatchdogOutcome,
754 child: &mut tokio::process::Child,
755 stop_holder: &Arc<Mutex<Option<String>>>,
756 idle_timeout: Duration,
757) -> Result<(), BackendError> {
758 let _ = child.start_kill();
759 match res {
760 WatchdogOutcome::PreIdleTimeout => {
761 tracing::warn!(
762 idle_timeout_ms = idle_timeout.as_millis() as u64,
763 "ACP session idle timeout (no protocol activity)"
764 );
765 Err(BackendError::Timeout)
766 }
767 WatchdogOutcome::ChannelClosed => {
768 tracing::debug!("ACP activity channel closed");
769 Err(BackendError::Timeout)
770 }
771 WatchdogOutcome::PostSubmissionTimeout => {
772 tracing::info!(
779 post_idle_ms = POST_SUBMISSION_IDLE.as_millis() as u64,
780 "ACP post-submission timeout; treating structured_output as result"
781 );
782 let mut guard = stop_holder.lock().unwrap();
787 if guard.is_none() {
788 *guard = Some(STOP_REASON_END_TURN.to_string());
789 }
790 Ok(())
791 }
792 }
793}
794
795fn classify_protocol_error(e: agent_client_protocol::Error) -> BackendError {
805 let s = e.to_string();
806 if is_connection_closed(&s) {
807 tracing::warn!("ACP connection closed");
808 BackendError::Protocol("connection closed".into())
809 } else {
810 tracing::error!(error = %s, "ACP protocol error");
811 BackendError::Protocol(s)
812 }
813}
814
815fn is_connection_closed(s: &str) -> bool {
816 s.contains("receiver dropped")
820 || s.contains("broken pipe")
821 || s.contains("unexpected eof")
822 || s.contains("connection closed")
823}
824
825fn collect_session_result(task: &AgentTask, state: &SessionState) -> AgentResult {
827 let stop = state.stop_holder.lock().unwrap().take().unwrap_or_default();
828 let message = std::mem::take(&mut *state.acc.message.lock().unwrap());
829 let tokens = *state.acc.tokens.lock().unwrap();
830 let structured = state.acc.structured_output.lock().unwrap().take();
831 result_collector::collect(task, &stop, message, tokens, structured)
832}
833
834async fn idle_watchdog(
857 pre_idle: Duration,
858 post_idle: Duration,
859 activity_rx: &mut tokio::sync::mpsc::UnboundedReceiver<()>,
860 submit_signal: Arc<tokio::sync::Notify>,
861) -> WatchdogOutcome {
862 let mut submitted = false;
863 loop {
864 if submitted {
865 while activity_rx.try_recv().is_ok() {}
874 tokio::time::sleep(post_idle).await;
875 return WatchdogOutcome::PostSubmissionTimeout;
876 }
877 tokio::select! {
878 biased;
879 _ = submit_signal.notified() => {
880 submitted = true;
881 tracing::debug!(
882 post_idle_ms = post_idle.as_millis() as u64,
883 "ACP watchdog entered post-submission mode"
884 );
885 }
886 msg = activity_rx.recv() => match msg {
887 Some(()) => { while activity_rx.try_recv().is_ok() {} }
888 None => return WatchdogOutcome::ChannelClosed,
889 },
890 _ = tokio::time::sleep(pre_idle) => {
891 return WatchdogOutcome::PreIdleTimeout;
892 }
893 }
894 }
895}
896
897#[cfg(test)]
898mod tests {
899 use super::*;
900
901 #[tokio::test]
908 async fn idle_watchdog_fires_after_idle_period() {
909 let (_atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
910 let submit = Arc::new(tokio::sync::Notify::new());
911 let r = tokio::time::timeout(
912 Duration::from_millis(500),
913 idle_watchdog(
914 Duration::from_millis(50),
915 Duration::from_millis(50),
916 &mut arx,
917 submit,
918 ),
919 )
920 .await;
921 let outcome = r.expect("should fire after idle period");
922 assert_eq!(outcome, WatchdogOutcome::PreIdleTimeout);
923 }
924
925 #[tokio::test]
926 async fn idle_watchdog_does_not_fire_with_activity() {
927 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
928 let submit = Arc::new(tokio::sync::Notify::new());
929 tokio::spawn(async move {
930 for _ in 0..5 {
931 tokio::time::sleep(Duration::from_millis(20)).await;
932 let _ = atx.send(());
933 }
934 });
935 let r = tokio::time::timeout(
936 Duration::from_millis(80),
937 idle_watchdog(
938 Duration::from_millis(50),
939 Duration::from_millis(50),
940 &mut arx,
941 submit,
942 ),
943 )
944 .await;
945 assert!(
946 r.is_err(),
947 "should not fire while activity is within idle window"
948 );
949 }
950
951 #[tokio::test]
952 async fn idle_watchdog_fires_after_activity_stops() {
953 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
954 let submit = Arc::new(tokio::sync::Notify::new());
955 let _ = atx.send(());
956 drop(atx);
957 let r = tokio::time::timeout(
958 Duration::from_millis(30),
959 idle_watchdog(
960 Duration::from_millis(80),
961 Duration::from_millis(80),
962 &mut arx,
963 submit,
964 ),
965 )
966 .await;
967 let outcome = r.expect("should return immediately when channel closes");
968 assert_eq!(outcome, WatchdogOutcome::ChannelClosed);
969 }
970
971 #[tokio::test]
983 async fn idle_watchdog_enters_post_mode_after_submit_signal() {
984 let (_atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
985 let submit = Arc::new(tokio::sync::Notify::new());
986 let submit_h = submit.clone();
987 tokio::spawn(async move {
988 tokio::time::sleep(Duration::from_millis(20)).await;
989 submit_h.notify_one();
990 });
991 let r = tokio::time::timeout(
992 Duration::from_millis(500),
993 idle_watchdog(
994 Duration::from_secs(60), Duration::from_millis(50), &mut arx,
997 submit,
998 ),
999 )
1000 .await;
1001 let outcome = r.expect("watchdog should return after post_idle");
1002 assert_eq!(outcome, WatchdogOutcome::PostSubmissionTimeout);
1003 }
1004
1005 #[tokio::test]
1006 async fn idle_watchdog_post_mode_is_not_reset_by_activity() {
1007 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
1011 let submit = Arc::new(tokio::sync::Notify::new());
1012 submit.notify_one();
1013 tokio::spawn(async move {
1014 for _ in 0..20 {
1015 tokio::time::sleep(Duration::from_millis(20)).await;
1016 let _ = atx.send(());
1017 }
1018 });
1019 let start = std::time::Instant::now();
1020 let r = tokio::time::timeout(
1021 Duration::from_millis(500),
1022 idle_watchdog(
1023 Duration::from_secs(60),
1024 Duration::from_millis(80),
1025 &mut arx,
1026 submit,
1027 ),
1028 )
1029 .await;
1030 let outcome = r.expect("watchdog should return after post_idle");
1031 let elapsed = start.elapsed();
1032 assert_eq!(outcome, WatchdogOutcome::PostSubmissionTimeout);
1033 assert!(
1037 elapsed < Duration::from_millis(300),
1038 "post-mode timer was reset by activity: elapsed={elapsed:?}"
1039 );
1040 }
1041
1042 #[tokio::test]
1043 async fn idle_watchdog_pre_mode_resets_on_activity() {
1044 let (atx, mut arx) = tokio::sync::mpsc::unbounded_channel::<()>();
1046 let submit = Arc::new(tokio::sync::Notify::new());
1047 tokio::spawn(async move {
1048 for _ in 0..10 {
1049 tokio::time::sleep(Duration::from_millis(30)).await;
1050 let _ = atx.send(());
1051 }
1052 });
1053 let r = tokio::time::timeout(
1054 Duration::from_millis(200),
1055 idle_watchdog(
1056 Duration::from_millis(60),
1057 Duration::from_millis(60),
1058 &mut arx,
1059 submit,
1060 ),
1061 )
1062 .await;
1063 assert!(
1064 r.is_err(),
1065 "pre-mode should not fire while activity keeps resetting timer"
1066 );
1067 }
1068
1069 #[test]
1077 fn stop_reason_as_str_end_turn_matches_constant() {
1078 assert_eq!(
1079 stop_reason_as_str(&StopReason::EndTurn),
1080 STOP_REASON_END_TURN
1081 );
1082 assert_eq!(stop_reason_as_str(&StopReason::EndTurn), "EndTurn");
1083 }
1084
1085 #[test]
1086 fn stop_reason_as_str_cancelled_contains_cancel() {
1087 assert_eq!(stop_reason_as_str(&StopReason::Cancelled), "Cancelled");
1088 }
1089
1090 #[test]
1091 fn stop_reason_as_str_other_variants_stable() {
1092 assert_eq!(stop_reason_as_str(&StopReason::MaxTokens), "MaxTokens");
1093 assert_eq!(
1094 stop_reason_as_str(&StopReason::MaxTurnRequests),
1095 "MaxTurnRequests"
1096 );
1097 assert_eq!(stop_reason_as_str(&StopReason::Refusal), "Refusal");
1098 }
1099
1100 #[tokio::test]
1106 async fn handle_watchdog_post_submission_synthesizes_end_turn() {
1107 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1108 let mut child = match tokio::process::Command::new("cmd")
1111 .arg("/C")
1112 .arg("exit 0")
1113 .stdin(Stdio::null())
1114 .stdout(Stdio::null())
1115 .stderr(Stdio::null())
1116 .kill_on_drop(true)
1117 .spawn()
1118 {
1119 Ok(c) => c,
1120 Err(_) => return, };
1122 let r = handle_watchdog_outcome(
1123 WatchdogOutcome::PostSubmissionTimeout,
1124 &mut child,
1125 &stop,
1126 Duration::from_secs(300),
1127 );
1128 assert!(
1129 r.is_ok(),
1130 "post-submission outcome should fall through to collect"
1131 );
1132 assert_eq!(stop.lock().unwrap().as_deref(), Some("EndTurn"));
1133 }
1134
1135 #[tokio::test]
1136 async fn handle_watchdog_post_submission_preserves_existing_stop() {
1137 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(Some("Cancelled".into())));
1138 let mut child = match tokio::process::Command::new("cmd")
1139 .arg("/C")
1140 .arg("exit 0")
1141 .stdin(Stdio::null())
1142 .stdout(Stdio::null())
1143 .stderr(Stdio::null())
1144 .kill_on_drop(true)
1145 .spawn()
1146 {
1147 Ok(c) => c,
1148 Err(_) => return,
1149 };
1150 let r = handle_watchdog_outcome(
1151 WatchdogOutcome::PostSubmissionTimeout,
1152 &mut child,
1153 &stop,
1154 Duration::from_secs(300),
1155 );
1156 assert!(r.is_ok());
1157 assert_eq!(stop.lock().unwrap().as_deref(), Some("Cancelled"));
1159 }
1160
1161 #[tokio::test]
1162 async fn handle_watchdog_pre_idle_returns_timeout_error() {
1163 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1164 let mut child = match tokio::process::Command::new("cmd")
1165 .arg("/C")
1166 .arg("exit 0")
1167 .stdin(Stdio::null())
1168 .stdout(Stdio::null())
1169 .stderr(Stdio::null())
1170 .kill_on_drop(true)
1171 .spawn()
1172 {
1173 Ok(c) => c,
1174 Err(_) => return,
1175 };
1176 let r = handle_watchdog_outcome(
1177 WatchdogOutcome::PreIdleTimeout,
1178 &mut child,
1179 &stop,
1180 Duration::from_secs(1),
1181 );
1182 assert!(matches!(r, Err(BackendError::Timeout)));
1183 }
1184
1185 #[tokio::test]
1186 async fn handle_watchdog_channel_closed_returns_timeout_error() {
1187 let stop: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
1188 let mut child = match tokio::process::Command::new("cmd")
1189 .arg("/C")
1190 .arg("exit 0")
1191 .stdin(Stdio::null())
1192 .stdout(Stdio::null())
1193 .stderr(Stdio::null())
1194 .kill_on_drop(true)
1195 .spawn()
1196 {
1197 Ok(c) => c,
1198 Err(_) => return,
1199 };
1200 let r = handle_watchdog_outcome(
1201 WatchdogOutcome::ChannelClosed,
1202 &mut child,
1203 &stop,
1204 Duration::from_secs(1),
1205 );
1206 assert!(matches!(r, Err(BackendError::Timeout)));
1207 }
1208
1209 #[test]
1212 fn is_connection_closed_matches_documented_substrings() {
1213 assert!(is_connection_closed("receiver dropped"));
1214 assert!(is_connection_closed("broken pipe"));
1215 assert!(is_connection_closed("unexpected eof"));
1216 assert!(is_connection_closed("connection closed"));
1217 assert!(is_connection_closed(
1218 "io error: broken pipe writing to stdin"
1219 ));
1220 assert!(!is_connection_closed("unknown protocol method"));
1221 assert!(!is_connection_closed(""));
1222 }
1223}