patx/relay-lang

Add WebSocket route and socket APIs

Commit 2f08acb · Harrison Erd · 2026-02-10T14:34:46-05:00

Changeset
2f08acbcc93214b83e2cda111071a783b91b6d9b
Parents
bc5515770cc7996e49c4b4391c36b3274017c519

View source at this commit

Comments

No comments yet.

Log in to comment

Diff

diff --git a/Cargo.toml b/Cargo.toml
index 9f942cd..c3b3fd7 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -10,7 +10,7 @@ anyhow = "1"
 serde = { version = "1", features = ["derive"] }
 serde_json = "1"
 reqwest = { version = "0.12", features = ["json", "rustls-tls"] }
-axum = "0.7"
+axum = { version = "0.7", features = ["ws"] }
 tower = "0.4"
 tower-http = "0.5"
 html-escape = "0.2"
diff --git a/README.md b/README.md
index 994129c..01aa093 100644
--- a/README.md
+++ b/README.md
@@ -752,6 +752,7 @@ Define HTTP endpoints using decorators:
 - `@app.put(path)`
 - `@app.patch(path)`
 - `@app.delete(path)`
+- `@app.ws(path)`
 
 ```relay
 app = WebApp()
@@ -765,6 +766,24 @@ fn create_user(name: str, email: str)
     return {"id": 123, "name": name, "email": email}
 ```
 
+#### WebSocket Routes
+
+Use `@app.ws(path)` for WebSocket upgrade endpoints.
+WebSocket handlers receive a `socket` object with:
+- `socket.recv()` → returns text (`str`), binary (`bytes`), or `None` on close
+- `socket.send(value)` → sends text (or binary when `bytes`)
+- `socket.close()` → closes the connection
+
+```relay
[email protected]("/chat")
+fn chat_room()
+    while True
+        msg = socket.recv()
+        if (msg == None)
+            break
+        socket.send("echo: " + str(msg))
+```
+
 #### Path Parameters
 
 Use `<name>` syntax to capture path segments:
