fix: replay protection on the pull protocol (AAD-bound seq + per-direction keys)

ober

ce4631822970760a351d9ba517bd39aa6d166f59

diff --git a/bin/collector.ss b/bin/collector.ss
index 3b0259b..d96a8f0 100644
--- a/bin/collector.ss
+++ b/bin/collector.ss
@@ -84,6 +84,27 @@
           (error 'recv-msg! "connection closed before frame body"))
         (psk-decrypt-transport psk-auth encrypted)))))
 
+;; Post-auth requests ride the replay-protected session (monotonic seq in the
+;; AAD + per-direction keys), mirroring the agent's listener.
+(define (send-session-msg! port session payload)
+  (let* ([encrypted (psk-session-seal session payload)]
+         [len-bytes (pack-u32-le (bytevector-length encrypted))])
+    (put-bytevector port len-bytes)
+    (put-bytevector port encrypted)
+    (flush-output-port port)))
+
+(define (recv-session-msg! port session)
+  (let* ([len-bv (get-bytevector-n port 4)])
+    (unless (and (bytevector? len-bv) (= (bytevector-length len-bv) 4))
+      (error 'recv-session-msg! "connection closed before frame length"))
+    (let ([len (unpack-u32-le len-bv 0)])
+      (when (> len 10485760)
+        (error 'recv-session-msg! "message too large"))
+      (let ([encrypted (get-bytevector-n port len)])
+        (unless (and (bytevector? encrypted) (= (bytevector-length encrypted) len))
+          (error 'recv-session-msg! "connection closed before frame body"))
+        (psk-session-open session encrypted)))))
+
 ;; --- Authenticate to agent ---
 
 (define (authenticate! in-port out-port psk-auth)
@@ -101,7 +122,8 @@
              [resp-bv (psk-response->bytevector response)]
              [msg (bv-tag MSG-CHALLENGE-RESPONSE resp-bv)])
         (send-msg! out-port psk-auth msg)
-        #t))))
+        (make-client-transport-session psk-auth
+          (psk-derive-channel-id (psk-challenge-nonce challenge)))))))
 
 (define (bv-tag tag . parts)
   (let* ([total (+ 1 (apply + (map bytevector-length parts)))]
@@ -115,11 +137,11 @@
 
 ;; --- Send request and get response ---
 
-(define (send-request! in-port out-port psk-auth req-tag . data-bvs)
+(define (send-request! in-port out-port session req-tag . data-bvs)
   (let ([request (apply bv-tag MSG-REQUEST
                    (make-bytevector 1 req-tag) data-bvs)])
-    (send-msg! out-port psk-auth request)
-    (let ([resp (recv-msg! in-port psk-auth)])
+    (send-session-msg! out-port session request)
+    (let ([resp (recv-session-msg! in-port session)])
       (unless resp
         (error 'send-request "no response"))
       (let ([tag (bytevector-u8-ref resp 0)])
@@ -271,12 +293,12 @@
     (let ([psk-auth (make-psk-auth psk)])
       (let-values ([(host port) (parse-host-port (car args))])
         (let-values ([(in out) (connect-to-agent host port)])
-          (authenticate! in out psk-auth)
-          (let ([resp (send-request! in out psk-auth REQ-STATUS)])
-            (displayln (format "Buffered events: ~a" (cadr resp)))
-            (displayln (format "Latest sequence: ~a" (caddr resp)))
-            (displayln (format "Uptime: ~as" (cadddr resp))))
-          (close-port out))))))
+          (let ([session (authenticate! in out psk-auth)])
+            (let ([resp (send-request! in out session REQ-STATUS)])
+              (displayln (format "Buffered events: ~a" (cadr resp)))
+              (displayln (format "Latest sequence: ~a" (caddr resp)))
+              (displayln (format "Uptime: ~as" (cadddr resp))))
+            (close-port out)))))))
 
 (define (run-poll args)
   (let ([format-json? (member "--format" args)]
@@ -302,16 +324,16 @@
             [decryptor (make-ecies-decryptor private-key)])
         (let-values ([(host port) (parse-host-port host-port)])
           (let-values ([(in out) (connect-to-agent host port)])
-            (authenticate! in out psk-auth)
-            (let ([resp (send-request! in out psk-auth
-                          REQ-GET-EVENTS-AFTER (pack-u64-le after-seq))])
-              (when (and (pair? resp) (eq? (car resp) 'events))
-                (for-each
-                  (lambda (stored-ev)
-                    (let ([ev (decrypt-event decryptor stored-ev)])
-                      (when ev (print-event ev format-json?))))
-                  (cadr resp))))
-            (close-port out)))))))
+            (let ([session (authenticate! in out psk-auth)])
+              (let ([resp (send-request! in out session
+                            REQ-GET-EVENTS-AFTER (pack-u64-le after-seq))])
+                (when (and (pair? resp) (eq? (car resp) 'events))
+                  (for-each
+                    (lambda (stored-ev)
+                      (let ([ev (decrypt-event decryptor stored-ev)])
+                        (when ev (print-event ev format-json?))))
+                    (cadr resp))))
+              (close-port out))))))))
 
 (define (run-watch args)
   (let ([format-json? (member "--format" args)]
@@ -366,27 +388,27 @@
                      (loop))])
           (let-values ([(in out) (connect-to-agent host port)])
             (set-box! cur-out out)
-            (authenticate! in out psk-auth)
-            (set-box! backoff 1)
-            (let poll-loop ()
-              (let ([resp (send-request! in out psk-auth
-                            REQ-GET-EVENTS-AFTER
-                            (pack-u64-le (unbox last-seq)))])
-                (when (and (pair? resp) (eq? (car resp) 'events))
-                  (for-each
-                    (lambda (stored-ev)
-                      (let ([ev (decrypt-event decryptor stored-ev)])
-                        (when ev
-                          (print-event ev format-json?)
-                          (when store
-                            (dynamic-wind
-                              (lambda () (mutex-acquire store-mtx))
-                              (lambda () (store-decrypted-event! store stored-ev ev host-str))
-                              (lambda () (mutex-release store-mtx)))))
-                        (set-box! last-seq (stored-event-seq stored-ev))))
-                    (cadr resp))))
-              (sleep (make-time 'time-duration 0 1))
-              (poll-loop))))))))
+            (let ([session (authenticate! in out psk-auth)])
+              (set-box! backoff 1)
+              (let poll-loop ()
+                (let ([resp (send-request! in out session
+                              REQ-GET-EVENTS-AFTER
+                              (pack-u64-le (unbox last-seq)))])
+                  (when (and (pair? resp) (eq? (car resp) 'events))
+                    (for-each
+                      (lambda (stored-ev)
+                        (let ([ev (decrypt-event decryptor stored-ev)])
+                          (when ev
+                            (print-event ev format-json?)
+                            (when store
+                              (dynamic-wind
+                                (lambda () (mutex-acquire store-mtx))
+                                (lambda () (store-decrypted-event! store stored-ev ev host-str))
+                                (lambda () (mutex-release store-mtx)))))
+                          (set-box! last-seq (stored-event-seq stored-ev))))
+                      (cadr resp))))
+                (sleep (make-time 'time-duration 0 1))
+                (poll-loop)))))))))
 
 (define (store-decrypted-event! store stored-ev ev host-str)
   (let ([data (security-event-data ev)])
diff --git a/lib/secmon/crypto/psk.sls b/lib/secmon/crypto/psk.sls
index be9dcf8..568ab56 100644
--- a/lib/secmon/crypto/psk.sls
+++ b/lib/secmon/crypto/psk.sls
@@ -7,7 +7,12 @@
     psk-challenge? psk-challenge-nonce psk-challenge-timestamp
     psk-response? psk-response-proof psk-response-counter-nonce
     psk-challenge->bytevector bytevector->psk-challenge
-    psk-response->bytevector bytevector->psk-response)
+    psk-response->bytevector bytevector->psk-response
+    psk-derive-transport-direction-key psk-derive-channel-id
+    make-psk-transport-session psk-transport-session?
+    make-server-transport-session make-client-transport-session
+    psk-session-seal psk-session-open
+    psk-session-send-seq psk-session-recv-top)
   (import
     (chezscheme)
     (jerboa prelude clean)
@@ -117,6 +122,103 @@
         (bytevector-copy! data 12 ct 0 ct-len)
         (rust-aead-open (psk-auth-transport-key auth) nonce ct #vu8()))))
 
+  ;; --- Replay-protected transport session ---
+  ;;
+  ;; The post-auth request loop must not accept a captured frame twice. Each
+  ;; direction gets its own HKDF-derived key (domain separation, so a frame
+  ;; sealed for one direction cannot decrypt as the other — reflection is
+  ;; rejected) and a monotonic sequence number bound into the AEAD AAD as
+  ;; direction ‖ channel-id ‖ seq. The receiver keeps a high-water mark and
+  ;; rejects any seq it has already advanced past, so a replayed or stale frame
+  ;; fails before it can purge buffered evidence. The channel-id binds every
+  ;; frame to the handshake that established the session, defeating
+  ;; cross-context splicing.
+
+  (define (psk-derive-transport-direction-key psk info)
+    (hkdf-sha256 psk #f (string->utf8 info) 32))
+
+  (define (psk-derive-channel-id challenge-nonce)
+    (rust-sha256 challenge-nonce))
+
+  (define-record-type psk-transport-session
+    (fields
+      (immutable send-key psk-session-send-key)
+      (immutable recv-key psk-session-recv-key)
+      (immutable send-dir psk-session-send-dir)
+      (immutable recv-dir psk-session-recv-dir)
+      (immutable channel-id psk-session-channel-id)
+      (mutable send-seq psk-session-send-seq psk-session-send-seq-set!)
+      (mutable recv-top psk-session-recv-top psk-session-recv-top-set!))
+    (nongenerative psk-transport-session-type)
+    (protocol
+      (lambda (new)
+        (lambda (send-key recv-key send-dir recv-dir channel-id)
+          (new send-key recv-key send-dir recv-dir channel-id 1 0)))))
+
+  (define (psk-transport-aad dir-byte channel-id seq)
+    (let* ([cid-len (bytevector-length channel-id)]
+           [aad (make-bytevector (+ 1 cid-len 8))])
+      (bytevector-u8-set! aad 0 dir-byte)
+      (bytevector-copy! channel-id 0 aad 1 cid-len)
+      (bytevector-u64-set! aad (+ 1 cid-len) seq (endianness little))
+      aad))
+
+  ;; Frame layout: seq(8 LE) ‖ nonce(12) ‖ ciphertext‖tag. The seq rides in the
+  ;; clear so the receiver can rebuild the AAD before opening, and is itself
+  ;; authenticated by the tag.
+  (define (psk-session-seal session plaintext)
+    (let* ([seq (psk-session-send-seq session)]
+           [_ (psk-session-send-seq-set! session (+ seq 1))]
+           [nonce (rust-random-bytes 12)]
+           [aad (psk-transport-aad (psk-session-send-dir session)
+                  (psk-session-channel-id session) seq)]
+           [ct (rust-aead-seal (psk-session-send-key session) nonce plaintext aad)]
+           [out (make-bytevector (+ 20 (bytevector-length ct)))])
+      (bytevector-u64-set! out 0 seq (endianness little))
+      (bytevector-copy! nonce 0 out 8 12)
+      (bytevector-copy! ct 0 out 20 (bytevector-length ct))
+      out))
+
+  ;; Open a frame, rejecting anything stale or replayed. Returns the plaintext
+  ;; bytevector, or #f on a short frame, a non-fresh seq, or a failed tag
+  ;; (wrong direction key, wrong channel-id, or tampering). The high-water mark
+  ;; advances only after a successful open, so a forged seq cannot poison it.
+  (define (psk-session-open session frame)
+    (let ([n (bytevector-length frame)])
+      (and (>= n 36) ;; seq(8) + nonce(12) + GCM tag(16)
+        (let* ([seq (bytevector-u64-ref frame 0 (endianness little))]
+               [top (psk-session-recv-top session)])
+          (and (> seq top) ;; fresh, strictly monotonic: rejects stale + replay
+            (let* ([nonce (make-bytevector 12)]
+                   [_ (bytevector-copy! frame 8 nonce 0 12)]
+                   [ct-len (- n 20)]
+                   [ct (make-bytevector ct-len)]
+                   [_ (bytevector-copy! frame 20 ct 0 ct-len)]
+                   [aad (psk-transport-aad (psk-session-recv-dir session)
+                          (psk-session-channel-id session) seq)]
+                   ;; rust-aead-open raises on a bad tag; a reflected, spliced,
+                   ;; or tampered frame must surface as #f, not an exception.
+                   [pt (guard (e [#t #f])
+                         (rust-aead-open (psk-session-recv-key session) nonce ct aad))])
+              (and pt
+                (begin (psk-session-recv-top-set! session seq) pt))))))))
+
+  ;; Direction byte 0 = client→server, 1 = server→client. The server sends on
+  ;; s2c and receives on c2s; the client is the mirror.
+  (define (make-server-transport-session auth channel-id)
+    (let ([psk (psk-auth-psk auth)])
+      (make-psk-transport-session
+        (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:s2c")
+        (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:c2s")
+        1 0 channel-id)))
+
+  (define (make-client-transport-session auth channel-id)
+    (let ([psk (psk-auth-psk auth)])
+      (make-psk-transport-session
+        (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:c2s")
+        (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:s2c")
+        0 1 channel-id)))
+
   ;; --- Serialization ---
 
   (define (psk-challenge->bytevector c)
diff --git a/lib/secmon/server/listener.sls b/lib/secmon/server/listener.sls
index c211cca..4abbb8c 100644
--- a/lib/secmon/server/listener.sls
+++ b/lib/secmon/server/listener.sls
@@ -208,6 +208,22 @@
         (poll-server-psk-auth server)
         (tcp-read-exact fd port len deadline-ms))))
 
+  ;; Post-auth frames ride a replay-protected session: a monotonic seq bound
+  ;; into the AAD plus per-direction keys. A captured request cannot be
+  ;; re-submitted (stale seq) nor reflected (wrong direction key).
+  (define (send-session-message! fd session payload deadline-ms)
+    (let* ([encrypted (psk-session-seal session payload)]
+           [len-bytes (pack-u32-le (bytevector-length encrypted))])
+      (tcp-write-all fd len-bytes deadline-ms)
+      (tcp-write-all fd encrypted deadline-ms)))
+
+  (define (recv-session-message! fd port session deadline-ms)
+    (let* ([len-bytes (tcp-read-exact fd port 4 deadline-ms)]
+           [len (unpack-u32-le len-bytes 0)])
+      (when (or (= len 0) (> len MAX-MESSAGE-SIZE))
+        (error 'recv-session-message! "invalid encrypted message length" len))
+      (psk-session-open session (tcp-read-exact fd port len deadline-ms))))
+
   (define (handle-request server req)
     (case (car req)
       [(get-events-after)
@@ -256,23 +272,28 @@
                                (poll-server-psk-auth server) challenge response
                                (poll-server-max-challenge-age server)))
                       (send-message! fd server (pack-auth-failed) deadline)
-                      (let request-loop ()
-                        (let* ([idle-ms (listener-policy-idle-ms policy)]
-                               [_ (tcp:set-socket-timeout! fd idle-ms)]
-                               [request-deadline (+ (current-time-ms) idle-ms)]
-                               [req-bytes
-                                 (recv-message! fd in-port server request-deadline)])
-                          (when (and (bytevector? req-bytes)
-                                     (> (bytevector-length req-bytes) 0)
-                                     (= (bytevector-u8-ref req-bytes 0) MSG-REQUEST))
-                            (let* ([req-data (make-bytevector
-                                              (- (bytevector-length req-bytes) 1))]
-                                   [_ (bytevector-copy! req-bytes 1 req-data 0
-                                        (bytevector-length req-data))]
-                                   [req (unpack-request req-data)])
-                              (send-message! fd server
-                                (handle-request server req) request-deadline)
-                              (request-loop)))))))))))
+                      (let ([session (make-server-transport-session
+                                       (poll-server-psk-auth server)
+                                       (psk-derive-channel-id
+                                         (psk-challenge-nonce challenge)))])
+                        (let request-loop ()
+                          (let* ([idle-ms (listener-policy-idle-ms policy)]
+                                 [_ (tcp:set-socket-timeout! fd idle-ms)]
+                                 [request-deadline (+ (current-time-ms) idle-ms)]
+                                 [req-bytes
+                                   (recv-session-message! fd in-port session
+                                     request-deadline)])
+                            (when (and (bytevector? req-bytes)
+                                       (> (bytevector-length req-bytes) 0)
+                                       (= (bytevector-u8-ref req-bytes 0) MSG-REQUEST))
+                              (let* ([req-data (make-bytevector
+                                                (- (bytevector-length req-bytes) 1))]
+                                     [_ (bytevector-copy! req-bytes 1 req-data 0
+                                          (bytevector-length req-data))]
+                                     [req (unpack-request req-data)])
+                                (send-session-message! fd session
+                                  (handle-request server req) request-deadline)
+                                (request-loop))))))))))))
           (lambda ()
             (guard (e [#t (void)]) (close-port in-port))))))
 
diff --git a/tests/transport-replay-test.ss b/tests/transport-replay-test.ss
new file mode 100644
index 0000000..49ebeaf
--- /dev/null
+++ b/tests/transport-replay-test.ss
@@ -0,0 +1,68 @@
+(import (chezscheme)
+        (secmon crypto psk))
+
+(define failures 0)
+(define (check name got want)
+  (let ([ok (equal? got want)])
+    (unless ok (set! failures (+ failures 1)))
+    (display (if ok "ok: " "FAIL: "))
+    (display name)
+    (unless ok (display (format " got=~s want=~s" got want)))
+    (newline)))
+
+(define psk (make-bytevector 32 #x42))
+(define auth (make-psk-auth psk))
+(define channel-id (psk-derive-channel-id (make-bytevector 32 #x07)))
+(define msg1 (string->utf8 "request-one"))
+(define msg2 (string->utf8 "request-two"))
+
+(define (fresh-pair)
+  (values (make-client-transport-session auth channel-id)
+          (make-server-transport-session auth channel-id)))
+
+;; Regression guard: fresh, in-order frames are accepted in both directions.
+(let-values ([(client server) (fresh-pair)])
+  (let* ([f1 (psk-session-seal client msg1)]   ;; seq 1
+         [f2 (psk-session-seal client msg2)])  ;; seq 2
+    (check "fresh in-order frame 1 accepted" (psk-session-open server f1) msg1)
+    (check "fresh in-order frame 2 accepted" (psk-session-open server f2) msg2))
+  (let ([r1 (psk-session-seal server (string->utf8 "reply"))])
+    (check "server->client frame accepted"
+           (psk-session-open client r1) (string->utf8 "reply"))))
+
+;; Replay: the same captured ciphertext is rejected the second time.
+(let-values ([(client server) (fresh-pair)])
+  (let ([f1 (psk-session-seal client msg1)])
+    (check "first submit accepted" (psk-session-open server f1) msg1)
+    (check "replayed frame rejected" (psk-session-open server f1) #f)))
+
+;; Stale/lower seq: a frame older than the high-water mark is rejected.
+(let-values ([(client server) (fresh-pair)])
+  (let* ([f1 (psk-session-seal client msg1)]   ;; seq 1
+         [f2 (psk-session-seal client msg2)])  ;; seq 2
+    (check "newer frame accepted first" (psk-session-open server f2) msg2)
+    (check "stale lower-seq frame rejected" (psk-session-open server f1) #f)))
+
+;; Reflection: a frame sealed for the opposite direction must not open.
+(let-values ([(client server) (fresh-pair)])
+  (let ([s2c (psk-session-seal server msg1)])  ;; sealed s2c (dir 1, s2c key)
+    (check "reflected frame rejected" (psk-session-open server s2c) #f)
+    (check "intended receiver accepts s2c" (psk-session-open client s2c) msg1)))
+
+;; Cross-context splicing: a frame from another channel does not open.
+(let ([other-channel (psk-derive-channel-id (make-bytevector 32 #x08))])
+  (let ([client-a (make-client-transport-session auth channel-id)]
+        [server-b (make-server-transport-session auth other-channel)])
+    (let ([fa (psk-session-seal client-a msg1)])
+      (check "cross-channel frame rejected" (psk-session-open server-b fa) #f))))
+
+;; Direction keys are domain-separated (distinct HKDF info -> distinct keys).
+(check "c2s and s2c keys differ"
+       (equal? (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:c2s")
+               (psk-derive-transport-direction-key psk "secmon-psk-transport-v2:s2c"))
+       #f)
+
+(if (= failures 0)
+  (begin (display "transport-replay-test: all replay/reflection checks passed") (newline))
+  (begin (display (format "transport-replay-test: ~a failures" failures))
+         (newline) (exit 1)))