1use std::io::{BufReader, Write};
4use std::os::unix::net::UnixStream;
5use std::path::Path;
6use std::time::Duration;
7
8use serde_json::Value;
9
10use crate::frame::{read_message, write_message, Framing};
11use crate::rpc::{
12 invalid_request, Error, JsonRpcError, JsonRpcId, JsonRpcRequest, JsonRpcResponse, Notification,
13 Request, Response,
14};
15
16pub struct Client {
19 writer: UnixStream,
20 reader: BufReader<UnixStream>,
21 framing: Framing,
22 next_id: i64,
23 on_notification: Option<Box<dyn FnMut(Notification) + Send>>,
24}
25
26impl Client {
27 pub fn connect(path: impl AsRef<Path>) -> Result<Self, Error> {
29 Self::connect_with_framing(path, Framing::Jsonl)
30 }
31
32 pub fn connect_with_framing(path: impl AsRef<Path>, framing: Framing) -> Result<Self, Error> {
34 let stream = UnixStream::connect(path.as_ref())?;
35 let reader = BufReader::new(stream.try_clone()?);
36 Ok(Self {
37 writer: stream,
38 reader,
39 framing,
40 next_id: 1,
41 on_notification: None,
42 })
43 }
44
45 pub fn framing(&self) -> Framing {
46 self.framing
47 }
48
49 pub fn on_notification<F>(&mut self, callback: F)
51 where
52 F: FnMut(Notification) + Send + 'static,
53 {
54 self.on_notification = Some(Box::new(callback));
55 }
56
57 pub fn request(&mut self, method: &str, params: Value) -> Result<Value, Error> {
59 let id = JsonRpcId::Number(self.next_id);
60 self.next_id += 1;
61 let msg = JsonRpcRequest::call(id.clone(), method, params);
62 self.write_rpc(&msg)?;
63 loop {
64 let (payload, _) = read_message(&mut self.reader)?;
65 let value: Value = serde_json::from_slice(&payload)?;
66 if is_notification(&value) {
67 let method = value
68 .get("method")
69 .and_then(Value::as_str)
70 .unwrap_or("")
71 .to_string();
72 let params = value.get("params").cloned().unwrap_or(Value::Null);
73 let note = Notification::parse(&method, params);
74 if let Some(cb) = &mut self.on_notification {
75 cb(note);
76 }
77 continue;
78 }
79 let resp: JsonRpcResponse = serde_json::from_value(value)?;
80 if resp.id.as_ref() != Some(&id) {
81 return Err(Error::Rpc(id_mismatch(resp.error)));
82 }
83 if let Some(err) = resp.error {
84 return Err(Error::Rpc(err));
85 }
86 return Ok(resp.result.unwrap_or(Value::Null));
87 }
88 }
89
90 pub fn request_typed(&mut self, req: &Request) -> Result<Value, Error> {
92 self.request(req.method().as_str(), req.to_params())
93 }
94
95 pub fn notify(&mut self, method: &str, params: Value) -> Result<(), Error> {
97 self.write_rpc(&JsonRpcRequest::notification(method, params))
98 }
99
100 pub fn wait_notification(&mut self, timeout: Duration) -> Result<Notification, Error> {
102 self.reader.get_ref().set_read_timeout(Some(timeout))?;
103 let result = read_next_notification(&mut self.reader);
104 let _ = self.reader.get_ref().set_read_timeout(None);
105 result
106 }
107
108 fn write_rpc(&mut self, msg: &JsonRpcRequest) -> Result<(), Error> {
109 let bytes = serde_json::to_vec(msg)?;
110 write_message(&mut self.writer, &bytes, self.framing)?;
111 self.writer.flush()?;
112 Ok(())
113 }
114}
115
116fn read_next_notification(reader: &mut BufReader<UnixStream>) -> Result<Notification, Error> {
117 loop {
118 let (payload, _) = read_message(reader)?;
119 if payload.iter().all(u8::is_ascii_whitespace) {
120 continue;
121 }
122 let value: Value = serde_json::from_slice(&payload)?;
123 if is_notification(&value) {
124 let method = value
125 .get("method")
126 .and_then(Value::as_str)
127 .unwrap_or("")
128 .to_string();
129 let params = value.get("params").cloned().unwrap_or(Value::Null);
130 return Ok(Notification::parse(&method, params));
131 }
132 let mut err = invalid_request();
133 err.message = "expected notification".into();
134 return Err(Error::Rpc(err));
135 }
136}
137
138fn is_notification(value: &Value) -> bool {
141 let obj = match value.as_object() {
142 Some(o) => o,
143 None => return false,
144 };
145 obj.contains_key("method") && !obj.contains_key("id") && !obj.contains_key("result")
146}
147
148fn id_mismatch(server_error: Option<JsonRpcError>) -> JsonRpcError {
149 match server_error {
150 Some(err) => err,
151 None => {
152 let mut err = invalid_request();
153 err.message = "response id does not match request".into();
154 err
155 }
156 }
157}
158
159pub fn decode_response(method: &str, value: Value) -> Result<Response, Error> {
161 match crate::rpc::Method::parse(method).map_err(Error::Rpc)? {
162 crate::rpc::Method::Initialize => Ok(Response::Initialize(serde_json::from_value(value)?)),
163 crate::rpc::Method::IdentityGet => {
164 Ok(Response::IdentityGet(serde_json::from_value(value)?))
165 }
166 crate::rpc::Method::IssueList => Ok(Response::IssueList(serde_json::from_value(value)?)),
167 crate::rpc::Method::IssueGet => Ok(Response::IssueGet(serde_json::from_value(value)?)),
168 crate::rpc::Method::IssueReady => Ok(Response::IssueReady(serde_json::from_value(value)?)),
169 crate::rpc::Method::IssueSearch => {
170 Ok(Response::IssueSearch(serde_json::from_value(value)?))
171 }
172 crate::rpc::Method::IssueClaims => {
173 Ok(Response::IssueClaims(serde_json::from_value(value)?))
174 }
175 crate::rpc::Method::IssueAgenda => {
176 Ok(Response::IssueAgenda(serde_json::from_value(value)?))
177 }
178 crate::rpc::Method::IssueShow => Ok(Response::IssueShow(serde_json::from_value(value)?)),
179 crate::rpc::Method::IssueExcerpt => {
180 Ok(Response::IssueExcerpt(serde_json::from_value(value)?))
181 }
182 crate::rpc::Method::IssueTree => Ok(Response::IssueTree(serde_json::from_value(value)?)),
183 crate::rpc::Method::IssueRelated => {
184 Ok(Response::IssueRelated(serde_json::from_value(value)?))
185 }
186 crate::rpc::Method::IssueChildren => {
187 Ok(Response::IssueChildren(serde_json::from_value(value)?))
188 }
189 crate::rpc::Method::IssueAncestors => {
190 Ok(Response::IssueAncestors(serde_json::from_value(value)?))
191 }
192 crate::rpc::Method::IssueImpact => {
193 Ok(Response::IssueImpact(serde_json::from_value(value)?))
194 }
195 crate::rpc::Method::IssueBacklinks => {
196 Ok(Response::IssueBacklinks(serde_json::from_value(value)?))
197 }
198 crate::rpc::Method::IssueOpen => Ok(Response::IssueOpen(serde_json::from_value(value)?)),
199 crate::rpc::Method::IssueCreate => {
200 Ok(Response::IssueCreate(serde_json::from_value(value)?))
201 }
202 crate::rpc::Method::IssueUpdate => {
203 Ok(Response::IssueUpdate(serde_json::from_value(value)?))
204 }
205 crate::rpc::Method::IssueClaim => Ok(Response::IssueClaim(serde_json::from_value(value)?)),
206 crate::rpc::Method::IssueNote => Ok(Response::IssueNote(serde_json::from_value(value)?)),
207 crate::rpc::Method::IssueRefile => {
208 Ok(Response::IssueRefile(serde_json::from_value(value)?))
209 }
210 crate::rpc::Method::ProjectList => {
211 Ok(Response::ProjectList(serde_json::from_value(value)?))
212 }
213 crate::rpc::Method::EventsSince => {
214 Ok(Response::EventsSince(serde_json::from_value(value)?))
215 }
216 crate::rpc::Method::EventsGen => Ok(Response::EventsGen(serde_json::from_value(value)?)),
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223 use crate::frame::{read_message, write_message, Framing};
224 use crate::rpc::{JsonRpcRequest, NOTIFY_VAULT_CHANGED};
225 use serde_json::json;
226 use std::io::{BufReader, Write};
227 use std::os::unix::net::UnixListener;
228 use std::sync::{Arc, Mutex};
229 use std::thread;
230
231 fn serve_one(
232 path: &Path,
233 framing: Framing,
234 reply: impl FnOnce(JsonRpcRequest) -> Value + Send + 'static,
235 ) {
236 let listener = UnixListener::bind(path).unwrap();
237 thread::spawn(move || {
238 let (stream, _) = listener.accept().unwrap();
239 let mut reader = BufReader::new(stream.try_clone().unwrap());
240 let mut writer = stream;
241 let (payload, got) = read_message(&mut reader).unwrap();
242 assert_eq!(got, framing);
243 let req: JsonRpcRequest = serde_json::from_slice(&payload).unwrap();
244 let body = reply(req);
245 let bytes = serde_json::to_vec(&body).unwrap();
246 write_message(&mut writer, &bytes, framing).unwrap();
247 writer.flush().unwrap();
248 });
249 }
250
251 #[test]
252 fn request_roundtrip_jsonl() {
253 let dir = tempfile::tempdir().unwrap();
254 let sock = dir.path().join("control.sock");
255 serve_one(
256 &sock,
257 Framing::Jsonl,
258 |req| json!({"jsonrpc":"2.0","id":req.id,"result":{"identity":"rg"}}),
259 );
260 let mut client = Client::connect(&sock).unwrap();
261 assert_eq!(client.framing(), Framing::Jsonl);
262 let result = client.request("identity/get", json!({})).unwrap();
263 assert_eq!(result["identity"], "rg");
264 }
265
266 #[test]
267 fn request_roundtrip_headers() {
268 let dir = tempfile::tempdir().unwrap();
269 let sock = dir.path().join("control.sock");
270 serve_one(
271 &sock,
272 Framing::Headers,
273 |req| json!({"jsonrpc":"2.0","id":req.id,"result":{"ok":true}}),
274 );
275 let mut client = Client::connect_with_framing(&sock, Framing::Headers).unwrap();
276 let result = client.request_typed(&Request::IdentityGet).unwrap();
277 assert_eq!(result["ok"], true);
278 }
279
280 #[test]
281 fn notification_callback_fires_before_result() {
282 let dir = tempfile::tempdir().unwrap();
283 let sock = dir.path().join("control.sock");
284 let listener = UnixListener::bind(&sock).unwrap();
285 thread::spawn(move || {
286 let (stream, _) = listener.accept().unwrap();
287 let mut reader = BufReader::new(stream.try_clone().unwrap());
288 let mut writer = stream;
289 let (payload, framing) = read_message(&mut reader).unwrap();
290 let req: JsonRpcRequest = serde_json::from_slice(&payload).unwrap();
291 let note = json!({
292 "jsonrpc":"2.0",
293 "method": NOTIFY_VAULT_CHANGED,
294 "params": {"generation": 9, "revision": 3, "projects": ["atlas"]}
295 });
296 write_message(&mut writer, &serde_json::to_vec(¬e).unwrap(), framing).unwrap();
297 let result = json!({"jsonrpc":"2.0","id":req.id,"result":{"ok":true}});
298 write_message(&mut writer, &serde_json::to_vec(&result).unwrap(), framing).unwrap();
299 writer.flush().unwrap();
300 });
301
302 let seen = Arc::new(Mutex::new(Vec::new()));
303 let seen_cb = Arc::clone(&seen);
304 let mut client = Client::connect(&sock).unwrap();
305 client.on_notification(move |n| seen_cb.lock().unwrap().push(n.method().to_string()));
306 let result = client.request("events/gen", json!({})).unwrap();
307 assert_eq!(result["ok"], true);
308 assert_eq!(seen.lock().unwrap().as_slice(), [NOTIFY_VAULT_CHANGED]);
309 }
310
311 #[test]
312 fn null_response_id_is_rpc_error() {
313 let dir = tempfile::tempdir().unwrap();
314 let sock = dir.path().join("control.sock");
315 serve_one(&sock, Framing::Jsonl, |_req| {
316 json!({
317 "jsonrpc":"2.0",
318 "id": null,
319 "error": {"code": -32600, "message": "invalid request"}
320 })
321 });
322 let mut client = Client::connect(&sock).unwrap();
323 let err = client.request("identity/get", json!({})).unwrap_err();
324 match err {
325 Error::Rpc(e) => {
326 assert_eq!(e.code, -32600);
327 assert_eq!(e.message, "invalid request");
328 }
329 other => panic!("{other:?}"),
330 }
331 }
332
333 #[test]
334 fn unmatched_response_id_is_rpc_error() {
335 let dir = tempfile::tempdir().unwrap();
336 let sock = dir.path().join("control.sock");
337 serve_one(
338 &sock,
339 Framing::Jsonl,
340 |_req| json!({"jsonrpc":"2.0","id": 99, "result":{"ok":true}}),
341 );
342 let mut client = Client::connect(&sock).unwrap();
343 let err = client.request("identity/get", json!({})).unwrap_err();
344 match err {
345 Error::Rpc(e) => {
346 assert_eq!(e.code, -32600);
347 assert_eq!(e.message, "response id does not match request");
348 }
349 other => panic!("{other:?}"),
350 }
351 }
352
353 #[test]
354 fn rpc_error_is_returned() {
355 let dir = tempfile::tempdir().unwrap();
356 let sock = dir.path().join("control.sock");
357 serve_one(&sock, Framing::Jsonl, |req| {
358 json!({
359 "jsonrpc":"2.0",
360 "id": req.id,
361 "error": {"code": -32601, "message": "method not found", "data": {"method": "nope"}}
362 })
363 });
364 let mut client = Client::connect(&sock).unwrap();
365 let err = client.request("nope", json!({})).unwrap_err();
366 match err {
367 Error::Rpc(e) => assert_eq!(e.code, -32601),
368 other => panic!("{other:?}"),
369 }
370 }
371
372 #[test]
373 fn wait_notification_reads_a_push() {
374 let dir = tempfile::tempdir().unwrap();
375 let sock = dir.path().join("control.sock");
376 let listener = UnixListener::bind(&sock).unwrap();
377 thread::spawn(move || {
378 let (stream, _) = listener.accept().unwrap();
379 let mut writer = stream;
380 let note = json!({
381 "jsonrpc":"2.0",
382 "method": NOTIFY_VAULT_CHANGED,
383 "params": {"generation": 2, "revision": 4, "projects": []}
384 });
385 write_message(
386 &mut writer,
387 &serde_json::to_vec(¬e).unwrap(),
388 Framing::Jsonl,
389 )
390 .unwrap();
391 writer.flush().unwrap();
392 thread::sleep(std::time::Duration::from_millis(50));
393 });
394 let mut client = Client::connect(&sock).unwrap();
395 let note = client
396 .wait_notification(std::time::Duration::from_secs(2))
397 .unwrap();
398 assert_eq!(note.method(), NOTIFY_VAULT_CHANGED);
399 }
400
401 #[test]
402 fn notify_writes_without_id() {
403 let dir = tempfile::tempdir().unwrap();
404 let sock = dir.path().join("control.sock");
405 let listener = UnixListener::bind(&sock).unwrap();
406 let handle = thread::spawn(move || {
407 let (stream, _) = listener.accept().unwrap();
408 let mut reader = BufReader::new(stream);
409 let (payload, _) = read_message(&mut reader).unwrap();
410 let req: JsonRpcRequest = serde_json::from_slice(&payload).unwrap();
411 assert!(req.is_notification());
412 assert_eq!(req.method, "serve/shutting_down");
413 });
414 let mut client = Client::connect(&sock).unwrap();
415 client.notify("serve/shutting_down", json!({})).unwrap();
416 handle.join().unwrap();
417 }
418
419 #[test]
420 fn decode_response_covers_methods() {
421 let value = json!({"protocolVersion":1,"capabilities":[],"root":"/","prefix":"Software","generation":1,"revision":1,"identity":"a"});
422 match decode_response("initialize", value).unwrap() {
423 Response::Initialize(r) => assert_eq!(r.protocol_version, 1),
424 other => panic!("{other:?}"),
425 }
426 let list = json!({"issues":[],"revision":1});
427 assert!(matches!(
428 decode_response("issue/list", list.clone()).unwrap(),
429 Response::IssueList(_)
430 ));
431 assert!(matches!(
432 decode_response("issue/ready", list).unwrap(),
433 Response::IssueReady(_)
434 ));
435 assert!(matches!(
436 decode_response("events/gen", json!({"generation":1,"revision":1})).unwrap(),
437 Response::EventsGen(_)
438 ));
439 assert!(decode_response("issue/fold", json!({})).is_err());
440
441 let detail = json!({
442 "id":"atlas-1a2b","project":"atlas","title":"t","state":"TODO","priority":"B",
443 "properties":{},"org_tags":[],"tags":[],"blocked_by":[],"parent":null,
444 "claimed_by":null,"claimed_at":null,"file":"f","line_start":1,"line_end":2,
445 "revision":1
446 });
447 for method in [
448 "issue/get",
449 "issue/show",
450 "issue/open",
451 "issue/excerpt",
452 "issue/search",
453 "issue/claims",
454 "issue/agenda",
455 "issue/tree",
456 "issue/related",
457 "issue/children",
458 "issue/ancestors",
459 "issue/impact",
460 "issue/backlinks",
461 "issue/create",
462 "issue/update",
463 "issue/claim",
464 "issue/note",
465 "issue/refile",
466 "project/list",
467 "events/since",
468 "identity/get",
469 ] {
470 let value = match method {
471 "issue/get" | "issue/show" | "issue/open" => detail.clone(),
472 "issue/excerpt" => json!({
473 "id":"atlas-1a2b","file":"f","line_start":1,"line_end":2,
474 "text":"","suppressed":false
475 }),
476 "issue/search" | "issue/claims" | "issue/agenda" | "issue/related"
477 | "issue/children" | "issue/ancestors" | "issue/impact" | "issue/backlinks" => {
478 json!([])
479 }
480 "issue/tree" => json!({"text": "* a"}),
481 "issue/create" | "issue/update" | "issue/claim" | "issue/note" | "issue/refile" => {
482 json!({"ok":true,"report":"","issue":null,"revision":1,"generation":1})
483 }
484 "project/list" => json!({"projects":[],"revision":1}),
485 "events/since" => json!({"events":[],"generation":1}),
486 "identity/get" => {
487 json!({"identity":"a","root":"/","prefix":"Software","version":"0.2.0"})
488 }
489 _ => json!({}),
490 };
491 decode_response(method, value).expect(method);
492 }
493 }
494}