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 | |
| 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.
| -rw-r--r-- | PLAN.md | 98 | ||||
| -rw-r--r-- | README.md | 23 | ||||
| -rw-r--r-- | cmd/mitmux/main.go | 187 | ||||
| -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 |
9 files changed, 860 insertions, 11 deletions
@@ -50,7 +50,8 @@ hudsucker) - same problem, worth studying even though this build is Go. streams over one connection need per-stream request/response boundaries, not just per-connection ones). - CA install UX per OS (Linux/macOS/Windows trust stores) -- Whether WebSocket interception is v1 or a later addition +- WebSocket interception landed as a later addition, not v1 - see its + own section below - Step 6 match-and-replace now covers headers and bodies. Body rules materialize the body into memory (bounded by the same maxCaptureBytes cap as history capture) rather than streaming it @@ -578,3 +579,98 @@ and left vi-mode state intact, an empty rules table's right-click correctly no-opped (no crash, no menu), and adding a real rule then right-clicking it and clicking "disable" correctly toggled it off (confirmed via the rendered checkmark disappearing). + +## WebSocket interception + +The last of the "known limitations" list. A WebSocket connection stops +being one-shot HTTP request/response the instant a `101 Switching +Protocols` comes back - it becomes a long-lived, bidirectional, +message-framed (RFC 6455) stream instead, which forward()'s normal +write-response-then-record flow has no way to represent. Scoped to +HTTP/1.1 client legs only: an HTTP/2 client connection can't be +hijacked for raw post-response access the way HTTP/1.1 can, and RFC +8441 (WebSocket-over-HTTP/2 Extended CONNECT) is rare enough in +practice - browsers open a dedicated HTTP/1.1 connection for a +WebSocket even when the surrounding page is HTTP/2 - that excluding it +isn't a real-world gap. + +`internal/proxy/websocket.go` holds a minimal RFC 6455 frame codec +(`relayWSFrame`, `pumpWS`) - deliberately relay-first: it decodes a +frame's opcode and payload for capture while writing the exact same +raw bytes it read to the other side, unmodified. This is capture, not +tampering, matching the rest of the codebase's "raw bytes are the +source of truth" stance; there's no live WS message editing. One row +per frame, not per reassembled logical message - RFC 6455 lets a +message span several frames (opcode 0x0 continuation, FIN unset until +the last one), which isn't reassembled here. Real-world WebSocket +traffic is overwhelmingly single-frame; buffering an unbounded number +of pending fragmented messages per connection to handle the rare case +wasn't a trade worth making. + +`forward()` branches right after the response comes back: a 101 +matching `isWebSocketUpgradeResponse` and an HTTP/1.1-negotiated +connection hands off to `handleWebSocketUpgrade` instead of the normal +body-copy path - none of match-and-replace, body-rule capture, or +`stripHopByHop`'s usual header stripping make sense for a protocol +upgrade. handleWebSocketUpgrade hijacks the client connection, writes +the 101 response's raw bytes through unmodified, records the upgrade +request/response pair to history exactly like a normal exchange, then +relays frames bidirectionally - each captured into a new `ws_messages` +table (entry_id, direction, opcode, payload), reachable from the TUI's +detail view via `w`. + +Two real bugs surfaced only by actually driving a WebSocket connection +through a running daemon, not by reading the code - exactly the "run +it, don't read it" pattern that's caught every prior bug like this in +this project: + +- `stripHopByHop` was already stripping `Connection` and `Upgrade` + from the outgoing request - correct for an ordinary request per RFC + 7230 (they're hop-by-hop headers), catastrophic for one asking to + upgrade, since those two headers *are* the upgrade request. Every + WebSocket connection attempt silently became a 426 before this was + caught: the origin never saw the upgrade at all. + `isWebSocketUpgradeRequest` + `stripHopByHopKeepingUpgrade` fix it - + every other hop-by-hop header still stripped, just not these two, + and only for a request that's actually asking to upgrade. +- The relay originally ended the whole handler the instant *either* + direction saw a close frame pass through. In practice this meant: a + client sends a close frame, that direction's pump relays it upstream + and returns, the handler tears the connection down immediately - + before the origin's own close-frame *reply* (which it sends after + receiving the client's) can be read and relayed back. The test + client's close handshake failed with an abrupt EOF instead of a + close frame. Fixed by waiting for the first direction to stop, then + giving the other one a bounded 5-second window to finish its own + close sequence before forcing both connections closed via a + deadline - long enough for a well-behaved peer's reply to get + through, bounded so a peer that never replies can't leak the + goroutine indefinitely. + +Verified live end to end against real servers, not mocks, on both +paths a real client actually uses: + +- `ws://` (plain HTTP forward-proxying): a Python `websockets`-based + echo server, and a hand-built raw-socket test client speaking RFC + 6455 directly (masked client frames, unmasked server frames, the + 16-bit extended-length form for a 500-byte message, a binary + message, and a full close handshake) sent through mitmux via an + absolute-form `GET http://host/ HTTP/1.1` - exactly how a real + proxy-configured WebSocket client negotiates one. Every message + round-tripped correctly and the close handshake completed with both + directions' close frames present. +- `wss://` (CONNECT-tunneled, TLS-intercepted): same echo server + behind TLS (a throwaway self-signed cert, trusted for the test via a + process-scoped `SSL_CERT_FILE`, never touching the real system trust + store), reached through mitmux's own CONNECT handling and MITM leaf + certificate. Confirmed the handshake, a message round-trip, and the + close handshake all work identically over the hijacked `*tls.Conn` + the CONNECT path hands back - this is the path real browsers + actually use for `wss://`, so this was worth checking separately + from the plain-HTTP path rather than assuming it'd behave the same. + +In both cases, `go run ./cmd/livetest` (a throwaway program, deleted +after use - never part of the build) confirmed the captured messages +in `ws_messages` via `ipc.Client.ListWSMessages`, matching what the +test client actually sent and received, correctly attributed to +direction and opcode. @@ -23,6 +23,9 @@ list of what's deliberately not implemented (and why), see independently on the client and upstream legs - a client that only speaks HTTP/1.1 and an origin that prefers HTTP/2 both work correctly in the same request. +- **WebSocket**: `ws://`/`wss://` connections are relayed byte-for-byte + unmodified with every frame captured for display, not just the + upgrade handshake - see WebSocket below. - **History**: every request/response captured to SQLite. Raw wire bytes are preserved byte-for-byte on HTTP/1.1 legs (what request smuggling and parser-differential analysis actually needs); HTTP/2 @@ -258,7 +261,22 @@ response (display-only - never touches the stored or resent bytes), `c` mark/compare (same as the history list), `r`/`i` jump straight to Repeater/Intruder seeded from this entry, `e` exports the entry (request and response, raw bytes, plain text - type a path and press -enter), `esc` back. +enter), `w` views captured WebSocket messages if this entry's +connection was upgraded (see WebSocket below), `esc` back. + +### WebSocket + +A `ws://` or `wss://` request that gets a matching `101 Switching +Protocols` back stops being one-shot request/response - mitmux relays +every frame byte-for-byte unmodified in both directions (this is +capture, not tampering) while decoding each one's payload for display. +Press `w` from an upgraded entry's detail view to see them: direction, +opcode (text/binary/close/ping/pong), size, and a preview; `enter` on +a row shows that frame's full decoded payload. One row per frame, not +per reassembled logical message - a message fragmented across several +frames (rare in real-world WebSocket traffic: JSON events, chat +messages, game state are almost always single-frame) shows up as +several rows rather than being stitched back together. ### Comparer @@ -571,7 +589,8 @@ reasoning behind each: 1000 requests per attack across all four modes - `mitmuxd -install-ca` prints per-OS trust-store install steps; it never runs them for you (see Quick start above for why) -- No WebSocket interception +- WebSocket messages are captured one row per frame, not reassembled + from fragments (rare in real-world traffic) - see WebSocket below - No active or passive vulnerability scanning, no plugin system - this is a manual-testing tool, not a scanner diff --git a/cmd/mitmux/main.go b/cmd/mitmux/main.go index 6e07069..afb345b 100644 --- a/cmd/mitmux/main.go +++ b/cmd/mitmux/main.go @@ -116,6 +116,7 @@ const ( viewDecoder viewScope viewClientCerts + viewWebSocket viewHelp ) @@ -236,6 +237,15 @@ type model struct { clientCertKeyPath textinput.Model clientCertIsRegex bool + // WebSocket messages captured for one history entry's upgraded + // connection (see proxy.go's handleWebSocketUpgrade) - reached from + // the detail view, not its own top-level list. + wsTable table.Model + wsRows []store.WSMessage + wsEntryID int64 + wsShowingDetail bool + wsDetailViewport viewport.Model + intruderScheme string intruderHost string intruderTemplate viTextarea @@ -389,6 +399,16 @@ func newModel(client *ipc.Client, subCh <-chan store.Summary, socketPath string) ccKeyPathIn := textinput.New() ccKeyPathIn.Placeholder = "path to PEM private key file" + wsCols := []table.Column{ + {Title: "#", Width: 4}, + {Title: "Dir", Width: 3}, + {Title: "Op", Width: 6}, + {Title: "Size", Width: 8}, + {Title: "Preview", Width: 50}, + } + wsTbl := table.New(table.WithColumns(wsCols), table.WithFocused(true)) + wsTbl.SetStyles(st) + itmpl := newViTextarea() itmpl.ta.Placeholder = "raw request bytes - wrap positions to fuzz in § markers, e.g. /users/§123§" itmpl.ta.ShowLineNumbers = false @@ -441,6 +461,7 @@ func newModel(client *ipc.Client, subCh <-chan store.Summary, socketPath string) clientCertPattern: ccPatternIn, clientCertCertPath: ccCertPathIn, clientCertKeyPath: ccKeyPathIn, + wsTable: wsTbl, ruleName: nameIn, ruleMatch: matchIn, ruleReplace: replaceIn, @@ -638,6 +659,19 @@ func (m *model) loadDetail(id int64, dest string) tea.Cmd { } } +type wsMessagesLoadedMsg struct { + entryID int64 + messages []store.WSMessage + err error +} + +func (m *model) loadWSMessages(entryID int64) tea.Cmd { + return func() tea.Msg { + msgs, err := m.client.ListWSMessages(entryID) + return wsMessagesLoadedMsg{entryID: entryID, messages: msgs, err: err} + } +} + // markOrCompare implements 'c': the first press on an entry marks it as // the comparison base (no fetch yet - cheap, no round trip until there's // actually something to compare). A second press on a *different* entry @@ -1154,6 +1188,9 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.searchInput.Width = msg.Width - 2 m.viewport = viewport.New(msg.Width, h-5) m.compareViewport = viewport.New(msg.Width, h-5) + m.wsTable.SetWidth(msg.Width) + m.wsTable.SetHeight(h - 5) + m.wsDetailViewport = viewport.New(msg.Width, h-5) decInHeight := (h - 8) / 2 m.decoderInput.SetWidth(msg.Width) @@ -1298,6 +1335,23 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.viewport.GotoTop() return m, nil + case wsMessagesLoadedMsg: + if msg.err != nil { + m.statusMsg = "websocket messages error: " + msg.err.Error() + return m, nil + } + if len(msg.messages) == 0 { + m.statusMsg = "no websocket messages captured for this entry" + return m, nil + } + m.wsRows = msg.messages + m.wsEntryID = msg.entryID + m.wsShowingDetail = false + setTableRows(&m.wsTable, wsRowsFor(m.wsRows)) + m.mode = viewWebSocket + m.statusMsg = "" + return m, nil + case compareLoadedMsg: if msg.err != nil { m.statusMsg = "compare error: " + msg.err.Error() @@ -1695,6 +1749,12 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if m.detail != nil { return m, m.markOrCompare(m.detail.ID) } + case "w": + if m.detail != nil { + m.statusMsg = "loading websocket messages..." + return m, m.loadWSMessages(m.detail.ID) + } + return m, nil case "tab": if m.activeTab == tabRequest { m.activeTab = tabResponse @@ -1978,6 +2038,41 @@ func (m *model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.clientCertTable, cmd = m.clientCertTable.Update(msg) return m, cmd + case viewWebSocket: + if m.wsShowingDetail { + switch msg.String() { + case "q", "esc": + m.wsShowingDetail = false + return m, nil + case "ctrl+c": + return m, tea.Quit + } + var cmd tea.Cmd + m.wsDetailViewport, cmd = m.wsDetailViewport.Update(msg) + return m, cmd + } + switch msg.String() { + case "q", "esc": + m.mode = viewDetail + return m, nil + case "ctrl+c": + return m, tea.Quit + case "?": + m.prevMode = viewWebSocket + m.mode = viewHelp + return m, nil + case "enter": + if row := m.wsTable.Cursor(); row >= 0 && row < len(m.wsRows) { + m.wsShowingDetail = true + m.wsDetailViewport.SetContent(wsMessageDetail(m.wsRows[row])) + m.wsDetailViewport.GotoTop() + } + return m, nil + } + var cmd tea.Cmd + m.wsTable, cmd = m.wsTable.Update(msg) + return m, cmd + case viewIntruder: // Editing a grep pattern is a modal overlay on top of the // normal template/payloads/results panes, same pattern as @@ -2235,6 +2330,12 @@ func (m *model) View() string { } else { body = m.clientCertView() } + case viewWebSocket: + if m.wsShowingDetail { + body = m.wsDetailView() + } else { + body = m.wsView() + } case viewIntruder: body = m.intruderView() case viewCompare: @@ -2264,7 +2365,7 @@ func (m *model) statusBar() string { view := map[viewMode]string{ viewList: "history", viewDetail: "detail", viewRepeater: "repeater", viewRules: "rules", viewIntruder: "intruder", viewCompare: "comparer", viewDecoder: "decoder", - viewScope: "scope", viewClientCerts: "client certs", + viewScope: "scope", viewClientCerts: "client certs", viewWebSocket: "websocket", }[m.mode] return statusBarStyle.Render(fmt.Sprintf(" mitmux · proxy %s%s · %s ", proxy, count, view)) } @@ -2323,8 +2424,16 @@ func (m *model) helpView() string { "c mark/compare (same as history list)", "r / i open in Repeater / Intruder", "e export this entry - .txt (raw request+response) or .sh/.curl (curl command)", + "w view captured WebSocket messages (only for an upgraded connection)", "esc / q back to history", ) + section("WebSocket messages", + "One row per captured frame (not per reassembled message - see", + "PLAN.md) for the entry's upgraded connection.", + "↑/↓ or j/k navigate (also g/G, ctrl+u/d)", + "enter view this frame's full decoded payload", + "esc / q back (from payload view: back to the message list)", + ) section("Comparer", "tab switch request/response diff", "↑/↓ or j/k scroll (also g/G, ctrl+u/d - same as history list)", @@ -2477,7 +2586,7 @@ func (m *model) detailView() string { b.WriteString(statusStyle.Render(sanitizeLine(m.statusMsg))) b.WriteString("\n") } - b.WriteString(helpStyle.Render("tab switch · p pretty-print · c compare · r repeater · i intruder · e export · esc back · ? help · q quit")) + b.WriteString(helpStyle.Render("tab switch · p pretty-print · c compare · r repeater · i intruder · e export · w websocket · esc back · ? help · q quit")) return b.String() } @@ -2733,6 +2842,80 @@ func clientCertRowsFor(cs []clientcert.Cert) []table.Row { return rows } +func (m *model) wsView() string { + var b strings.Builder + title := fmt.Sprintf(" websocket messages (%d) - entry #%d ", len(m.wsRows), m.wsEntryID) + b.WriteString(titleStyle.Render(title)) + b.WriteString("\n") + b.WriteString(m.wsTable.View()) + b.WriteString("\n") + if m.statusMsg != "" { + b.WriteString(statusStyle.Render(sanitizeLine(m.statusMsg))) + b.WriteString("\n") + } + b.WriteString(helpStyle.Render("enter view payload · esc back · q quit")) + return b.String() +} + +func (m *model) wsDetailView() string { + var b strings.Builder + b.WriteString(titleStyle.Render(" websocket message ")) + b.WriteString("\n") + b.WriteString(m.wsDetailViewport.View()) + b.WriteString("\n") + b.WriteString(helpStyle.Render("↑/↓ scroll · esc back · q quit")) + return b.String() +} + +func wsOpcodeName(opcode int) string { + switch opcode { + case 0x1: + return "text" + case 0x2: + return "binary" + case 0x8: + return "close" + case 0x9: + return "ping" + case 0xa: + return "pong" + default: + return fmt.Sprintf("0x%x", opcode) + } +} + +func wsRowsFor(msgs []store.WSMessage) []table.Row { + rows := make([]table.Row, len(msgs)) + for i, m := range msgs { + dir := "->" + if m.Direction == "server_to_client" { + dir = "<-" + } + rows[i] = table.Row{ + fmt.Sprintf("%d", i+1), + dir, + wsOpcodeName(m.Opcode), + humanBytes(len(m.Payload)), + sanitizeLine(string(m.Payload)), + } + } + return rows +} + +// wsMessageDetail is the full, sanitized content shown when viewing one +// captured frame's payload - sanitizeBlock because, same as request/ +// response bodies elsewhere, this is attacker- or origin-controlled data +// reaching the operator's real terminal, not just the display width +// truncation the table preview gets away with. +func wsMessageDetail(m store.WSMessage) string { + dir := "client -> server" + if m.Direction == "server_to_client" { + dir = "server -> client" + } + return fmt.Sprintf("%s · opcode: %s · %s\n\n%s", + dir, wsOpcodeName(m.Opcode), humanBytes(len(m.Payload)), sanitizeBlock(string(m.Payload))) +} + // nextAttackMode cycles Sniper -> BatteringRam -> Pitchfork -> ClusterBomb // -> Sniper. func nextAttackMode(mode proxy.AttackMode) proxy.AttackMode { 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() |