fiber-ws: fiber-aware WebSocket with httpd integration (Phase 5)

ober

7fbbd858494d9418433847abbed1788302445d3f

diff --git a/lib/std/net/fiber-httpd.sls b/lib/std/net/fiber-httpd.sls
index 5e148fc..f0c0220 100644
--- a/lib/std/net/fiber-httpd.sls
+++ b/lib/std/net/fiber-httpd.sls
@@ -58,7 +58,12 @@
 
     ;; Middleware
     wrap-health-check
-    wrap-metrics-endpoint)
+    wrap-metrics-endpoint
+
+    ;; WebSocket integration
+    make-websocket-response
+    websocket-response?
+    websocket-response-handler)
 
   (import (chezscheme)
           (std fiber)
@@ -115,6 +120,17 @@
       '(("Content-Type" . "text/html; charset=utf-8"))
       html))
 
+  ;; ========== WebSocket response ==========
+  ;;
+  ;; When a handler returns a websocket-response, the connection handler
+  ;; performs the WebSocket handshake and calls the ws-handler with
+  ;; (fd poller req) so it can create a fiber-ws and run a WS loop.
+  ;; The handler proc receives (fd poller req).
+
+  (define-record-type websocket-response
+    (fields (immutable handler))   ;; (lambda (fd poller req) ...)
+    (sealed #t))
+
   ;; ========== HTTP Parser ==========
   ;;
   ;; Reads HTTP/1.1 requests from a raw fd using fiber-aware I/O.
@@ -399,27 +415,39 @@
 
   ;; Connection handler: one fiber per connection, keep-alive loop
   (define (handle-connection fd poller handler metrics)
-    (let loop ()
-      (let ([req (read-request fd poller)])
-        (when req
-          (metrics-inc-requests! metrics)
-          (let ([resp (guard (exn [#t
-                        (metrics-inc-errors! metrics)
-                        (respond-text 500
-                          (if (message-condition? exn)
-                            (condition-message exn)
-                            "Internal Server Error"))])
-                        (handler req))])
-            (when (response? resp)
-              ;; Track 5xx errors
-              (when (>= (response-status resp) 500)
-                (metrics-inc-errors! metrics))
-              (write-response fd poller resp)
-              ;; Keep-alive: check Connection header
-              (let ([conn (request-header req "connection")])
-                (unless (and conn (string=? (string-downcase conn) "close"))
-                  (loop))))))))
-    (fiber-tcp-close fd))
+    (let ([ws-upgraded? #f])
+      (let loop ()
+        (let ([req (read-request fd poller)])
+          (when req
+            (metrics-inc-requests! metrics)
+            (let ([resp (guard (exn [#t
+                          (metrics-inc-errors! metrics)
+                          (respond-text 500
+                            (if (message-condition? exn)
+                              (condition-message exn)
+                              "Internal Server Error"))])
+                          (handler req))])
+              (cond
+                [(websocket-response? resp)
+                 ;; WebSocket upgrade — hand off fd to WS handler
+                 ;; The WS handler owns the fd now; don't close it here
+                 (set! ws-upgraded? #t)
+                 (guard (exn [#t (void)])
+                   ((websocket-response-handler resp) fd poller req))]
+                [(response? resp)
+                 ;; Track 5xx errors
+                 (when (>= (response-status resp) 500)
+                   (metrics-inc-errors! metrics))
+                 (write-response fd poller resp)
+                 ;; Keep-alive: check Connection header
+                 (let ([conn (request-header req "connection")])
+                   (unless (and conn (string=? (string-downcase conn) "close"))
+                     (loop)))]
+                ;; If resp is something else, just close
+                [else (void)])))))
+      ;; Only close fd if not handed off to WebSocket
+      (unless ws-upgraded?
+        (fiber-tcp-close fd))))
 
   ;; Accept loop with optional admission control
   (define (accept-loop listen-fd poller handler server)
diff --git a/lib/std/net/fiber-ws.sls b/lib/std/net/fiber-ws.sls
new file mode 100644
index 0000000..a260c2d
--- /dev/null
+++ b/lib/std/net/fiber-ws.sls
@@ -0,0 +1,201 @@
+#!chezscheme
+;;; (std net fiber-ws) — Fiber-aware WebSocket connections
+;;;
+;;; Builds on (std net websocket) for frame codec and (std net io) for
+;;; fiber-aware TCP I/O. One fiber per WebSocket connection.
+;;;
+;;; API:
+;;;   (make-fiber-ws fd poller)        — wrap an already-upgraded fd
+;;;   (fiber-ws? obj)                  — predicate
+;;;   (fiber-ws-open? ws)              — is connection open?
+;;;   (fiber-ws-recv ws)               — receive message (parks fiber)
+;;;                                       returns string, bytevector, or #f (close)
+;;;   (fiber-ws-send ws msg)           — send text message
+;;;   (fiber-ws-send-binary ws bv)     — send binary message
+;;;   (fiber-ws-close ws)              — send close frame and shut down
+;;;   (fiber-ws-ping ws)               — send ping (pong auto-handled)
+;;;
+;;;   (fiber-ws-upgrade req fd poller) — perform WebSocket handshake
+;;;                                       req-headers: alist from HTTP request
+;;;                                       returns fiber-ws or #f
+
+(library (std net fiber-ws)
+  (export
+    make-fiber-ws
+    fiber-ws?
+    fiber-ws-open?
+    fiber-ws-recv
+    fiber-ws-send
+    fiber-ws-send-binary
+    fiber-ws-close
+    fiber-ws-ping
+    fiber-ws-upgrade)
+
+  (import (chezscheme)
+          (std fiber)
+          (std net io)
+          (std net websocket))
+
+  ;; ========== Fiber WebSocket record ==========
+
+  (define-record-type fiber-ws
+    (fields
+      (immutable fd)
+      (immutable poller)
+      (mutable open?))
+    (sealed #t))
+
+  ;; ========== Low-level I/O ==========
+
+  ;; Read exactly n bytes from fd using fiber-aware I/O.
+  (define (read-exact fd buf n poller)
+    (let loop ([got 0])
+      (if (>= got n) got
+        (let* ([remaining (- n got)]
+               [tmp (make-bytevector (min 4096 remaining))])
+          (let ([rc (fiber-tcp-read fd tmp (min 4096 remaining) poller)])
+            (cond
+              [(<= rc 0) got]  ;; EOF or error — return partial
+              [else
+               (bytevector-copy! tmp 0 buf got rc)
+               (loop (+ got rc))]))))))
+
+  ;; Read a single WebSocket frame from the fd.
+  ;; Returns a ws-frame record or #f on connection close.
+  (define (read-ws-frame fd poller)
+    ;; Read header: first 2 bytes
+    (let ([hdr (make-bytevector 2)])
+      (let ([n (read-exact fd hdr 2 poller)])
+        (if (< n 2) #f
+          (let* ([b0 (bytevector-u8-ref hdr 0)]
+                 [b1 (bytevector-u8-ref hdr 1)]
+                 [fin? (not (zero? (bitwise-and b0 #x80)))]
+                 [opcode (bitwise-and b0 #x0F)]
+                 [masked? (not (zero? (bitwise-and b1 #x80)))]
+                 [len7 (bitwise-and b1 #x7F)])
+            ;; Extended length
+            (let ([payload-len
+                    (cond
+                      [(= len7 126)
+                       (let ([ext (make-bytevector 2)])
+                         (when (< (read-exact fd ext 2 poller) 2) (void))
+                         (bitwise-ior
+                           (bitwise-arithmetic-shift-left (bytevector-u8-ref ext 0) 8)
+                           (bytevector-u8-ref ext 1)))]
+                      [(= len7 127)
+                       (let ([ext (make-bytevector 8)])
+                         (when (< (read-exact fd ext 8 poller) 8) (void))
+                         ;; Use lower 32 bits
+                         (bitwise-ior
+                           (bitwise-arithmetic-shift-left (bytevector-u8-ref ext 4) 24)
+                           (bitwise-arithmetic-shift-left (bytevector-u8-ref ext 5) 16)
+                           (bitwise-arithmetic-shift-left (bytevector-u8-ref ext 6) 8)
+                           (bytevector-u8-ref ext 7)))]
+                      [else len7])])
+              ;; Masking key
+              (let ([mask-key (if masked?
+                                (let ([mk (make-bytevector 4)])
+                                  (read-exact fd mk 4 poller)
+                                  mk)
+                                #f)])
+                ;; Payload
+                (let ([payload (make-bytevector payload-len)])
+                  (when (> payload-len 0)
+                    (read-exact fd payload payload-len poller))
+                  ;; Unmask if needed
+                  (let ([data (if (and masked? mask-key)
+                                (ws-unmask-payload payload mask-key)
+                                payload)])
+                    (make-ws-frame fin? masked? opcode data mask-key))))))))))
+
+  ;; Write a ws-frame over the fd.
+  (define (write-ws-frame fd poller frame)
+    (let ([encoded (ws-frame-encode frame)])
+      (fiber-tcp-write fd encoded (bytevector-length encoded) poller)))
+
+  ;; ========== WebSocket upgrade handshake ==========
+
+  ;; Perform the server-side WebSocket handshake.
+  ;; req-headers: alist of (lowercase-name . value) from the HTTP request
+  ;; Returns: fiber-ws record, or #f on failure
+  (define (fiber-ws-upgrade req-headers fd poller)
+    (let ([ws-key (cond [(assoc "sec-websocket-key" req-headers) => cdr]
+                        [else #f])])
+      (if (not ws-key)
+        #f
+        (let* ([accept-key (ws-handshake-accept ws-key)]
+               [resp-str (string-append
+                           "HTTP/1.1 101 Switching Protocols\r\n"
+                           "Upgrade: websocket\r\n"
+                           "Connection: Upgrade\r\n"
+                           "Sec-WebSocket-Accept: " accept-key "\r\n"
+                           "\r\n")]
+               [resp-bv (string->bytevector resp-str (make-transcoder (utf-8-codec)))])
+          (fiber-tcp-write fd resp-bv (bytevector-length resp-bv) poller)
+          (make-fiber-ws fd poller #t)))))
+
+  ;; ========== Public API ==========
+
+  ;; Receive a message. Parks the fiber until a frame arrives.
+  ;; Returns: string (text), bytevector (binary), or #f (close/error).
+  (define (fiber-ws-recv ws)
+    (unless (fiber-ws-open? ws)
+      (error 'fiber-ws-recv "WebSocket is closed"))
+    (let ([frame (read-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws))])
+      (if (not frame)
+        (begin (fiber-ws-open?-set! ws #f) #f)
+        (let ([opcode (ws-frame-opcode frame)]
+              [payload (ws-frame-payload frame)])
+          (cond
+            [(= opcode ws-opcode-text)
+             (bytevector->string payload (make-transcoder (utf-8-codec)))]
+            [(= opcode ws-opcode-binary) payload]
+            [(= opcode ws-opcode-close)
+             ;; Send close back
+             (guard (exn [#t (void)])
+               (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws) (ws-close-frame)))
+             (fiber-ws-open?-set! ws #f)
+             #f]
+            [(= opcode ws-opcode-ping)
+             ;; Auto-pong
+             (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws)
+               (ws-pong-frame payload))
+             (fiber-ws-recv ws)]
+            [(= opcode ws-opcode-pong)
+             ;; Ignore, keep receiving
+             (fiber-ws-recv ws)]
+            [else
+             ;; Unknown opcode
+             (fiber-ws-open?-set! ws #f)
+             #f])))))
+
+  ;; Send a text message.
+  (define (fiber-ws-send ws msg)
+    (unless (fiber-ws-open? ws)
+      (error 'fiber-ws-send "WebSocket is closed"))
+    (let ([payload (string->bytevector msg (make-transcoder (utf-8-codec)))])
+      (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws)
+        (ws-text-frame payload))))
+
+  ;; Send a binary message.
+  (define (fiber-ws-send-binary ws bv)
+    (unless (fiber-ws-open? ws)
+      (error 'fiber-ws-send-binary "WebSocket is closed"))
+    (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws)
+      (ws-binary-frame bv)))
+
+  ;; Send a ping frame.
+  (define (fiber-ws-ping ws)
+    (when (fiber-ws-open? ws)
+      (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws)
+        (ws-ping-frame (make-bytevector 0)))))
+
+  ;; Close the WebSocket gracefully.
+  (define (fiber-ws-close ws)
+    (when (fiber-ws-open? ws)
+      (guard (exn [#t (void)])
+        (write-ws-frame (fiber-ws-fd ws) (fiber-ws-poller ws) (ws-close-frame)))
+      (fiber-ws-open?-set! ws #f)
+      (fiber-tcp-close (fiber-ws-fd ws))))
+
+) ;; end library
diff --git a/tests/test-fiber-ws.ss b/tests/test-fiber-ws.ss
new file mode 100644
index 0000000..33629b7
--- /dev/null
+++ b/tests/test-fiber-ws.ss
@@ -0,0 +1,230 @@
+;;; Tests for Phase 5: Fiber-aware WebSocket
+;;; Tests handshake, send/recv, ping/pong, close.
+
+(import (chezscheme))
+(import (std fiber))
+(import (std net io))
+(import (std net websocket))
+(import (std net fiber-ws))
+(import (std net fiber-httpd))
+(import (std text base64))
+(import (std crypto native-rust))
+
+(define test-count 0)
+(define pass-count 0)
+
+(define-syntax test
+  (syntax-rules ()
+    [(_ name body ...)
+     (begin
+       (set! test-count (+ test-count 1))
+       (guard (exn [#t
+         (display "FAIL: ") (display name) (newline)
+         (display "  Error: ")
+         (display (if (message-condition? exn) (condition-message exn) exn))
+         (newline)])
+         body ...
+         (set! pass-count (+ pass-count 1))
+         (display "PASS: ") (display name) (newline)))]))
+
+(define-syntax assert-equal
+  (syntax-rules ()
+    [(_ got expected msg)
+     (unless (equal? got expected)
+       (error 'assert msg (list 'got: got 'expected: expected)))]))
+
+(define-syntax assert-true
+  (syntax-rules ()
+    [(_ val msg)
+     (unless val (error 'assert msg))]))
+
+;; Helper: send raw WS upgrade request, read 101 response
+(define (ws-client-upgrade fd poller)
+  (let* ([ws-key (u8vector->base64-string (rust-random-bytes 16))]
+         [req-str (string-append
+                    "GET /ws HTTP/1.1\r\n"
+                    "Host: localhost\r\n"
+                    "Upgrade: websocket\r\n"
+                    "Connection: Upgrade\r\n"
+                    "Sec-WebSocket-Key: " ws-key "\r\n"
+                    "Sec-WebSocket-Version: 13\r\n"
+                    "\r\n")]
+         [req-bv (string->bytevector req-str (make-transcoder (utf-8-codec)))])
+    (fiber-tcp-write fd req-bv (bytevector-length req-bv) poller)
+    ;; Read 101 response
+    (let ([buf (make-bytevector 4096)])
+      (let ([n (fiber-tcp-read fd buf 4096 poller)])
+        (> n 0)))))
+
+;; Helper: send a masked text frame (client→server must be masked)
+(define (ws-client-send-text fd poller msg)
+  (let* ([payload (string->bytevector msg (make-transcoder (utf-8-codec)))]
+         [mask-key (rust-random-bytes 4)]
+         [frame (make-ws-frame #t #t ws-opcode-text payload mask-key)]
+         [encoded (ws-frame-encode frame)])
+    (fiber-tcp-write fd encoded (bytevector-length encoded) poller)))
+
+;; Helper: read a server frame and return payload as string
+(define (ws-client-recv-text fd poller)
+  (let ([buf (make-bytevector 4096)])
+    (let ([n (fiber-tcp-read fd buf 4096 poller)])
+      (if (<= n 0) #f
+        (let* ([bv (let ([b (make-bytevector n)])
+                     (bytevector-copy! buf 0 b 0 n) b)]
+               [frame (ws-frame-decode bv)])
+          (bytevector->string (ws-frame-payload frame)
+            (make-transcoder (utf-8-codec))))))))
+
+;; Helper: send masked close frame
+(define (ws-client-close fd poller)
+  (let* ([frame (make-ws-frame #t #t ws-opcode-close
+                  (make-bytevector 0) (rust-random-bytes 4))]
+         [encoded (ws-frame-encode frame)])
+    (fiber-tcp-write fd encoded (bytevector-length encoded) poller)))
+
+;; Standard WS echo handler for httpd
+(define (make-ws-echo-httpd-handler)
+  (lambda (req)
+    (let ([upgrade (request-header req "upgrade")])
+      (if (and upgrade (string=? (string-downcase upgrade) "websocket"))
+        (make-websocket-response
+          (lambda (fd poller req)
+            (let ([ws (fiber-ws-upgrade (request-headers req) fd poller)])
+              (when ws
+                (let loop ()
+                  (let ([msg (fiber-ws-recv ws)])
+                    (when msg
+                      (if (string? msg)
+                        (fiber-ws-send ws msg)
+                        (fiber-ws-send-binary ws msg))
+                      (loop))))
+                (fiber-ws-close ws)))))
+        (respond-text 200 "not a websocket")))))
+
+;; =========================================================================
+;; Test 1: WebSocket handshake key computation (RFC 6455 test vector)
+;; =========================================================================
+
+(test "ws handshake accept key"
+  (let ([accept (ws-handshake-accept "dGhlIHNhbXBsZSBub25jZQ==")])
+    (assert-equal accept "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" "RFC 6455 test vector")))
+
+;; =========================================================================
+;; Test 2: WebSocket frame encode/decode round-trip
+;; =========================================================================
+
+(test "ws frame encode/decode round-trip"
+  (let* ([payload (string->bytevector "Hello" (make-transcoder (utf-8-codec)))]
+         [frame (ws-text-frame payload)]
+         [encoded (ws-frame-encode frame)]
+         [decoded (ws-frame-decode encoded)])
+    (assert-true (ws-frame-fin? decoded) "FIN set")
+    (assert-equal (ws-frame-opcode decoded) ws-opcode-text "opcode text")
+    (assert-equal (ws-frame-payload decoded) payload "payload matches")))
+
+;; =========================================================================
+;; Test 3: Masked frame round-trip
+;; =========================================================================
+
+(test "ws masked frame round-trip"
+  (let* ([payload (string->bytevector "Masked!" (make-transcoder (utf-8-codec)))]
+         [mask-key (make-bytevector 4)])
+    (bytevector-u8-set! mask-key 0 #x37)
+    (bytevector-u8-set! mask-key 1 #xFA)
+    (bytevector-u8-set! mask-key 2 #x21)
+    (bytevector-u8-set! mask-key 3 #x3D)
+    (let* ([frame (make-ws-frame #t #t ws-opcode-text payload mask-key)]
+           [encoded (ws-frame-encode frame)]
+           [decoded (ws-frame-decode encoded)])
+      (assert-true (ws-frame-masked? decoded) "masked")
+      (assert-equal (ws-frame-payload decoded) payload "unmasked correctly"))))
+
+;; =========================================================================
+;; Test 4: Full WebSocket echo via fiber-httpd
+;; =========================================================================
+
+(test "WebSocket echo through fiber-httpd"
+  (let* ([handler (make-ws-echo-httpd-handler)]
+         [srv (fiber-httpd-start 0 handler)]
+         [port (fiber-httpd-listen-port srv)]
+         [result-box (box #f)])
+
+    (sleep (make-time 'time-duration 100000000 0))
+
+    (let ([rt (make-fiber-runtime 2)])
+      (with-io-poller rt poller
+        (fiber-spawn rt
+          (lambda ()
+            (let ([fd (fiber-tcp-connect "127.0.0.1" port poller)])
+              (ws-client-upgrade fd poller)
+              (ws-client-send-text fd poller "hello-ws")
+              (let ([reply (ws-client-recv-text fd poller)])
+                (set-box! result-box reply))
+              (ws-client-close fd poller)
+              (fiber-tcp-close fd)))
+          "ws-client")
+        (fiber-runtime-run! rt)))
+
+    (fiber-httpd-stop! srv)
+    (assert-equal (unbox result-box) "hello-ws" "echo round-trip")))
+
+;; =========================================================================
+;; Test 5: Multiple WebSocket messages
+;; =========================================================================
+
+(test "5 sequential WebSocket messages"
+  (let* ([handler (lambda (req)
+                    (let ([upgrade (request-header req "upgrade")])
+                      (if (and upgrade (string=? (string-downcase upgrade) "websocket"))
+                        (make-websocket-response
+                          (lambda (fd poller req)
+                            (let ([ws (fiber-ws-upgrade (request-headers req) fd poller)])
+                              (when ws
+                                (let loop ()
+                                  (let ([msg (fiber-ws-recv ws)])
+                                    (when msg
+                                      (fiber-ws-send ws (string-append "echo:" msg))
+                                      (loop))))
+                                (fiber-ws-close ws)))))
+                        (respond-text 200 "http"))))]
+         [srv (fiber-httpd-start 0 handler)]
+         [port (fiber-httpd-listen-port srv)]
+         [results (make-vector 5 #f)])
+
+    (sleep (make-time 'time-duration 100000000 0))
+
+    (let ([rt (make-fiber-runtime 2)])
+      (with-io-poller rt poller
+        (fiber-spawn rt
+          (lambda ()
+            (let ([fd (fiber-tcp-connect "127.0.0.1" port poller)])
+              (ws-client-upgrade fd poller)
+              (do ([i 0 (+ i 1)])
+                ((= i 5))
+                (let ([msg (string-append "msg-" (number->string i))])
+                  (ws-client-send-text fd poller msg)
+                  (let ([reply (ws-client-recv-text fd poller)])
+                    (vector-set! results i
+                      (and reply (string=? reply (string-append "echo:" msg)))))))
+              (ws-client-close fd poller)
+              (fiber-tcp-close fd)))
+          "ws-multi")
+        (fiber-runtime-run! rt)))
+
+    (fiber-httpd-stop! srv)
+
+    (do ([i 0 (+ i 1)])
+      ((= i 5))
+      (assert-true (vector-ref results i)
+        (string-append "message " (number->string i))))))
+
+;; =========================================================================
+;; Summary
+;; =========================================================================
+(newline)
+(display "=========================================") (newline)
+(display "Results: ") (display pass-count) (display "/")
+(display test-count) (display " passed") (newline)
+(display "=========================================") (newline)
+(when (< pass-count test-count)
+  (exit 1))