@@ -1871,7 +1890,7 @@ Found a bug? [Open an issue](https://github.com/patx/relay-lang/issues) with:
 - [x] List comprehensions
 - [x] Destructuring assignment
 - [x] Error handling with try/except
-- [ ] WebSocket support
+- [x] WebSocket support
 - [X] Static file serving
 - [X] Middleware support
 
diff --git a/src/main.rs b/src/main.rs
index ef87a15..80e9ce0 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -21,6 +21,7 @@ use std::{
 };
 
 use axum::{
+    extract::ws::{Message as WsMessage, WebSocket, WebSocketUpgrade},
     extract::{Path, Query},
     http::{HeaderMap, Method, StatusCode},
     response::IntoResponse,
@@ -29,7 +30,7 @@ use axum::{
 };
 
 use async_recursion::async_recursion;
-use futures::stream::TryStreamExt;
+use futures::{stream::TryStreamExt, SinkExt, StreamExt};
 use html_escape::encode_safe;
 use indexmap::IndexMap;
 use mongodb::{
@@ -1267,6 +1268,12 @@ enum Object {
     MongoClient(MongoClientHandle),
     MongoDatabase(MongoDatabaseHandle),
     MongoCollection(MongoCollectionHandle),
+    WebSocket(WebSocketHandle),
+}
+
+#[derive(Clone)]
+struct WebSocketHandle {
+    inner: Arc<tokio::sync::Mutex<WebSocket>>,
 }
 
 #[derive(Clone)]
@@ -2861,6 +2868,7 @@ struct RequestParts {
     headers: HashMap<String, String>,
     cookies: HashMap<String, String>,
     session_id: Option<String>,
+    websocket: Option<WebSocketHandle>,
 }
 
 type HandlerCb = Arc<dyn Fn(WebAppHandle, String, RequestParts) -> BoxFut + Send + Sync>;
@@ -2915,7 +2923,11 @@ async fn run_app(app: WebAppHandle) -> RResult<()> {
             let route_prefix = route_for_index.clone();
             async move {
                 let response = serve_static_file(&route_prefix, &dir, "index.html").await;
-                Ok::<_, (StatusCode, String)>(value_to_axum_response(Value::Response(response)).await.into_response())
+                Ok::<_, (StatusCode, String)>(
+                    value_to_axum_response(Value::Response(response))
+                        .await
+                        .into_response(),
+                )
             }
         };
 
@@ -2926,7 +2938,11 @@ async fn run_app(app: WebAppHandle) -> RResult<()> {
             let route_prefix = route_for_files.clone();
             async move {
                 let response = serve_static_file(&route_prefix, &dir, &file).await;
-                Ok::<_, (StatusCode, String)>(value_to_axum_response(Value::Response(response)).await.into_response())
+                Ok::<_, (StatusCode, String)>(
+                    value_to_axum_response(Value::Response(response))
+                        .await
+                        .into_response(),
+                )
             }
         };
 
@@ -2942,15 +2958,22 @@ async fn run_app(app: WebAppHandle) -> RResult<()> {
         let route_path = r.path.clone();
         let app_handle = app.clone();
 
-        // define handler in THIS scope so `get(handler)` / `post(handler)` can see it
+        let http_fn_name = fn_name.clone();
+        let http_route_path = route_path.clone();
+        let http_app_handle = app_handle.clone();
+        let ws_fn_name = fn_name.clone();
+        let ws_route_path = route_path.clone();
+        let ws_app_handle = app_handle.clone();
+
+        // define handlers in THIS scope so route registration can capture them
         let handler = move |method: Method,
                             Path(ax_path): Path<HashMap<String, String>>,
                             Query(q): Query<HashMap<String, String>>,
                             headers: HeaderMap,
                             body: Option<AxumJson<J>>| {
-            let fn_name = fn_name.clone();
-            let route_path = route_path.clone();
-            let app_handle = app_handle.clone();
+            let fn_name = http_fn_name.clone();
+            let route_path = http_route_path.clone();
+            let app_handle = http_app_handle.clone();
             async move {
                 let mut h = HashMap::new();
                 for (k, v) in headers.iter() {
@@ -2970,6 +2993,7 @@ async fn run_app(app: WebAppHandle) -> RResult<()> {
                     headers: h,
                     cookies,
                     session_id,
+                    websocket: None,
                 };
 
                 let cb = get_global_callback().ok_or_else(|| {
@@ -2987,12 +3011,62 @@ async fn run_app(app: WebAppHandle) -> RResult<()> {
             }
         };
 
+        let ws_handler = move |ws: WebSocketUpgrade,
+                               Path(ax_path): Path<HashMap<String, String>>,
+                               Query(q): Query<HashMap<String, String>>,
+                               headers: HeaderMap| {
+            let fn_name = ws_fn_name.clone();
+            let route_path = ws_route_path.clone();
+            let app_handle = ws_app_handle.clone();
+            async move {
+                let mut h = HashMap::new();
+                for (k, v) in headers.iter() {
+                    if let Ok(s) = v.to_str() {
+                        h.insert(k.to_string(), s.to_string());
+                    }
+                }
+                let cookies = parse_cookie_header(h.get("cookie").cloned());
+                let session_id = cookies.get("relay_sid").cloned();
+
+                let response = ws.on_upgrade(move |socket| async move {
+                    let cb = match get_global_callback() {
+                        Some(cb) => cb,
+                        None => {
+                            eprintln!("WebSocket route handler callback is not set");
+                            return;
+                        }
+                    };
+
+                    let req = RequestParts {
+                        method: "WS".to_string(),
+                        path: route_path,
+                        path_params: ax_path,
+                        query: q,
+                        json: None,
+                        headers: h,
+                        cookies,
+                        session_id,
+                        websocket: Some(WebSocketHandle {
+                            inner: Arc::new(tokio::sync::Mutex::new(socket)),
+                        }),
+                    };
+
+                    if let Err(e) = cb(app_handle, fn_name, req).await {
+                        eprintln!("WebSocket handler error: {e}");
+                    }
+                });
+
+                Ok::<_, (StatusCode, String)>(response.into_response())
+            }
+        };
+
         router = match method.as_str() {
             "get" => router.route(&path, get(handler)),
             "post" => router.route(&path, post(handler)),
             "put" => router.route(&path, axum::routing::put(handler)),
             "patch" => router.route(&path, axum::routing::patch(handler)),
             "delete" => router.route(&path, axum::routing::delete(handler)),
+            "ws" => router.route(&path, get(ws_handler)),
             _ => {
                 return Err(RelayError::Runtime(format!(
                     "Unsupported HTTP method: {method}"
@@ -3711,6 +3785,61 @@ impl Object {
                 }
                 _ => Value::None,
             },
+            Object::WebSocket(ws) => match name {
+                "recv" => {
+                    let ws = ws.clone();
+                    Value::Builtin(Arc::new(move |_args, _| {
+                        let ws = ws.clone();
+                        Box::pin(async move {
+                            let mut socket = ws.inner.lock().await;
+                            match socket.next().await {
+                                Some(Ok(WsMessage::Text(s))) => Ok(Value::Str(s.to_string())),
+                                Some(Ok(WsMessage::Binary(b))) => Ok(Value::Bytes(b.to_vec())),
+                                Some(Ok(WsMessage::Ping(_))) | Some(Ok(WsMessage::Pong(_))) => {
+                                    Ok(Value::None)
+                                }
+                                Some(Ok(WsMessage::Close(_))) | None => Ok(Value::None),
+                                Some(Err(e)) => Err(RelayError::Runtime(format!(
+                                    "WebSocket receive error: {e}"
+                                ))),
+                            }
+                        })
+                    }))
+                }
+                "send" => {
+                    let ws = ws.clone();
+                    Value::Builtin(Arc::new(move |args, _| {
+                        let ws = ws.clone();
+                        Box::pin(async move {
+                            let msg = args.get(0).cloned().unwrap_or(Value::None);
+                            let frame = match msg {
+                                Value::Bytes(b) => WsMessage::Binary(b.into()),
+                                Value::None => WsMessage::Text(String::new().into()),
+                                other => WsMessage::Text(other.repr().into()),
+                            };
+                            let mut socket = ws.inner.lock().await;
+                            socket.send(frame).await.map_err(|e| {
+                                RelayError::Runtime(format!("WebSocket send error: {e}"))
+                            })?;
+                            Ok(Value::None)
+                        })
+                    }))
+                }
+                "close" => {
+                    let ws = ws.clone();
+                    Value::Builtin(Arc::new(move |_args, _| {
+                        let ws = ws.clone();
+                        Box::pin(async move {
+                            let mut socket = ws.inner.lock().await;
+                            socket.send(WsMessage::Close(None)).await.map_err(|e| {
+                                RelayError::Runtime(format!("WebSocket close error: {e}"))
+                            })?;
+                            Ok(Value::None)
+                        })
+                    }))
+                }
+                _ => Value::None,
+            },
             Object::Http(h) => h.get_member(name),
             Object::HttpResponse(r) => r.get_member(name),
             Object::MongoClient(m) => m.get_member(name),
@@ -3765,6 +3894,7 @@ impl Evaluator {
             headers,
             cookies,
             session_id,
+            websocket,
         } = req;
 
         let query_map = query.clone();
@@ -3773,6 +3903,10 @@ impl Evaluator {
         // build arg map by param list
         let mut bound: HashMap<String, Value> = HashMap::new();
 
+        if let Some(ws) = websocket.clone() {
+            bound.insert("socket".into(), Value::Obj(Object::WebSocket(ws)));
+        }
+
         // start with Query
         for (k, v) in query {
             bound.insert(k, Value::Str(v));
@@ -3877,6 +4011,9 @@ impl Evaluator {
                         .collect(),
                 ),
             );
+            if let Some(ws) = websocket.clone() {
+                env.set("socket", Value::Obj(Object::WebSocket(ws)));
+            }
             env.set("session", Value::Dict(existing_session.clone()));
         }
 
@@ -3932,6 +4069,13 @@ impl Evaluator {
             env.pop();
         }
 
+        if websocket.is_some() {
+            if let Some(sid) = session_id {
+                app.put_session(sid, final_session);
+            }
+            return Ok(Value::None);
+        }
+
         let mut response = match out {
             Value::Response(r) => r,
             Value::Json(j) => Response {