srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/ipc/ipc.go20
-rw-r--r--internal/ipc/server.go8
-rw-r--r--internal/proxy/proxy.go156
-rw-r--r--internal/proxy/websocket.go159
-rw-r--r--internal/proxy/websocket_test.go159
-rw-r--r--internal/store/store.go61
6 files changed, 557 insertions, 6 deletions
diff --git a/internal/ipc/ipc.go b/internal/ipc/ipc.go
index 68217c2..dc7a1df 100644
--- a/internal/ipc/ipc.go
+++ b/internal/ipc/ipc.go
@@ -114,6 +114,7 @@ type Response struct {
Rules []rules.Rule `json:"rules,omitempty"` // for "rules"
ScopeRules []scope.Rule `json:"scope_rules,omitempty"` // for "scope_rules"
ClientCerts []clientcert.Cert `json:"client_certs,omitempty"` // for "client_certs"
+ WSMessages []store.WSMessage `json:"ws_messages,omitempty"` // for "ws_messages"
Status *StatusMsg `json:"status,omitempty"` // for "status"
// For "import_done": how many entries were actually inserted (a
@@ -331,6 +332,25 @@ func (c *Client) Get(id int64) (*EntryDetail, error) {
return resp.Detail, nil
}
+// ListWSMessages returns every WebSocket frame captured for entryID's
+// connection, in the order they were sent - empty (not an error) if the
+// entry wasn't a WebSocket upgrade or nothing was captured.
+func (c *Client) ListWSMessages(entryID int64) ([]store.WSMessage, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ if err := c.enc.Encode(Request{Type: "ws_messages", ID: entryID}); err != nil {
+ return nil, err
+ }
+ var resp Response
+ if err := c.dec.Decode(&resp); err != nil {
+ return nil, err
+ }
+ if resp.Type == "error" {
+ return nil, errors.New(resp.Error)
+ }
+ return resp.WSMessages, nil
+}
+
// Repeat sends raw to scheme://host exactly as given (no re-serialization,
// no header injection) and returns the resulting entry, including the raw
// response bytes. The exchange is also recorded to history.
diff --git a/internal/ipc/server.go b/internal/ipc/server.go
index 4db27b9..378c1a9 100644
--- a/internal/ipc/server.go
+++ b/internal/ipc/server.go
@@ -143,6 +143,14 @@ func (s *Server) handleConn(conn net.Conn) {
}
enc.Encode(Response{Type: "get", Detail: detailFromEntry(e)})
+ case "ws_messages":
+ msgs, err := s.db.ListWSMessages(req.ID)
+ if err != nil {
+ enc.Encode(Response{Type: "error", Error: err.Error()})
+ continue
+ }
+ enc.Encode(Response{Type: "ws_messages", WSMessages: msgs})
+
case "repeat":
if s.repeater == nil {
enc.Encode(Response{Type: "error", Error: "repeater not available"})
diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go
index a98a35a..382205f 100644
--- a/internal/proxy/proxy.go
+++ b/internal/proxy/proxy.go
@@ -22,6 +22,7 @@ package proxy
import (
"bufio"
+ "bytes"
"context"
"crypto/tls"
"errors"
@@ -494,7 +495,11 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr
outReq.URL.Scheme = scheme
outReq.URL.Host = hostname
outReq.RequestURI = ""
- stripHopByHop(outReq.Header)
+ if isWebSocketUpgradeRequest(r) {
+ stripHopByHopKeepingUpgrade(outReq.Header)
+ } else {
+ stripHopByHop(outReq.Header)
+ }
// Header rules are applied to outReq only, after cloning and header
// stripping - history's request_raw keeps showing what the client
// actually sent (clientTee/reqBodyCap already capture from r, not
@@ -591,6 +596,17 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr
}
defer resp.Body.Close()
+ // A WebSocket upgrade stops being one-shot request/response the
+ // instant the 101 lands - match-and-replace rules, body capture, and
+ // the normal write-response-then-record flow below all assume a
+ // bounded response with a body, none of which applies here. HTTP/2
+ // client legs are excluded: they can't be hijacked for raw access
+ // the way an HTTP/1.1 connection can (see handleWebSocketUpgrade).
+ if negotiated != http2.NextProtoTLS && isWebSocketUpgradeResponse(resp) {
+ s.handleWebSocketUpgrade(w, r, scheme, hostname, started, duration, reqRaw, reqExact, reqTrunc, resp, upstreamTee)
+ return
+ }
+
respRules, err := s.enabledRules("response")
if err != nil {
log.Printf("load response rules: %v", err)
@@ -653,6 +669,115 @@ func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWr
s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, reqTrunc, respRaw, respExact, respTrunc, resp.StatusCode, "")
}
+// handleWebSocketUpgrade takes over the connection after resp (a 101
+// matching isWebSocketUpgradeResponse) comes back from the origin. It
+// records the upgrade request/response pair to history exactly like a
+// normal exchange, then relays WebSocket frames bidirectionally, byte-
+// for-byte unmodified, until either side closes - decoding each frame's
+// payload along the way for capture into the ws_messages table, tagged
+// to this exchange's own history entry.
+func (s *Server) handleWebSocketUpgrade(w http.ResponseWriter, r *http.Request, scheme, hostname string,
+ started time.Time, duration time.Duration, reqRaw []byte, reqExact, reqTrunc bool,
+ resp *http.Response, upstreamTee *teeConn) {
+
+ hijacker, ok := w.(http.Hijacker)
+ if !ok {
+ s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, reqTrunc, nil, false, false, resp.StatusCode,
+ "websocket: client connection doesn't support hijacking")
+ return
+ }
+ clientConn, brw, err := hijacker.Hijack()
+ if err != nil {
+ s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, reqTrunc, nil, false, false, resp.StatusCode, err.Error())
+ return
+ }
+ defer clientConn.Close()
+
+ // A WebSocket connection is expected to live far longer than one
+ // request/response - unlike the bounded upstreamTimeout the rest of
+ // forward() uses, there's no natural cutoff here.
+ clientConn.SetDeadline(time.Time{})
+ upstreamTee.SetDeadline(time.Time{})
+
+ // Everything upstreamTee has captured so far is exactly the raw 101
+ // response bytes, possibly with some already-arrived WebSocket frame
+ // bytes tacked on the end (bufio's own read-ahead inside
+ // roundTripH1) - split at the header/body boundary so the header
+ // portion can be relayed and recorded as this exchange's
+ // response_raw, and any leftover treated as the start of the frame
+ // stream rather than lost.
+ respRaw, _ := upstreamTee.Take()
+ headerEnd := len(respRaw)
+ if idx := bytes.Index(respRaw, []byte("\r\n\r\n")); idx >= 0 {
+ headerEnd = idx + 4
+ }
+ headerBytes, upstreamLeftover := respRaw[:headerEnd], respRaw[headerEnd:]
+
+ if _, err := clientConn.Write(headerBytes); err != nil {
+ s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, reqTrunc, headerBytes, true, false, resp.StatusCode, err.Error())
+ return
+ }
+
+ entryID := s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, reqTrunc, headerBytes, true, false, resp.StatusCode, "")
+
+ var clientLeftover []byte
+ if brw.Reader.Buffered() > 0 {
+ clientLeftover = make([]byte, brw.Reader.Buffered())
+ io.ReadFull(brw.Reader, clientLeftover)
+ }
+
+ clientReader := io.MultiReader(bytes.NewReader(clientLeftover), clientConn)
+ upstreamReader := io.MultiReader(bytes.NewReader(upstreamLeftover), upstreamTee)
+
+ // Each direction pumps independently.
+ done := make(chan struct{}, 2)
+ go func() {
+ pumpWS(clientReader, upstreamTee, func(opcode byte, payload []byte) {
+ s.recordWSMessage(entryID, "client_to_server", opcode, payload)
+ })
+ done <- struct{}{}
+ }()
+ go func() {
+ pumpWS(upstreamReader, clientConn, func(opcode byte, payload []byte) {
+ s.recordWSMessage(entryID, "server_to_client", opcode, payload)
+ })
+ done <- struct{}{}
+ }()
+
+ // Wait for the first direction to stop, then give the other one a
+ // bounded window to finish its own close sequence too - typically
+ // relaying the peer's own close-frame reply - rather than tearing
+ // the connection down the instant either side sees a close frame
+ // pass through. Without this, a client that closes gracefully would
+ // see its own close frame answered with an abrupt EOF instead of
+ // the origin's actual close reply. If the other direction doesn't
+ // finish in time (a slow or non-compliant peer), the deadlines below
+ // force it to unblock rather than leak the goroutine indefinitely.
+ <-done
+ deadline := time.Now().Add(5 * time.Second)
+ clientConn.SetDeadline(deadline)
+ upstreamTee.SetDeadline(deadline)
+ select {
+ case <-done:
+ case <-time.After(5 * time.Second):
+ }
+}
+
+func (s *Server) recordWSMessage(entryID int64, direction string, opcode byte, payload []byte) {
+ if s.store == nil || entryID == 0 {
+ return
+ }
+ if _, err := s.store.AddWSMessage(store.WSMessage{
+ EntryID: entryID,
+ StartedAt: time.Now(),
+ Direction: direction,
+ Opcode: int(opcode),
+ Payload: payload,
+ }); err != nil {
+ log.Printf("store websocket message: %v", err)
+ }
+}
+
// enabledRules fetches the current enabled match-and-replace rules for
// scope ("request" or "response") fresh from the store on every call -
// simple and always current, and cheap enough (a local, in-process
@@ -664,11 +789,15 @@ func (s *Server) enabledRules(scope string) ([]rules.Rule, error) {
return s.store.EnabledRules(scope)
}
-// record stores one history entry and notifies OnEntry.
+// record stores one history entry and notifies OnEntry, returning the
+// entry's assigned ID (0 if it wasn't stored at all - no store attached,
+// scope excluded it, or the insert itself failed) so a caller that needs
+// to attach more data to this specific entry afterward (see
+// handleWebSocketUpgrade's ws_messages rows) can do so.
func (s *Server) record(started time.Time, duration time.Duration, scheme, host string, r *http.Request,
- reqRaw []byte, reqExact, reqTruncated bool, respRaw []byte, respExact, respTruncated bool, status int, errMsg string) {
+ reqRaw []byte, reqExact, reqTruncated bool, respRaw []byte, respExact, respTruncated bool, status int, errMsg string) int64 {
if s.store == nil {
- return
+ return 0
}
// Scope only filters what gets recorded here - the request has
// already been forwarded and its response already written to the
@@ -680,7 +809,7 @@ func (s *Server) record(started time.Time, duration time.Duration, scheme, host
// the result regardless of scope, which exists to cut passive-
// capture noise, not to second-guess a deliberate action.
if scopeRules, err := s.store.ListScopeRules(); err == nil && !scope.InScope(scopeRules, host) {
- return
+ return 0
}
e := &store.Entry{
@@ -702,7 +831,7 @@ func (s *Server) record(started time.Time, duration time.Duration, scheme, host
id, err := s.store.Insert(e)
if err != nil {
log.Printf("store history entry: %v", err)
- return
+ return 0
}
if s.OnEntry != nil {
@@ -721,6 +850,7 @@ func (s *Server) record(started time.Time, duration time.Duration, scheme, host
Source: "proxy",
})
}
+ return id
}
// singleConnListener adapts one already-accepted net.Conn into a
@@ -821,3 +951,17 @@ func stripHopByHop(h http.Header) {
h.Del(k)
}
}
+
+// stripHopByHopKeepingUpgrade is stripHopByHop for a request that's
+// asking to upgrade the connection (see isWebSocketUpgradeRequest):
+// every other hop-by-hop header is still stripped, but Connection and
+// Upgrade are left alone since they're the upgrade request itself, not
+// leftover framing from the client's hop to mitmux.
+func stripHopByHopKeepingUpgrade(h http.Header) {
+ for _, k := range hopByHopHeaders {
+ if k == "Connection" || k == "Upgrade" {
+ continue
+ }
+ h.Del(k)
+ }
+}
diff --git a/internal/proxy/websocket.go b/internal/proxy/websocket.go
new file mode 100644
index 0000000..10a934c
--- /dev/null
+++ b/internal/proxy/websocket.go
@@ -0,0 +1,159 @@
+// WebSocket interception: after a client's Upgrade: websocket request
+// gets a matching 101 Switching Protocols response back from the origin
+// (see forward's WS branch), the connection stops being HTTP request/
+// response and becomes a long-lived, bidirectional, message-framed
+// stream (RFC 6455) instead. mitmux relays every frame byte-for-byte
+// unmodified in both directions - this is capture, not tampering - while
+// decoding each one's payload for display, recorded to the ws_messages
+// table tagged to the upgrade request's own history entry.
+//
+// Deliberately one row per frame, not per logical message: RFC 6455
+// lets a single message span several frames (opcode 0x0 continuation,
+// FIN unset until the last one), which mitmux does not reassemble.
+// Real-world WebSocket traffic - JSON events, chat messages, game state
+// - is overwhelmingly single-frame; reassembly would need buffering an
+// unbounded number of pending fragmented messages per connection for a
+// case that's rare in practice, which isn't a trade worth making here.
+package proxy
+
+import (
+ "encoding/binary"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+)
+
+// WebSocket opcodes (RFC 6455 section 5.2).
+const (
+ wsOpContinuation = 0x0
+ wsOpText = 0x1
+ wsOpBinary = 0x2
+ wsOpClose = 0x8
+ wsOpPing = 0x9
+ wsOpPong = 0xa
+)
+
+// isWebSocketUpgradeRequest reports whether r is asking to upgrade to a
+// WebSocket connection (Connection: Upgrade plus Upgrade: websocket).
+// forward() uses this to keep those two headers off stripHopByHop's list
+// for this one request - RFC 7230 correctly treats Connection/Upgrade as
+// hop-by-hop for a normal request, but stripping them here would delete
+// the very signal the origin needs to recognize the upgrade at all,
+// turning every WebSocket connection attempt into a silent 426.
+func isWebSocketUpgradeRequest(r *http.Request) bool {
+ return headerHasToken(r.Header, "Connection", "upgrade") &&
+ strings.EqualFold(r.Header.Get("Upgrade"), "websocket")
+}
+
+// isWebSocketUpgradeResponse reports whether resp is a successful
+// WebSocket upgrade (101 Switching Protocols, with Connection: Upgrade
+// and Upgrade: websocket) - checking the response rather than the
+// request it answers, since a 101 only ever comes back from an origin
+// that accepted the upgrade, and that's the one thing forward() actually
+// needs to know before handing the connection off.
+func isWebSocketUpgradeResponse(resp *http.Response) bool {
+ return resp.StatusCode == http.StatusSwitchingProtocols &&
+ headerHasToken(resp.Header, "Connection", "upgrade") &&
+ strings.EqualFold(resp.Header.Get("Upgrade"), "websocket")
+}
+
+func headerHasToken(h http.Header, name, token string) bool {
+ for _, v := range h.Values(name) {
+ for _, part := range strings.Split(v, ",") {
+ if strings.EqualFold(strings.TrimSpace(part), token) {
+ return true
+ }
+ }
+ }
+ return false
+}
+
+// relayWSFrame reads exactly one RFC 6455 frame from src, writes the
+// same raw bytes to dst unmodified, and returns the frame's opcode and
+// decoded (unmasked) payload for capture. Masking is direction-
+// dependent - client-to-server frames are always masked, server-to-
+// client frames never are - but relayed bytes are whatever was actually
+// read, so this works correctly regardless of which direction it's
+// called for.
+func relayWSFrame(src io.Reader, dst io.Writer) (opcode byte, payload []byte, err error) {
+ hdr := make([]byte, 2)
+ if _, err = io.ReadFull(src, hdr); err != nil {
+ return 0, nil, err
+ }
+ opcode = hdr[0] & 0x0f
+ masked := hdr[1]&0x80 != 0
+ length := uint64(hdr[1] & 0x7f)
+
+ raw := append([]byte(nil), hdr...)
+
+ switch length {
+ case 126:
+ ext := make([]byte, 2)
+ if _, err = io.ReadFull(src, ext); err != nil {
+ return 0, nil, err
+ }
+ raw = append(raw, ext...)
+ length = uint64(binary.BigEndian.Uint16(ext))
+ case 127:
+ ext := make([]byte, 8)
+ if _, err = io.ReadFull(src, ext); err != nil {
+ return 0, nil, err
+ }
+ raw = append(raw, ext...)
+ length = binary.BigEndian.Uint64(ext)
+ }
+
+ var maskKey [4]byte
+ if masked {
+ if _, err = io.ReadFull(src, maskKey[:]); err != nil {
+ return 0, nil, err
+ }
+ raw = append(raw, maskKey[:]...)
+ }
+
+ // Bounded the same way request/response body capture is (see
+ // maxCaptureBytes) - a length field mitmux doesn't control shouldn't
+ // be able to force an unbounded read/allocation. Relaying (not just
+ // capturing) is refused too: a frame this large is already well
+ // outside normal WebSocket usage, and guessing at a partial relay
+ // would corrupt the stream's framing for whichever side reads next.
+ if length > maxCaptureBytes {
+ return 0, nil, fmt.Errorf("websocket frame too large (%d bytes, over the %d limit)", length, maxCaptureBytes)
+ }
+
+ body := make([]byte, length)
+ if _, err = io.ReadFull(src, body); err != nil {
+ return 0, nil, err
+ }
+ raw = append(raw, body...)
+
+ if _, err = dst.Write(raw); err != nil {
+ return 0, nil, err
+ }
+
+ if !masked {
+ return opcode, body, nil
+ }
+ payload = make([]byte, length)
+ for i := range payload {
+ payload[i] = body[i] ^ maskKey[i%4]
+ }
+ return opcode, payload, nil
+}
+
+// pumpWS relays frames from src to dst until one fails to read/write or
+// a close frame (opcode 0x8) passes through, calling capture with each
+// frame's opcode and decoded payload as it goes.
+func pumpWS(src io.Reader, dst io.Writer, capture func(opcode byte, payload []byte)) {
+ for {
+ opcode, payload, err := relayWSFrame(src, dst)
+ if err != nil {
+ return
+ }
+ capture(opcode, payload)
+ if opcode == wsOpClose {
+ return
+ }
+ }
+}
diff --git a/internal/proxy/websocket_test.go b/internal/proxy/websocket_test.go
new file mode 100644
index 0000000..8256567
--- /dev/null
+++ b/internal/proxy/websocket_test.go
@@ -0,0 +1,159 @@
+package proxy
+
+import (
+ "bytes"
+ "testing"
+)
+
+// maskedFrame builds a masked (client-to-server style) RFC 6455 frame
+// for opcode carrying payload, mask applied per spec (XOR with a 4-byte
+// key repeated across the payload).
+func maskedFrame(t *testing.T, opcode byte, payload []byte, key [4]byte) []byte {
+ t.Helper()
+ var buf bytes.Buffer
+ buf.WriteByte(0x80 | opcode) // FIN=1
+ if len(payload) > 125 {
+ t.Fatalf("test helper only supports short payloads")
+ }
+ buf.WriteByte(0x80 | byte(len(payload))) // MASK=1
+ buf.Write(key[:])
+ masked := make([]byte, len(payload))
+ for i, b := range payload {
+ masked[i] = b ^ key[i%4]
+ }
+ buf.Write(masked)
+ return buf.Bytes()
+}
+
+// unmaskedFrame builds an unmasked (server-to-client style) frame.
+func unmaskedFrame(t *testing.T, opcode byte, payload []byte) []byte {
+ t.Helper()
+ var buf bytes.Buffer
+ buf.WriteByte(0x80 | opcode)
+ if len(payload) > 125 {
+ t.Fatalf("test helper only supports short payloads")
+ }
+ buf.WriteByte(byte(len(payload))) // MASK=0
+ buf.Write(payload)
+ return buf.Bytes()
+}
+
+func TestRelayWSFrameMaskedRoundTrips(t *testing.T) {
+ src := maskedFrame(t, wsOpText, []byte("hello"), [4]byte{1, 2, 3, 4})
+ var dst bytes.Buffer
+
+ opcode, payload, err := relayWSFrame(bytes.NewReader(src), &dst)
+ if err != nil {
+ t.Fatalf("unexpected error: %v", err)
+ }
+ if opcode != wsOpText {
+ t.Errorf("opcode = %d, want %d", opcode, wsOpText)
+ }
+ if string(payload) != "hello" {
+ t.Errorf("decoded payload = %q, want %q", payload, "hello")
+ }
+ // The relayed bytes must be byte-for-byte identical to the input -
+ // capture decodes for display, it never re-encodes what's on the
+ // wire.
+ if !bytes.Equal(dst.Bytes(), src) {
+ t.Errorf("relayed bytes = %x, want exactly %x (byte-exact passthrough)", dst.Bytes(), src)
+ }
+}
+
+func TestRelayWSFrameUnmaskedRoundTrips(t *testing.T) {
+ src := unmaskedFrame(t, wsOpBinary, []byte{0xde, 0xad, 0xbe, 0xef})
+ var dst bytes.Buffer
+
+ opcode, payload, err := relayWSFrame(bytes.NewReader(src), &dst)
+ if err != nil {
+ t.Fatalf("unexpected error: %v", err)
+ }
+ if opcode != wsOpBinary {
+ t.Errorf("opcode = %d, want %d", opcode, wsOpBinary)
+ }
+ if !bytes.Equal(payload, []byte{0xde, 0xad, 0xbe, 0xef}) {
+ t.Errorf("decoded payload = %x, want deadbeef", payload)
+ }
+ if !bytes.Equal(dst.Bytes(), src) {
+ t.Error("relayed bytes don't match input exactly")
+ }
+}
+
+func TestRelayWSFrameExtended16BitLength(t *testing.T) {
+ payload := bytes.Repeat([]byte("a"), 200) // over 125, forces the 126 extended-length form
+ var buf bytes.Buffer
+ buf.WriteByte(0x80 | wsOpBinary)
+ buf.WriteByte(126)
+ buf.Write([]byte{0x00, 0xc8}) // 200 in 16-bit big-endian
+ buf.Write(payload)
+
+ var dst bytes.Buffer
+ opcode, got, err := relayWSFrame(bytes.NewReader(buf.Bytes()), &dst)
+ if err != nil {
+ t.Fatalf("unexpected error: %v", err)
+ }
+ if opcode != wsOpBinary {
+ t.Errorf("opcode = %d, want %d", opcode, wsOpBinary)
+ }
+ if !bytes.Equal(got, payload) {
+ t.Errorf("payload length = %d, want %d", len(got), len(payload))
+ }
+ if !bytes.Equal(dst.Bytes(), buf.Bytes()) {
+ t.Error("relayed bytes don't match input exactly")
+ }
+}
+
+func TestRelayWSFrameOverCapRejected(t *testing.T) {
+ var buf bytes.Buffer
+ buf.WriteByte(0x80 | wsOpBinary)
+ buf.WriteByte(127)
+ lenBytes := make([]byte, 8)
+ big := uint64(maxCaptureBytes + 1)
+ for i := 7; i >= 0; i-- {
+ lenBytes[i] = byte(big)
+ big >>= 8
+ }
+ buf.Write(lenBytes)
+ // Deliberately no payload bytes written - relayWSFrame must reject
+ // based on the length field alone, before trying to read a body this
+ // large.
+
+ var dst bytes.Buffer
+ _, _, err := relayWSFrame(bytes.NewReader(buf.Bytes()), &dst)
+ if err == nil {
+ t.Fatal("expected an error for a frame over the capture limit")
+ }
+}
+
+func TestPumpWSCapturesUntilClose(t *testing.T) {
+ var src bytes.Buffer
+ src.Write(unmaskedFrame(t, wsOpText, []byte("one")))
+ src.Write(unmaskedFrame(t, wsOpText, []byte("two")))
+ src.Write(unmaskedFrame(t, wsOpClose, nil))
+ // A frame after close must never be reached.
+ src.Write(unmaskedFrame(t, wsOpText, []byte("unreachable")))
+
+ var dst bytes.Buffer
+ var captured [][]byte
+ pumpWS(&src, &dst, func(opcode byte, payload []byte) {
+ captured = append(captured, append([]byte(nil), payload...))
+ })
+
+ if len(captured) != 3 {
+ t.Fatalf("captured %d frames, want 3 (two messages + close)", len(captured))
+ }
+ if string(captured[0]) != "one" || string(captured[1]) != "two" {
+ t.Errorf("captured = %q, %q, want \"one\", \"two\"", captured[0], captured[1])
+ }
+}
+
+func TestPumpWSStopsOnReadError(t *testing.T) {
+ // Truncated frame: header claims a payload that never arrives.
+ src := bytes.NewReader([]byte{0x81, 0x05, 'h', 'i'})
+ var dst bytes.Buffer
+ called := false
+ pumpWS(src, &dst, func(byte, []byte) { called = true })
+ if called {
+ t.Error("capture should never run for a frame that fails to read fully")
+ }
+}
diff --git a/internal/store/store.go b/internal/store/store.go
index 0e57c3c..e0f27d1 100644
--- a/internal/store/store.go
+++ b/internal/store/store.go
@@ -74,6 +74,17 @@ CREATE TABLE IF NOT EXISTS client_certs (
cert_pem BLOB NOT NULL,
key_pem BLOB NOT NULL
);
+
+CREATE TABLE IF NOT EXISTS ws_messages (
+ id INTEGER PRIMARY KEY AUTOINCREMENT,
+ entry_id INTEGER NOT NULL,
+ started_at INTEGER NOT NULL,
+ direction TEXT NOT NULL,
+ opcode INTEGER NOT NULL,
+ payload BLOB NOT NULL
+);
+
+CREATE INDEX IF NOT EXISTS ws_messages_entry_id ON ws_messages(entry_id);
`
// Store is a handle to the history database. Safe for concurrent use.
@@ -683,6 +694,56 @@ func (s *Store) DeleteClientCert(id int64) error {
return nil
}
+// WSMessage is one captured WebSocket frame, tagged to the history entry
+// of the upgrade request/response that started its connection - see
+// internal/proxy/websocket.go for why it's one row per frame rather than
+// per reassembled logical message.
+type WSMessage struct {
+ ID int64
+ EntryID int64
+ StartedAt time.Time
+ Direction string // "client_to_server" or "server_to_client"
+ Opcode int // RFC 6455 opcode: 1 text, 2 binary, 8 close, 9 ping, 10 pong
+ Payload []byte
+}
+
+// AddWSMessage stores one captured frame and returns its assigned ID.
+func (s *Store) AddWSMessage(m WSMessage) (int64, error) {
+ res, err := s.db.Exec(
+ `INSERT INTO ws_messages (entry_id, started_at, direction, opcode, payload) VALUES (?, ?, ?, ?, ?)`,
+ m.EntryID, m.StartedAt.UnixMilli(), m.Direction, m.Opcode, m.Payload,
+ )
+ if err != nil {
+ return 0, fmt.Errorf("add ws message: %w", err)
+ }
+ return res.LastInsertId()
+}
+
+// ListWSMessages returns every frame captured for entryID's WebSocket
+// connection, in the order they were sent.
+func (s *Store) ListWSMessages(entryID int64) ([]WSMessage, error) {
+ rows, err := s.db.Query(
+ `SELECT id, entry_id, started_at, direction, opcode, payload FROM ws_messages WHERE entry_id = ? ORDER BY id`,
+ entryID,
+ )
+ if err != nil {
+ return nil, fmt.Errorf("list ws messages: %w", err)
+ }
+ defer rows.Close()
+
+ var out []WSMessage
+ for rows.Next() {
+ var m WSMessage
+ var startedAt int64
+ if err := rows.Scan(&m.ID, &m.EntryID, &startedAt, &m.Direction, &m.Opcode, &m.Payload); err != nil {
+ return nil, fmt.Errorf("scan ws message row: %w", err)
+ }
+ m.StartedAt = time.UnixMilli(startedAt)
+ out = append(out, m)
+ }
+ return out, rows.Err()
+}
+
// DeleteEntry removes a single history entry and its search index row.
func (s *Store) DeleteEntry(id int64) error {
tx, err := s.db.Begin()