Skip to main content

agency_proxy_client/
lib.rs

1use 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}