fiber-ws: fiber-aware WebSocket with httpd integration (Phase 5)
ober
7fbbd858494d9418433847abbed1788302445d3f
--- 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) new file mode 100644 --- /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 new file mode 100644 --- /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))