1use 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)] async 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)] fn 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}