1use std::collections::{HashMap, VecDeque};
12use std::sync::Arc;
13use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
14
15use arc_swap::ArcSwap;
16use tokio::sync::mpsc;
17
18use agent_client_protocol::Responder;
19use agent_client_protocol::schema::v1::{
20 PermissionOption, PermissionOptionId, PermissionOptionKind, RequestPermissionOutcome,
21 RequestPermissionRequest, RequestPermissionResponse, SelectedPermissionOutcome, SessionUpdate,
22 ToolCallContent, ToolCallStatus, ToolKind,
23};
24use lattice_agent::{
25 AiLogLevel, AiLogSource, AiLogger, DiffReviewRequest, SessionKey, review_diff,
26};
27use lattice_diff::ProgrammaticDiffBus;
28use lattice_diff::subsystem::DiffOutcome;
29
30use crate::Result;
31use crate::acp::connection::{Connection, PermissionRequest, SessionId, SessionNotification};
32use crate::acp::conversation::{ConversationStore, SessionStatus};
33use crate::acp::error::AiError;
34use crate::acp::handle::{AiClientHandle, AiCmd, AiState};
35use crate::acp::providers::ProviderConfig;
36
37impl AiClientHandle {
38 pub fn spawn(
46 runtime: &tokio::runtime::Handle,
47 logger: AiLogger,
48 conv_store: ConversationStore,
49 diff_bus: Option<ProgrammaticDiffBus>,
50 ) -> AiClientHandle {
51 let (cmd_tx, cmd_rx) = mpsc::unbounded_channel::<AiCmd>();
52 let state = Arc::new(ArcSwap::from_pointee(AiState::default()));
53 let queue_len = Arc::new(AtomicUsize::new(0));
54 let ql = queue_len.clone();
55 runtime.spawn(supervisor_loop(
56 cmd_rx,
57 state.clone(),
58 ql,
59 logger,
60 conv_store,
61 diff_bus,
62 ));
63 AiClientHandle {
64 cmd_tx,
65 state,
66 queue_len,
67 }
68 }
69}
70
71enum SupervisorEvent {
74 Cmd(AiCmd),
75 ChildExited,
79}
80
81async fn next_event(
91 cmd_rx: &mut mpsc::UnboundedReceiver<AiCmd>,
92 child: &mut Option<tokio::process::Child>,
93) -> Option<SupervisorEvent> {
94 match child.as_mut() {
95 Some(child) => {
96 tokio::select! {
97 cmd = cmd_rx.recv() => cmd.map(SupervisorEvent::Cmd),
98 _ = child.wait() => Some(SupervisorEvent::ChildExited),
99 }
100 }
101 None => cmd_rx.recv().await.map(SupervisorEvent::Cmd),
102 }
103}
104
105fn kill_child(mut child: tokio::process::Child) {
111 let _ = child.start_kill();
112 tokio::spawn(async move {
113 let _ = child.wait().await;
114 });
115}
116
117async fn supervisor_loop(
124 mut cmd_rx: mpsc::UnboundedReceiver<AiCmd>,
125 state: Arc<ArcSwap<AiState>>,
126 queue_len: Arc<AtomicUsize>,
127 logger: AiLogger,
128 conv_store: ConversationStore,
129 diff_bus: Option<ProgrammaticDiffBus>,
130) {
131 let mut conn: Option<Arc<Connection>> = None;
132 let mut sess: Option<SessionId> = None;
133 let mut active_key: Option<SessionKey> = None;
134 let mut child: Option<tokio::process::Child> = None;
139 let mut indices: HashMap<&'static str, u32> = HashMap::new();
143 let auto_accept = Arc::new(AtomicBool::new(false));
147 let mut prompt_queue: VecDeque<String> = VecDeque::new();
149 let mut prompt_in_flight = false;
150 let (prompt_done_tx, mut prompt_done_rx) = mpsc::unbounded_channel::<()>();
153
154 while let Some(event) = next_event(&mut cmd_rx, &mut child).await {
155 match event {
156 SupervisorEvent::ChildExited => {
163 prompt_queue.clear();
165 prompt_in_flight = false;
166 queue_len.store(0, Ordering::Relaxed);
167 logger.log(
168 active_key.as_ref(),
169 AiLogLevel::Warn,
170 AiLogSource::Lifecycle,
171 "agent exited",
172 );
173 child = None;
174 conn = None;
175 sess = None;
176 active_key = None;
177 state.store(Arc::new(AiState::default()));
178 }
179 SupervisorEvent::Cmd(AiCmd::Start(provider)) => {
180 prompt_queue.clear();
182 prompt_in_flight = false;
183 queue_len.store(0, Ordering::Relaxed);
184 if let Some(c) = child.take() {
197 kill_child(c);
198 }
199 conn = None;
200 sess = None;
201 active_key = None;
202 auto_accept.store(false, Ordering::Relaxed);
205 state.store(Arc::new(AiState::default()));
206
207 let idx = indices.entry(provider.display_name).or_insert(0);
208 *idx += 1;
209 let key = SessionKey::new(provider.display_name, *idx);
210
211 logger.log(
212 Some(&key),
213 AiLogLevel::Info,
214 AiLogSource::Lifecycle,
215 format!("starting {}", provider.display_name),
216 );
217
218 match start_provider(
219 &provider,
220 key.clone(),
221 conv_store.clone(),
222 diff_bus.clone(),
223 auto_accept.clone(),
224 logger.clone(),
225 )
226 .await
227 {
228 Ok((new_conn, new_sess, new_child)) => {
229 conn = Some(new_conn);
230 sess = Some(new_sess);
231 child = Some(new_child);
232 active_key = Some(key.clone());
233 state.store(Arc::new(AiState {
234 running: true,
235 provider: Some(provider.display_name),
236 session: Some(key.clone()),
237 auto_accept: false,
238 queue_len: 0,
239 }));
240 logger.log(
241 Some(&key),
242 AiLogLevel::Info,
243 AiLogSource::Lifecycle,
244 "session opened",
245 );
246 }
247 Err(e) => {
248 logger.log(
249 Some(&key),
250 AiLogLevel::Error,
251 AiLogSource::Lifecycle,
252 format!("start failed: {e}"),
253 );
254 }
255 }
256 }
257 SupervisorEvent::Cmd(AiCmd::Prompt(text)) => {
258 if let (Some(c), Some(s)) = (conn.clone(), sess.clone()) {
259 if let Some(key) = active_key.as_ref() {
263 conv_store.push_user_text(key, &text);
264 }
265 if prompt_in_flight {
266 let ql = queue_len.fetch_add(1, Ordering::Relaxed) + 1;
268 prompt_queue.push_back(text);
269 logger.log(
270 active_key.as_ref(),
271 AiLogLevel::Info,
272 AiLogSource::Lifecycle,
273 format!("prompt queued ({} pending)", ql),
274 );
275 let mut next = (**state.load()).clone();
277 next.queue_len = ql;
278 state.store(Arc::new(next));
279 } else {
280 prompt_in_flight = true;
282 let done_tx = prompt_done_tx.clone();
283 tokio::spawn(async move {
284 let _ = crate::acp::session::prompt(&c, &s, &text).await;
285 let _ = done_tx.send(());
287 });
288 }
289 } else {
290 logger.log(
291 None,
292 AiLogLevel::Warn,
293 AiLogSource::Lifecycle,
294 "prompt dropped: no active session",
295 );
296 }
297 }
298 SupervisorEvent::Cmd(AiCmd::SetAutoAccept(on)) => {
299 auto_accept.store(on, Ordering::Relaxed);
303 let mut next = (**state.load()).clone();
304 next.auto_accept = on;
305 next.queue_len = queue_len.load(Ordering::Relaxed);
306 state.store(Arc::new(next));
307 logger.log(
308 active_key.as_ref(),
309 AiLogLevel::Info,
310 AiLogSource::Lifecycle,
311 if on {
312 "trust mode on (auto-accept)"
313 } else {
314 "review mode"
315 },
316 );
317 }
318 SupervisorEvent::Cmd(AiCmd::Interrupt) => {
319 if let (Some(c), Some(s)) = (conn.clone(), sess.clone()) {
324 let key = active_key.clone();
325 let logger = logger.clone();
326 tokio::spawn(async move {
327 if let Err(e) = c.cancel(&s).await {
328 logger.log(
329 key.as_ref(),
330 AiLogLevel::Warn,
331 AiLogSource::Lifecycle,
332 format!("interrupt failed: {e}"),
333 );
334 }
335 });
336 } else {
337 logger.log(
338 None,
339 AiLogLevel::Warn,
340 AiLogSource::Lifecycle,
341 "interrupt dropped: no active session",
342 );
343 }
344 }
345 SupervisorEvent::Cmd(AiCmd::Stop) => {
346 prompt_queue.clear();
348 prompt_in_flight = false;
349 queue_len.store(0, Ordering::Relaxed);
350 logger.log(
351 active_key.as_ref(),
352 AiLogLevel::Info,
353 AiLogSource::Lifecycle,
354 "stopped",
355 );
356 if let Some(c) = child.take() {
357 kill_child(c);
358 }
359 conn = None;
360 sess = None;
361 active_key = None;
362 state.store(Arc::new(AiState::default()));
363 }
364 }
365 while prompt_done_rx.try_recv().is_ok() {
367 prompt_in_flight = false;
368 let next_text = prompt_queue.pop_front();
369 if let Some(text) = next_text {
370 let ql = queue_len.fetch_sub(1, Ordering::Relaxed).saturating_sub(1);
371 prompt_in_flight = true;
372 if let (Some(c), Some(s)) = (conn.clone(), sess.clone()) {
373 let done_tx = prompt_done_tx.clone();
374 tokio::spawn(async move {
375 let _ = crate::acp::session::prompt(&c, &s, &text).await;
376 let _ = done_tx.send(());
377 });
378 logger.log(
379 active_key.as_ref(),
380 AiLogLevel::Info,
381 AiLogSource::Lifecycle,
382 format!("dequeued prompt ({} remaining)", ql),
383 );
384 }
385 let mut next = (**state.load()).clone();
387 next.queue_len = ql;
388 state.store(Arc::new(next));
389 } else {
390 let mut next = (**state.load()).clone();
392 next.queue_len = 0;
393 state.store(Arc::new(next));
394 }
395 }
396 }
397}
398
399pub(crate) async fn drain_notifications(
406 mut rx: mpsc::UnboundedReceiver<SessionNotification>,
407 conv_store: ConversationStore,
408 session: SessionKey,
409) {
410 let mut tool_names: std::collections::HashMap<String, String> =
414 std::collections::HashMap::new();
415 while let Some(notification) = rx.recv().await {
416 conv_store.apply(&session, ¬ification.update);
417 match ¬ification.update {
419 SessionUpdate::AgentMessageChunk(_) | SessionUpdate::AgentThoughtChunk(_) => {
420 conv_store.set_status(&session, SessionStatus::Thinking);
421 }
422 SessionUpdate::ToolCall(tc) => {
423 let title = tc.title.clone();
424 tool_names.insert(tc.tool_call_id.0.to_string(), title.clone());
425 conv_store.set_status(&session, SessionStatus::Executing { tool: title });
426 }
427 SessionUpdate::ToolCallUpdate(u) => {
428 let tid = u.tool_call_id.0.to_string();
429 if let Some(ToolCallStatus::InProgress) | Some(ToolCallStatus::Pending) =
430 u.fields.status
431 {
432 let title = tool_names
434 .get(&tid)
435 .cloned()
436 .unwrap_or_else(|| "tool".to_string());
437 conv_store.set_status(&session, SessionStatus::Executing { tool: title });
438 } else if matches!(
439 u.fields.status,
440 Some(ToolCallStatus::Completed) | Some(ToolCallStatus::Failed)
441 ) {
442 tool_names.remove(&tid);
443 conv_store.set_status(&session, SessionStatus::Idle);
444 }
445 }
446 _ => {}
447 }
448 }
449 conv_store.set_status(&session, SessionStatus::Idle);
451}
452
453async fn drain_permissions(
459 mut rx: mpsc::UnboundedReceiver<PermissionRequest>,
460 conv_store: ConversationStore,
461 session: SessionKey,
462 diff_bus: Option<ProgrammaticDiffBus>,
463 origin_session: u64,
464 auto_accept: Arc<AtomicBool>,
465) {
466 while let Some(pr) = rx.recv().await {
467 tokio::spawn(handle_permission(
468 pr,
469 conv_store.clone(),
470 session.clone(),
471 diff_bus.clone(),
472 origin_session,
473 auto_accept.clone(),
474 ));
475 }
476}
477
478async fn handle_permission(
485 pr: PermissionRequest,
486 conv_store: ConversationStore,
487 session: SessionKey,
488 diff_bus: Option<ProgrammaticDiffBus>,
489 origin_session: u64,
490 auto_accept: Arc<AtomicBool>,
491) {
492 let PermissionRequest { request, responder } = pr;
493 let trusted = auto_accept.load(Ordering::Relaxed);
496 match resolve_decision(trusted, &request, origin_session) {
497 PermissionDecision::AutoAllow => {
498 respond(responder, allow_outcome(&request.options));
499 }
500 PermissionDecision::Deny => {
501 respond(responder, deny_outcome(&request.options));
505 }
506 PermissionDecision::AskUser => {
507 conv_store.set_status(&session, SessionStatus::AwaitingPermission);
510 let id = request.tool_call.tool_call_id.0.to_string();
512 let title = request
513 .tool_call
514 .fields
515 .title
516 .clone()
517 .unwrap_or_else(|| "agent action".to_string());
518 let (tx, rx) = tokio::sync::oneshot::channel();
519 conv_store.push_permission_request(
520 &session,
521 id,
522 title,
523 None,
524 request.options.clone(),
525 tx,
526 );
527 match rx.await {
532 Ok(option_id) => {
533 respond(
534 responder,
535 RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(
536 option_id,
537 )),
538 );
539 }
540 Err(_) => respond(responder, RequestPermissionOutcome::Cancelled),
541 }
542 conv_store.set_status(&session, SessionStatus::Thinking);
545 }
546 PermissionDecision::Review(review) => {
547 let Some(bus) = diff_bus else {
548 tracing::debug!("ACP edit permission denied: no programmatic diff bus");
551 respond(responder, deny_outcome(&request.options));
552 return;
553 };
554 let outcome = match review_diff(&bus, review).await {
555 Ok(DiffOutcome::Accept) => allow_outcome(&request.options),
556 Ok(DiffOutcome::Reject) => deny_outcome(&request.options),
557 _ => RequestPermissionOutcome::Cancelled,
560 };
561 respond(responder, outcome);
562 }
563 }
564}
565
566enum PermissionDecision {
568 AutoAllow,
570 Review(DiffReviewRequest),
572 AskUser,
575 Deny,
579}
580
581fn resolve_decision(
585 trusted: bool,
586 request: &RequestPermissionRequest,
587 origin_session: u64,
588) -> PermissionDecision {
589 if trusted {
590 PermissionDecision::AutoAllow
591 } else {
592 classify_permission(request, origin_session)
593 }
594}
595
596fn classify_permission(
610 request: &RequestPermissionRequest,
611 origin_session: u64,
612) -> PermissionDecision {
613 let kind = request.tool_call.fields.kind.unwrap_or_default();
614 let read_only = matches!(
615 kind,
616 ToolKind::Read
617 | ToolKind::Search
618 | ToolKind::Fetch
619 | ToolKind::Think
620 | ToolKind::SwitchMode
621 );
622 if read_only {
623 return PermissionDecision::AutoAllow;
624 }
625 match diff_review_from(request, origin_session) {
626 Some(review) => PermissionDecision::Review(review),
627 None => match kind {
628 ToolKind::Edit => PermissionDecision::Deny,
630 _ => PermissionDecision::AskUser,
632 },
633 }
634}
635
636fn diff_review_from(
639 request: &RequestPermissionRequest,
640 origin_session: u64,
641) -> Option<DiffReviewRequest> {
642 let title = request
643 .tool_call
644 .fields
645 .title
646 .clone()
647 .unwrap_or_else(|| "agent edit".to_string());
648 request
649 .tool_call
650 .fields
651 .content
652 .as_ref()?
653 .iter()
654 .find_map(|content| match content {
655 ToolCallContent::Diff(diff) => Some(DiffReviewRequest {
656 old_file_path: diff.path.clone(),
657 new_file_path: diff.path.clone(),
658 new_contents: diff.new_text.clone(),
659 tab_name: title.clone(),
660 origin_session,
661 }),
662 _ => None,
663 })
664}
665
666fn pick_option(options: &[PermissionOption], allow: bool) -> Option<PermissionOptionId> {
669 let (once, always) = if allow {
670 (
671 PermissionOptionKind::AllowOnce,
672 PermissionOptionKind::AllowAlways,
673 )
674 } else {
675 (
676 PermissionOptionKind::RejectOnce,
677 PermissionOptionKind::RejectAlways,
678 )
679 };
680 options
681 .iter()
682 .find(|o| o.kind == once)
683 .or_else(|| options.iter().find(|o| o.kind == always))
684 .map(|o| o.option_id.clone())
685}
686
687fn allow_outcome(options: &[PermissionOption]) -> RequestPermissionOutcome {
689 match pick_option(options, true) {
690 Some(id) => RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(id)),
691 None => RequestPermissionOutcome::Cancelled,
692 }
693}
694
695fn deny_outcome(options: &[PermissionOption]) -> RequestPermissionOutcome {
697 match pick_option(options, false) {
698 Some(id) => RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(id)),
699 None => RequestPermissionOutcome::Cancelled,
700 }
701}
702
703fn respond(responder: Responder<RequestPermissionResponse>, outcome: RequestPermissionOutcome) {
705 let _ = responder.respond(RequestPermissionResponse::new(outcome));
706}
707
708async fn start_provider(
712 provider: &ProviderConfig,
713 session: SessionKey,
714 conv_store: ConversationStore,
715 diff_bus: Option<ProgrammaticDiffBus>,
716 auto_accept: Arc<AtomicBool>,
717 logger: AiLogger,
718) -> Result<(Arc<Connection>, SessionId, tokio::process::Child)> {
719 let mut child = tokio::process::Command::new(&provider.command)
720 .args(&provider.args)
721 .envs(provider.env.iter().cloned())
722 .stdin(std::process::Stdio::piped())
723 .stdout(std::process::Stdio::piped())
724 .stderr(std::process::Stdio::piped())
728 .kill_on_drop(true)
734 .spawn()
735 .map_err(|e| AiError::Process(e.to_string()))?;
736
737 let stdin = child
738 .stdin
739 .take()
740 .ok_or_else(|| AiError::Process("no stdin".to_string()))?;
741 let stdout = child
742 .stdout
743 .take()
744 .ok_or_else(|| AiError::Process("no stdout".to_string()))?;
745
746 if let Some(stderr) = child.stderr.take() {
750 let logger = logger.clone();
751 let key = session.clone();
752 tokio::spawn(async move {
753 use tokio::io::AsyncBufReadExt;
754 let mut lines = tokio::io::BufReader::new(stderr).lines();
755 while let Ok(Some(line)) = lines.next_line().await {
756 let line = line.trim();
757 if line.is_empty() {
758 continue;
759 }
760 logger.log(
761 Some(&key),
762 AiLogLevel::Warn,
763 AiLogSource::Client,
764 line.to_string(),
765 );
766 }
767 });
768 }
769
770 let (conn, notif_rx, perm_rx) = Connection::spawn(stdout, stdin);
771 let origin_session = u64::from(session.index);
775 let session_for_perms = session.clone();
776 tokio::spawn(drain_notifications(notif_rx, conv_store.clone(), session));
777 tokio::spawn(drain_permissions(
778 perm_rx,
779 conv_store,
780 session_for_perms,
781 diff_bus,
782 origin_session,
783 auto_accept,
784 ));
785
786 let cwd = lattice_core::project::root_from_cwd()
795 .map(|p| p.display().to_string())
796 .unwrap_or_default();
797 let session_id = crate::acp::session::handshake(&conn, &cwd).await?;
798
799 Ok((conn, session_id, child))
800}
801
802#[cfg(test)]
803mod tests {
804 use std::time::Duration;
805
806 use std::sync::Mutex;
807
808 use agent_client_protocol::schema::v1::{
809 ContentBlock, ContentChunk, SessionId as AcpSessionId, SessionUpdate, TextContent,
810 };
811 use tokio::sync::mpsc;
812
813 use super::*;
814 use agent_client_protocol::schema::v1::{
815 Diff, ToolCallContent, ToolCallId, ToolCallUpdate, ToolCallUpdateFields, ToolKind,
816 };
817
818 fn perm_options() -> Vec<PermissionOption> {
821 vec![
822 PermissionOption::new("allow-once", "Allow", PermissionOptionKind::AllowOnce),
823 PermissionOption::new("allow-always", "Always", PermissionOptionKind::AllowAlways),
824 PermissionOption::new("reject-once", "Reject", PermissionOptionKind::RejectOnce),
825 ]
826 }
827
828 fn perm_req(kind: ToolKind, content: Option<Vec<ToolCallContent>>) -> RequestPermissionRequest {
829 let fields = ToolCallUpdateFields::new()
830 .kind(Some(kind))
831 .title(Some("edit parse.rs".to_string()))
832 .content(content);
833 RequestPermissionRequest::new(
834 "s",
835 ToolCallUpdate::new(ToolCallId::new("t1"), fields),
836 perm_options(),
837 )
838 }
839
840 #[test]
841 fn read_class_kinds_auto_allow() {
842 for kind in [
843 ToolKind::Read,
844 ToolKind::Search,
845 ToolKind::Fetch,
846 ToolKind::Think,
847 ] {
848 assert!(matches!(
849 classify_permission(&perm_req(kind, None), 1),
850 PermissionDecision::AutoAllow
851 ));
852 }
853 }
854
855 #[test]
856 fn edit_with_diff_goes_to_review_with_path_and_contents() {
857 let content = vec![ToolCallContent::Diff(Diff::new(
858 "/w/parse.rs",
859 "fn new() {}\n",
860 ))];
861 match classify_permission(&perm_req(ToolKind::Edit, Some(content)), 7) {
862 PermissionDecision::Review(dr) => {
863 assert_eq!(dr.old_file_path, std::path::PathBuf::from("/w/parse.rs"));
864 assert_eq!(dr.new_file_path, std::path::PathBuf::from("/w/parse.rs"));
865 assert_eq!(dr.new_contents, "fn new() {}\n");
866 assert_eq!(dr.tab_name, "edit parse.rs");
867 assert_eq!(dr.origin_session, 7);
868 }
869 _ => panic!("expected a Review decision"),
870 }
871 }
872
873 #[test]
874 fn edit_without_diff_is_denied_fail_closed() {
875 assert!(
878 matches!(
879 classify_permission(&perm_req(ToolKind::Edit, None), 1),
880 PermissionDecision::Deny
881 ),
882 "Edit without a diff must be denied",
883 );
884 }
885
886 #[test]
887 fn non_read_ops_without_diff_ask_user() {
888 for kind in [ToolKind::Execute, ToolKind::Delete, ToolKind::Other] {
890 assert!(
891 matches!(
892 classify_permission(&perm_req(kind, None), 1),
893 PermissionDecision::AskUser
894 ),
895 "{kind:?} without a diff must be AskUser",
896 );
897 }
898 }
899
900 #[test]
901 fn allow_and_deny_pick_matching_options_preferring_once() {
902 let opts = perm_options();
903 match allow_outcome(&opts) {
905 RequestPermissionOutcome::Selected(sel) => {
906 assert_eq!(sel.option_id, PermissionOptionId::new("allow-once"));
907 }
908 _ => panic!("expected a Selected allow outcome"),
909 }
910 match deny_outcome(&opts) {
911 RequestPermissionOutcome::Selected(sel) => {
912 assert_eq!(sel.option_id, PermissionOptionId::new("reject-once"));
913 }
914 _ => panic!("expected a Selected reject outcome"),
915 }
916 }
917
918 #[test]
919 fn trust_mode_auto_allows_without_review() {
920 let content = vec![ToolCallContent::Diff(Diff::new("/w/a.rs", "x\n"))];
923 assert!(matches!(
924 resolve_decision(true, &perm_req(ToolKind::Edit, Some(content)), 1),
925 PermissionDecision::AutoAllow
926 ));
927 assert!(matches!(
928 resolve_decision(true, &perm_req(ToolKind::Execute, None), 1),
929 PermissionDecision::AutoAllow
930 ));
931 }
932
933 #[test]
934 fn review_mode_defers_to_classification() {
935 assert!(matches!(
938 resolve_decision(false, &perm_req(ToolKind::Read, None), 1),
939 PermissionDecision::AutoAllow
940 ));
941 let content = vec![ToolCallContent::Diff(Diff::new("/w/a.rs", "x\n"))];
942 assert!(matches!(
943 resolve_decision(false, &perm_req(ToolKind::Edit, Some(content)), 1),
944 PermissionDecision::Review(_)
945 ));
946 assert!(matches!(
947 resolve_decision(false, &perm_req(ToolKind::Execute, None), 1),
948 PermissionDecision::AskUser
949 ));
950 }
951
952 #[test]
953 fn missing_option_kind_yields_cancelled() {
954 let only_reject = vec![PermissionOption::new(
956 "r",
957 "Reject",
958 PermissionOptionKind::RejectOnce,
959 )];
960 assert!(matches!(
961 allow_outcome(&only_reject),
962 RequestPermissionOutcome::Cancelled
963 ));
964 }
965
966 fn text_update(text: &str, thought: bool) -> SessionUpdate {
967 let chunk = ContentChunk::new(ContentBlock::Text(TextContent::new(text)));
968 if thought {
969 SessionUpdate::AgentThoughtChunk(chunk)
970 } else {
971 SessionUpdate::AgentMessageChunk(chunk)
972 }
973 }
974
975 fn test_conv_store() -> ConversationStore {
978 ConversationStore::new(Arc::new(|_| {}))
979 }
980
981 #[tokio::test]
982 async fn drain_applies_agent_text_to_conversation_store() {
983 let store = test_conv_store();
984 let (tx, rx) = mpsc::unbounded_channel();
985 let key = SessionKey::new("opencode", 1);
986
987 let notification =
988 SessionNotification::new(AcpSessionId::new("sess-1"), text_update("pong", false));
989 tx.send(notification).expect("send should succeed");
990 drop(tx);
991
992 drain_notifications(rx, store.clone(), key.clone()).await;
993
994 let conv = store.snapshot();
997 assert_eq!(conv.turns.len(), 1);
998 assert_eq!(
999 conv.turns[0].blocks,
1000 vec![crate::acp::conversation::Block::Text("pong".to_string())],
1001 "expected a single assistant Text block \"pong\", got {conv:?}"
1002 );
1003 }
1004
1005 #[tokio::test]
1013 async fn start_failure_still_increments_index_and_logs_per_session() {
1014 let logger = AiLogger::with_defaults();
1015 let handle = AiClientHandle::spawn(
1016 &tokio::runtime::Handle::current(),
1017 logger.clone(),
1018 test_conv_store(),
1019 None,
1020 );
1021 let cfg = ProviderConfig {
1022 command: "/nonexistent/definitely-not-a-real-binary".into(),
1023 args: vec![],
1024 env: vec![],
1025 display_name: "fakeprov",
1026 };
1027
1028 handle.start(cfg.clone());
1029 handle.start(cfg.clone());
1030
1031 let key1 = SessionKey::new("fakeprov", 1);
1032 let key2 = SessionKey::new("fakeprov", 2);
1033
1034 let deadline = std::time::Instant::now() + Duration::from_secs(5);
1035 loop {
1036 let s1 = logger.snapshot_session(&key1);
1037 let s2 = logger.snapshot_session(&key2);
1038 if !s1.is_empty() && !s2.is_empty() {
1039 assert!(
1040 s1.iter().any(|r| r.message.contains("start failed")),
1041 "expected a start-failure record for fakeprov:1, got {s1:?}"
1042 );
1043 assert!(
1044 s2.iter().any(|r| r.message.contains("start failed")),
1045 "expected a start-failure record for fakeprov:2, got {s2:?}"
1046 );
1047 break;
1048 }
1049 assert!(
1050 std::time::Instant::now() < deadline,
1051 "expected both fakeprov:1 and fakeprov:2 start-failure records before the timeout"
1052 );
1053 tokio::time::sleep(Duration::from_millis(20)).await;
1054 }
1055 }
1056
1057 #[tokio::test]
1062 async fn start_failure_leaves_idle_state() {
1063 let logger = AiLogger::with_defaults();
1064 let handle = AiClientHandle::spawn(
1065 &tokio::runtime::Handle::current(),
1066 logger.clone(),
1067 test_conv_store(),
1068 None,
1069 );
1070 let cfg = ProviderConfig {
1071 command: "/nonexistent/definitely-not-a-real-binary".into(),
1072 args: vec![],
1073 env: vec![],
1074 display_name: "fakeprov",
1075 };
1076
1077 handle.start(cfg);
1078
1079 let key = SessionKey::new("fakeprov", 1);
1080 let deadline = std::time::Instant::now() + Duration::from_secs(5);
1081 loop {
1082 let records = logger.snapshot_session(&key);
1083 if records.iter().any(|r| r.message.contains("start failed")) {
1084 break;
1085 }
1086 assert!(
1087 std::time::Instant::now() < deadline,
1088 "expected a start-failure record for fakeprov:1 before the timeout"
1089 );
1090 tokio::time::sleep(Duration::from_millis(20)).await;
1091 }
1092
1093 assert_eq!(
1094 handle.snapshot(),
1095 AiState::default(),
1096 "state must be idle after a failed Start"
1097 );
1098 }
1099
1100 #[cfg(unix)]
1115 const MOCK_AGENT_EXITS_AFTER_SESSION: &str = r#"
1116while IFS= read -r line; do
1117 id=$(printf '%s' "$line" | sed -n 's/.*"id":"\([^"]*\)".*/\1/p')
1118 case "$line" in
1119 *'"method":"initialize"'*)
1120 printf '{"jsonrpc":"2.0","id":"%s","result":{"protocolVersion":1}}\n' "$id" ;;
1121 *'"method":"session/new"'*)
1122 printf '{"jsonrpc":"2.0","id":"%s","result":{"sessionId":"sess-1"}}\n' "$id"
1123 exit 0 ;;
1124 esac
1125done
1126"#;
1127
1128 #[cfg(unix)]
1129 fn mock_agent_provider() -> ProviderConfig {
1130 ProviderConfig {
1131 command: "/bin/sh".into(),
1132 args: vec!["-c".into(), MOCK_AGENT_EXITS_AFTER_SESSION.into()],
1133 env: vec![],
1134 display_name: "mockprov",
1135 }
1136 }
1137
1138 #[cfg(unix)]
1145 async fn wait_for_record(logger: &AiLogger, session: &SessionKey, needle: &str) {
1146 let deadline = std::time::Instant::now() + Duration::from_secs(10);
1147 loop {
1148 let records = logger.snapshot_session(session);
1149 if records.iter().any(|r| r.message.contains(needle)) {
1150 return;
1151 }
1152 assert!(
1153 std::time::Instant::now() < deadline,
1154 "timed out waiting for a {needle:?} record on {session:?}, got {records:?}"
1155 );
1156 tokio::time::sleep(Duration::from_millis(20)).await;
1157 }
1158 }
1159
1160 fn tracing_conv_store() -> (ConversationStore, Arc<Mutex<Vec<SessionStatus>>>) {
1173 let seen = Arc::new(Mutex::new(Vec::new()));
1174 let cell = Arc::new(std::sync::Mutex::new(None::<ConversationStore>));
1179 let c2 = cell.clone();
1180 let s2 = seen.clone();
1181 let store = ConversationStore::new(Arc::new(move |_ev| {
1182 if let Some(ref store) = *c2.lock().unwrap() {
1183 s2.lock().unwrap().push(store.snapshot().status);
1184 }
1185 }));
1186 *cell.lock().unwrap() = Some(store.clone());
1187 (store, seen)
1188 }
1189
1190 #[tokio::test]
1191 async fn drain_transitions_through_thinking_to_idle() {
1192 let (store, seen) = tracing_conv_store();
1193 let key = SessionKey::new("opencode", 1);
1194 let (tx, rx) = mpsc::unbounded_channel();
1195
1196 tx.send(SessionNotification::new(
1198 AcpSessionId::new("sess-1"),
1199 text_update("hello", false),
1200 ))
1201 .expect("send");
1202 drop(tx);
1203
1204 drain_notifications(rx, store.clone(), key.clone()).await;
1205
1206 let trail: Vec<SessionStatus> = seen.lock().unwrap().clone();
1207 assert!(
1208 trail.contains(&SessionStatus::Thinking),
1209 "must have passed through Thinking, got: {trail:?}",
1210 );
1211 assert_eq!(store.snapshot().status, SessionStatus::Idle, "final → Idle");
1212 }
1213
1214 #[tokio::test]
1215 async fn drain_transitions_through_executing_to_idle() {
1216 use agent_client_protocol::schema::v1::ToolCall as AcpToolCall;
1217 let (store, seen) = tracing_conv_store();
1218 let key = SessionKey::new("opencode", 1);
1219 let (tx, rx) = mpsc::unbounded_channel();
1220
1221 let mut tc = AcpToolCall::new("tc-1", "edit parse.rs");
1222 tc.status = ToolCallStatus::InProgress;
1223 tx.send(SessionNotification::new(
1224 AcpSessionId::new("sess-1"),
1225 SessionUpdate::ToolCall(tc),
1226 ))
1227 .expect("send");
1228 drop(tx);
1229
1230 drain_notifications(rx, store.clone(), key.clone()).await;
1231
1232 let trail: Vec<SessionStatus> = seen.lock().unwrap().clone();
1233 assert!(
1234 trail.contains(&SessionStatus::Executing {
1235 tool: "edit parse.rs".into()
1236 }),
1237 "must have passed through Executing, got: {trail:?}",
1238 );
1239 assert_eq!(store.snapshot().status, SessionStatus::Idle, "final → Idle");
1240 }
1241
1242 #[tokio::test]
1243 async fn drain_sets_idle_when_stream_ends() {
1244 let store = test_conv_store();
1245 let key = SessionKey::new("opencode", 1);
1246 let (tx, rx) = mpsc::unbounded_channel();
1247
1248 tx.send(SessionNotification::new(
1249 AcpSessionId::new("sess-1"),
1250 text_update("working", false),
1251 ))
1252 .expect("send");
1253 drop(tx);
1254
1255 drain_notifications(rx, store.clone(), key.clone()).await;
1256
1257 assert_eq!(
1258 store.snapshot().status,
1259 SessionStatus::Idle,
1260 "stream closed → Idle"
1261 );
1262 }
1263
1264 #[tokio::test]
1265 async fn drain_ends_in_idle_when_tool_completed_and_stream_closed() {
1266 use agent_client_protocol::schema::v1::{
1267 ToolCall as AcpToolCall, ToolCallUpdate, ToolCallUpdateFields,
1268 };
1269 let (store, seen) = tracing_conv_store();
1270 let key = SessionKey::new("opencode", 1);
1271 let (tx, rx) = mpsc::unbounded_channel();
1272
1273 let mut tc = AcpToolCall::new("tc-1", "edit");
1274 tc.status = ToolCallStatus::InProgress;
1275 tx.send(SessionNotification::new(
1276 AcpSessionId::new("sess-1"),
1277 SessionUpdate::ToolCall(tc),
1278 ))
1279 .expect("send");
1280
1281 let update = ToolCallUpdate::new(
1282 "tc-1",
1283 ToolCallUpdateFields::new().status(ToolCallStatus::Completed),
1284 );
1285 tx.send(SessionNotification::new(
1286 AcpSessionId::new("sess-1"),
1287 SessionUpdate::ToolCallUpdate(update),
1288 ))
1289 .expect("send");
1290 drop(tx);
1291
1292 drain_notifications(rx, store.clone(), key.clone()).await;
1293
1294 let trail: Vec<SessionStatus> = seen.lock().unwrap().clone();
1295 assert!(
1297 trail.contains(&SessionStatus::Idle),
1298 "must have Idle after tool completed, got: {trail:?}",
1299 );
1300 assert_eq!(store.snapshot().status, SessionStatus::Idle, "final → Idle");
1301 }
1302
1303 #[cfg(unix)]
1309 #[tokio::test(flavor = "multi_thread")]
1310 async fn unexpected_child_exit_resets_state_and_logs() {
1311 let logger = AiLogger::with_defaults();
1312 let handle = AiClientHandle::spawn(
1313 &tokio::runtime::Handle::current(),
1314 logger.clone(),
1315 test_conv_store(),
1316 None,
1317 );
1318 let key = SessionKey::new("mockprov", 1);
1319
1320 handle.start(mock_agent_provider());
1321
1322 wait_for_record(&logger, &key, "session opened").await;
1324 wait_for_record(&logger, &key, "agent exited").await;
1326
1327 assert_eq!(
1328 handle.snapshot(),
1329 AiState::default(),
1330 "state must be idle after the agent exits on its own"
1331 );
1332 }
1333
1334 #[cfg(unix)]
1337 #[tokio::test(flavor = "multi_thread")]
1338 async fn start_after_child_exit_opens_the_next_session() {
1339 let logger = AiLogger::with_defaults();
1340 let handle = AiClientHandle::spawn(
1341 &tokio::runtime::Handle::current(),
1342 logger.clone(),
1343 test_conv_store(),
1344 None,
1345 );
1346 let first = SessionKey::new("mockprov", 1);
1347 let second = SessionKey::new("mockprov", 2);
1348
1349 handle.start(mock_agent_provider());
1350 wait_for_record(&logger, &first, "agent exited").await;
1351
1352 handle.start(mock_agent_provider());
1353 wait_for_record(&logger, &second, "session opened").await;
1354 wait_for_record(&logger, &second, "agent exited").await;
1355
1356 assert_eq!(handle.snapshot(), AiState::default());
1357 }
1358
1359 #[ignore]
1365 #[tokio::test(flavor = "multi_thread")]
1366 async fn opencode_supervisor_end_to_end() {
1367 let logger = AiLogger::with_defaults();
1368 let store = test_conv_store();
1369 let handle = AiClientHandle::spawn(
1370 &tokio::runtime::Handle::current(),
1371 logger.clone(),
1372 store.clone(),
1373 None,
1374 );
1375
1376 handle.start(ProviderConfig::opencode());
1377
1378 let deadline = std::time::Instant::now() + Duration::from_secs(30);
1379 loop {
1380 if handle.snapshot().session.is_some() {
1381 break;
1382 }
1383 assert!(
1384 std::time::Instant::now() < deadline,
1385 "session did not open before the timeout"
1386 );
1387 tokio::time::sleep(Duration::from_millis(100)).await;
1388 }
1389
1390 handle.prompt("reply with the single word: pong".into());
1391
1392 let deadline = std::time::Instant::now() + Duration::from_secs(30);
1393 loop {
1394 let has_agent_text = store.snapshot().turns.iter().any(|t| {
1396 t.role == crate::acp::conversation::Role::Assistant
1397 && t.blocks
1398 .iter()
1399 .any(|b| matches!(b, crate::acp::conversation::Block::Text(_)))
1400 });
1401 if has_agent_text {
1402 break;
1403 }
1404 assert!(
1405 std::time::Instant::now() < deadline,
1406 "no assistant text arrived in the conversation before the timeout"
1407 );
1408 tokio::time::sleep(Duration::from_millis(200)).await;
1409 }
1410
1411 let deadline = std::time::Instant::now() + Duration::from_secs(30);
1414 loop {
1415 if let Some(usage) = store.snapshot().usage {
1416 assert!(usage.size > 0, "context size should be non-zero: {usage:?}");
1417 break;
1418 }
1419 assert!(
1420 std::time::Instant::now() < deadline,
1421 "no usage_update reached the conversation store before the timeout"
1422 );
1423 tokio::time::sleep(Duration::from_millis(200)).await;
1424 }
1425 }
1426}