srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal/proxy/proxy.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/proxy/proxy.go')
-rw-r--r--internal/proxy/proxy.go156
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)
+ }
+}