diff options
| author | srdusr <[email protected]> | 2026-06-30 14:52:00 +0200 |
|---|---|---|
| committer | srdusr <[email protected]> | 2026-06-30 14:52:00 +0200 |
| commit | 384573a2dc5e3b8e2a7bdfe2ce949f2c52ba2c52 (patch) | |
| tree | 13bfe97b1f5bcf90e3fad33d90b522b29f3e439d /internal | |
| parent | 2ade8c807584bff0b60d6b6f278dbde29b13a5ff (diff) | |
| download | mitmux-384573a2dc5e3b8e2a7bdfe2ce949f2c52ba2c52.tar.gz mitmux-384573a2dc5e3b8e2a7bdfe2ce949f2c52ba2c52.zip | |
WebSocket interception
The last "known limitation": a ws://wss:// connection stops being
one-shot request/response the instant its 101 Switching Protocols
lands, and forward()'s normal write-response-then-record flow has no
way to represent that. Scoped to HTTP/1.1 client legs (HTTP/2 can't be
hijacked for raw post-response access the way HTTP/1.1 can, and
browsers open a dedicated HTTP/1.1 connection for WebSocket regardless
of the surrounding page's protocol, so this isn't a real-world gap).
internal/proxy/websocket.go decodes each RFC 6455 frame's opcode and
payload for capture while relaying the exact same raw bytes it read
unmodified - this is capture, not tampering, matching the rest of the
codebase's raw-bytes-as-source-of-truth stance. One row per frame, not
per reassembled message (fragmentation is rare in real-world
WebSocket traffic; not worth buffering an unbounded number of pending
fragments to handle it). forward() branches on a matching 101 into
handleWebSocketUpgrade, which hijacks the client connection, relays
the handshake response raw, records the upgrade request/response to
history normally, then relays frames bidirectionally into a new
ws_messages table - reachable from the TUI's detail view via `w`.
Found and fixed two real bugs by actually driving a WebSocket
connection through a running daemon, not by reading the code:
stripHopByHop was deleting Connection/Upgrade from every outgoing
request (correct for an ordinary request per RFC 7230, catastrophic
for one asking to upgrade - every WebSocket attempt silently became a
426); and the relay tore the whole connection down the instant either
side saw a close frame, before the peer's own close-frame reply could
be relayed back, producing an abrupt EOF instead of a clean close.
Verified live end to end on both paths a real client uses: ws://
(plain HTTP forward-proxying) against a Python websockets echo
server, and wss:// (CONNECT-tunneled, TLS-intercepted) against the
same server behind TLS - text, binary, and extended-length frames,
plus a full close handshake with both directions' close frames
present, confirmed via the actual bytes captured in ws_messages.
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/ipc/ipc.go | 20 | ||||
| -rw-r--r-- | internal/ipc/server.go | 8 | ||||
| -rw-r--r-- | internal/proxy/proxy.go | 156 | ||||
| -rw-r--r-- | internal/proxy/websocket.go | 159 | ||||
| -rw-r--r-- | internal/proxy/websocket_test.go | 159 | ||||
| -rw-r--r-- | internal/store/store.go | 61 |
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() |