agency_proxy_client/
lib.rs1use agency_proxy_protocol::{
2 ClientFrame, ClientMessage, MAX_FRAME_BYTES, PROTOCOL_VERSION, ProtocolVersion, ServerFrame,
3 ServerResponse,
4};
5use endpoint_libs::libs::ws::{WireMessage, transport::framed::framed_json_with_max_frame};
6use futures::{SinkExt, StreamExt};
7use std::{collections::BTreeMap, path::Path};
8use thiserror::Error;
9use tokio::{
10 net::UnixStream,
11 sync::{broadcast, mpsc, oneshot},
12};
13
14#[derive(Debug, Error)]
15pub enum ClientError {
16 #[error("I/O error: {0}")]
17 Io(#[from] std::io::Error),
18 #[error("proxy transport failed: {0}")]
19 Transport(String),
20 #[error("proxy protocol failed: {0}")]
21 Protocol(String),
22 #[error("proxy client task stopped")]
23 Closed,
24}
25
26struct PendingRequest {
27 message: ClientMessage,
28 response: oneshot::Sender<Result<ServerResponse, ClientError>>,
29}
30
31#[derive(Clone, Debug)]
32pub struct Client {
33 requests: mpsc::Sender<PendingRequest>,
34 events: broadcast::Sender<ServerFrame>,
35 version: ProtocolVersion,
36}
37
38impl Client {
39 pub async fn connect(socket_path: impl AsRef<Path>) -> Result<Self, ClientError> {
40 let stream = UnixStream::connect(socket_path).await?;
41 let transport = framed_json_with_max_frame(stream, MAX_FRAME_BYTES + 1);
42 let (requests, request_rx) = mpsc::channel(64);
43 let (events, _) = broadcast::channel(512);
44 tokio::spawn(drive(transport, request_rx, events.clone()));
45 let mut client = Self {
46 requests,
47 events,
48 version: PROTOCOL_VERSION,
49 };
50 match client
51 .request(ClientMessage::Hello {
52 client_name: "agency-proxy-client".into(),
53 version: PROTOCOL_VERSION,
54 })
55 .await?
56 {
57 ServerResponse::Hello { version, .. } if version.major == PROTOCOL_VERSION.major => {
58 client.version = version;
59 Ok(client)
60 }
61 ServerResponse::Error { message, .. } => Err(ClientError::Protocol(message)),
62 response => Err(ClientError::Protocol(format!(
63 "unexpected hello response: {response:?}"
64 ))),
65 }
66 }
67
68 pub fn subscribe(&self) -> broadcast::Receiver<ServerFrame> {
69 self.events.subscribe()
70 }
71
72 #[must_use]
73 pub const fn version(&self) -> ProtocolVersion {
74 self.version
75 }
76
77 pub async fn request(&self, message: ClientMessage) -> Result<ServerResponse, ClientError> {
78 let (response, received) = oneshot::channel();
79 self.requests
80 .send(PendingRequest { message, response })
81 .await
82 .map_err(|_| ClientError::Closed)?;
83 received.await.map_err(|_| ClientError::Closed)?
84 }
85}
86
87async fn drive<T, E>(
88 mut transport: T,
89 mut requests: mpsc::Receiver<PendingRequest>,
90 events: broadcast::Sender<ServerFrame>,
91) where
92 T: futures::Sink<WireMessage, Error = E>
93 + futures::Stream<Item = Result<WireMessage, E>>
94 + Unpin,
95 E: std::error::Error,
96{
97 let mut next_request_id = 1u64;
98 let mut pending = BTreeMap::<u64, oneshot::Sender<Result<ServerResponse, ClientError>>>::new();
99 let stopped = loop {
100 tokio::select! {
101 request = requests.recv() => {
102 let Some(request) = request else { break "request channel closed".to_string() };
103 let request_id = next_request_id;
104 next_request_id = next_request_id.wrapping_add(1).max(1);
105 let frame = ClientFrame { request_id, message: request.message };
106 let encoded = match serde_json::to_string(&frame) {
107 Ok(encoded) => encoded,
108 Err(error) => {
109 let _ = request.response.send(Err(ClientError::Protocol(error.to_string())));
110 continue;
111 }
112 };
113 if let Err(error) = transport.send(WireMessage::Text(encoded)).await {
114 let detail = error.to_string();
115 let _ = request.response.send(Err(ClientError::Transport(detail.clone())));
116 break detail;
117 }
118 pending.insert(request_id, request.response);
119 }
120 incoming = transport.next() => {
121 let Some(incoming) = incoming else { break "proxy closed the connection".into() };
122 let message = match incoming {
123 Ok(message) => message,
124 Err(error) => break error.to_string(),
125 };
126 if message.is_close() { break "proxy closed the connection".into(); }
127 let Some(text) = message.as_text() else { continue };
128 let frame: ServerFrame = match serde_json::from_str(text) {
129 Ok(frame) => frame,
130 Err(error) => break error.to_string(),
131 };
132 match frame {
133 ServerFrame::Response { request_id, response } => {
134 if let Some(waiter) = pending.remove(&request_id) {
135 let _ = waiter.send(Ok(response));
136 }
137 }
138 event @ ServerFrame::Event { .. } => {
139 let _ = events.send(event);
140 }
141 }
142 }
143 }
144 };
145 for (_, waiter) in pending {
146 let _ = waiter.send(Err(ClientError::Transport(stopped.clone())));
147 }
148}