Skip to main content

piw/
server.rs

1//! `piw serve`: a loopback-only, one-WebSocket-to-one-server-socket relay.
2
3use anyhow::{Context, Result};
4use futures_util::{SinkExt, StreamExt};
5use std::path::PathBuf;
6use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
7#[cfg(windows)]
8use tokio::net::windows::named_pipe::ClientOptions;
9#[cfg(unix)]
10use tokio::net::UnixStream;
11use tokio::net::{TcpListener, TcpStream};
12use tokio_tungstenite::tungstenite::handshake::server::{ErrorResponse, Request, Response};
13use tokio_tungstenite::tungstenite::http::StatusCode;
14use tokio_tungstenite::tungstenite::Message;
15
16pub struct ServeOptions {
17    pub socket_path: PathBuf,
18    pub bind: String,
19}
20
21pub async fn serve(options: ServeOptions) -> Result<()> {
22    let listener = TcpListener::bind(&options.bind)
23        .await
24        .with_context(|| format!("binding {}", options.bind))?;
25    eprintln!(
26        "piw serve: relaying {} on ws://{}/ws",
27        options.socket_path.display(),
28        listener.local_addr()?
29    );
30    serve_on(listener, options.socket_path).await
31}
32
33pub async fn serve_on(listener: TcpListener, socket_path: PathBuf) -> Result<()> {
34    let local = listener.local_addr()?;
35    if !local.ip().is_loopback() {
36        anyhow::bail!(
37            "refusing to serve on non-loopback address {local}: the client protocol is unauthenticated"
38        );
39    }
40    loop {
41        let (stream, _) = listener.accept().await?;
42        let socket_path = socket_path.clone();
43        tokio::spawn(async move {
44            if let Err(error) = relay(stream, socket_path).await {
45                eprintln!("piw serve connection: {error:#}");
46            }
47        });
48    }
49}
50
51#[allow(clippy::result_large_err)] // Required by tungstenite's handshake callback type.
52async fn relay(stream: TcpStream, socket_path: PathBuf) -> Result<()> {
53    let websocket =
54        tokio_tungstenite::accept_hdr_async(stream, |request: &Request, response: Response| {
55            validate_request(request, response)
56        })
57        .await
58        .context("accepting WebSocket")?;
59    #[cfg(unix)]
60    {
61        let server = UnixStream::connect(&socket_path)
62            .await
63            .with_context(|| format!("connecting to workflow server {}", socket_path.display()))?;
64        relay_server(websocket, server).await
65    }
66    #[cfg(windows)]
67    {
68        let server = ClientOptions::new()
69            .open(&socket_path)
70            .with_context(|| format!("connecting to workflow server {}", socket_path.display()))?;
71        relay_server(websocket, server).await
72    }
73    #[cfg(not(any(unix, windows)))]
74    anyhow::bail!("local workflow server transport is not supported on this platform");
75}
76
77async fn relay_server<S>(
78    websocket: tokio_tungstenite::WebSocketStream<TcpStream>,
79    server: S,
80) -> Result<()>
81where
82    S: AsyncRead + AsyncWrite + Unpin,
83{
84    let (mut ws_sink, mut ws_stream) = websocket.split();
85    let (server_read, mut server_write) = tokio::io::split(server);
86    let mut server_lines = BufReader::new(server_read).lines();
87    loop {
88        tokio::select! {
89            server_line = server_lines.next_line() => {
90                let Some(line) = server_line.context("reading workflow server")? else { break };
91                ws_sink.send(Message::Text(line.into())).await.context("sending WebSocket frame")?;
92            }
93            ws_message = ws_stream.next() => {
94                let Some(message) = ws_message else { break };
95                match message.context("reading WebSocket frame")? {
96                    Message::Text(text) => {
97                        if text.contains('\n') || text.contains('\r') {
98                            anyhow::bail!("client frame contains a line break");
99                        }
100                        server_write.write_all(text.as_bytes()).await?;
101                        server_write.write_all(b"\n").await?;
102                        server_write.flush().await?;
103                    }
104                    Message::Close(_) => break,
105                    Message::Ping(payload) => ws_sink.send(Message::Pong(payload)).await?,
106                    Message::Pong(_) => {}
107                    Message::Binary(_) | Message::Frame(_) => {
108                        anyhow::bail!("client frame must be text");
109                    }
110                }
111            }
112        }
113    }
114    Ok(())
115}
116
117#[allow(clippy::result_large_err)] // Required by tungstenite's handshake callback type.
118fn validate_request(request: &Request, response: Response) -> Result<Response, ErrorResponse> {
119    if request.uri().path() != "/ws" {
120        return Err(error_response(
121            StatusCode::NOT_FOUND,
122            "WebSocket path must be /ws",
123        ));
124    }
125    if let Some(origin) = request.headers().get("origin") {
126        let loopback = origin
127            .to_str()
128            .ok()
129            .and_then(|value| {
130                value
131                    .parse::<tokio_tungstenite::tungstenite::http::Uri>()
132                    .ok()
133            })
134            .and_then(|uri| uri.host().map(str::to_string))
135            .is_some_and(|host| {
136                matches!(host.as_str(), "127.0.0.1" | "localhost" | "[::1]" | "::1")
137            });
138        if !loopback {
139            return Err(error_response(
140                StatusCode::FORBIDDEN,
141                "origin is not loopback",
142            ));
143        }
144    }
145    Ok(response)
146}
147
148fn error_response(status: StatusCode, message: &str) -> ErrorResponse {
149    tokio_tungstenite::tungstenite::http::Response::builder()
150        .status(status)
151        .body(Some(message.to_string()))
152        .expect("valid relay error response")
153}
154
155#[cfg(test)]
156mod tests {
157    use super::*;
158    use futures_util::{SinkExt, StreamExt};
159    #[cfg(unix)]
160    use tempfile::tempdir;
161    #[cfg(unix)]
162    use tokio::net::UnixListener;
163
164    #[test]
165    fn handshake_rejects_a_lookalike_loopback_origin() {
166        let request = Request::builder()
167            .uri("/ws")
168            .header("origin", "http://localhost.evil")
169            .body(())
170            .unwrap();
171        let response = Response::new(());
172        let error = validate_request(&request, response).unwrap_err();
173        assert_eq!(error.status(), StatusCode::FORBIDDEN);
174    }
175
176    #[tokio::test]
177    async fn relay_rejects_a_non_loopback_listener() {
178        let listener = TcpListener::bind("0.0.0.0:0").await.unwrap();
179        let error = serve_on(listener, PathBuf::from("unused"))
180            .await
181            .unwrap_err();
182        assert!(error.to_string().contains("non-loopback"));
183    }
184
185    #[cfg(unix)]
186    #[tokio::test]
187    async fn relay_couples_one_websocket_to_one_server_socket() {
188        let temporary = tempdir().unwrap();
189        let socket_path = temporary.path().join("host.sock");
190        let server_listener = UnixListener::bind(&socket_path).unwrap();
191        let server_task = tokio::spawn(async move {
192            let (server, _) = server_listener.accept().await.unwrap();
193            let (read, mut write) = server.into_split();
194            let hello = format!(
195                "{{\"connectionId\":\"one\",\"packageVersion\":\"{}\",\"schema\":\"pi-workflows.client.v1\",\"type\":\"hello\"}}\n",
196                env!("CARGO_PKG_VERSION")
197            );
198            write.write_all(hello.as_bytes()).await.unwrap();
199            let mut lines = BufReader::new(read).lines();
200            lines.next_line().await.unwrap().unwrap()
201        });
202        let tcp = TcpListener::bind("127.0.0.1:0").await.unwrap();
203        let address = tcp.local_addr().unwrap();
204        let relay_task = tokio::spawn(serve_on(tcp, socket_path));
205        let (mut websocket, _) = tokio_tungstenite::connect_async(format!("ws://{address}/ws"))
206            .await
207            .unwrap();
208        let hello = websocket.next().await.unwrap().unwrap();
209        assert!(hello.into_text().unwrap().contains("\"type\":\"hello\""));
210        let request = "{\"clientId\":\"client\",\"idempotencyKey\":\"key\",\"operation\":\"server.status\",\"payload\":{},\"requestId\":\"request\",\"schema\":\"pi-workflows.client.v1\",\"type\":\"request\"}";
211        websocket
212            .send(Message::Text(request.to_string().into()))
213            .await
214            .unwrap();
215        assert_eq!(server_task.await.unwrap(), request);
216        websocket.close(None).await.unwrap();
217        relay_task.abort();
218    }
219}