diff options
Diffstat (limited to 'internal/proxy/proxy.go')
| -rw-r--r-- | internal/proxy/proxy.go | 156 |
1 files changed, 150 insertions, 6 deletions
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) + } +} |