lattice_ai/acp/
session.rs1use crate::acp::connection::Connection;
5use crate::acp::error::Result;
6
7pub use crate::acp::connection::SessionId;
8
9pub async fn handshake(conn: &Connection, cwd: &str) -> Result<SessionId> {
11 conn.initialize().await?;
12 conn.new_session(cwd).await
13}
14
15pub async fn prompt(conn: &Connection, session: &SessionId, text: &str) -> Result<()> {
17 conn.prompt(session, text).await
18}
19
20#[cfg(test)]
21mod tests {
22 use std::sync::Arc;
23 use std::time::Duration;
24
25 use serde_json::{Value, json};
26 use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
27 use tokio::sync::mpsc;
28
29 use super::*;
30 use crate::acp::connection::SessionNotification;
31
32 async fn run_mock_peer(peer: tokio::io::DuplexStream) {
37 let mut reader = BufReader::new(peer);
38 let mut line = String::new();
39 loop {
40 line.clear();
41 match reader.read_line(&mut line).await {
42 Ok(0) | Err(_) => break,
43 Ok(_) => {}
44 }
45 let trimmed = line.trim();
46 if trimmed.is_empty() {
47 continue;
48 }
49 let request: Value = match serde_json::from_str(trimmed) {
50 Ok(v) => v,
51 Err(_) => continue,
52 };
53 let method = request.get("method").and_then(Value::as_str).unwrap_or("");
54 let id = request.get("id").cloned().unwrap_or(Value::Null);
55
56 let response = match method {
57 "initialize" => Some(json!({
58 "jsonrpc": "2.0",
59 "id": id,
60 "result": { "protocolVersion": 1 },
61 })),
62 "session/new" => Some(json!({
63 "jsonrpc": "2.0",
64 "id": id,
65 "result": { "sessionId": "sess-42" },
66 })),
67 _ => None,
68 };
69
70 if let Some(response) = response {
71 write_line(reader.get_mut(), &response).await;
72 }
73 }
74 }
75
76 async fn write_line(writer: &mut tokio::io::DuplexStream, value: &Value) {
77 let mut line = value.to_string();
78 line.push('\n');
79 let _ = writer.write_all(line.as_bytes()).await;
80 }
81
82 fn spawn_connection_with_mock_peer() -> (
83 Arc<Connection>,
84 mpsc::UnboundedReceiver<SessionNotification>,
85 ) {
86 let (ours, mock) = tokio::io::duplex(8192);
87 let (reader, writer) = tokio::io::split(ours);
88 let (connection, notif_rx, _perm_rx) = Connection::spawn(reader, writer);
89 tokio::spawn(run_mock_peer(mock));
90 (connection, notif_rx)
91 }
92
93 #[tokio::test]
94 async fn handshake_returns_session_id() {
95 let (connection, _notif_rx) = spawn_connection_with_mock_peer();
96
97 let session = handshake(&connection, "/work")
98 .await
99 .expect("handshake should succeed");
100
101 assert_eq!(session, SessionId("sess-42".to_string()));
102 }
103
104 #[ignore]
109 #[tokio::test(flavor = "multi_thread")]
110 async fn opencode_end_to_end() {
111 use tokio::process::Command;
112
113 let mut child = Command::new("/Users/dhruva/.opencode/bin/opencode")
114 .arg("acp")
115 .stdin(std::process::Stdio::piped())
116 .stdout(std::process::Stdio::piped())
117 .stderr(std::process::Stdio::null())
118 .spawn()
119 .expect("opencode acp should spawn");
120
121 let stdin = child.stdin.take().expect("child stdin should be piped");
122 let stdout = child.stdout.take().expect("child stdout should be piped");
123
124 let (connection, mut notif_rx, _perm_rx) = Connection::spawn(stdout, stdin);
125
126 let session = handshake(&connection, ".")
127 .await
128 .expect("handshake should succeed");
129
130 prompt(&connection, &session, "reply with the single word: pong")
131 .await
132 .expect("prompt should succeed");
133
134 let notification = tokio::time::timeout(Duration::from_secs(30), notif_rx.recv())
135 .await
136 .expect("a session/update notification should arrive before the timeout")
137 .expect("the notification channel should still be open");
138
139 assert_eq!(notification.session_id.0.to_string(), session.0);
140
141 let _ = child.kill().await;
142 }
143}