diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/ipc/ipc.go | 126 | ||||
| -rw-r--r-- | internal/ipc/server.go | 136 | ||||
| -rw-r--r-- | internal/proxy/capture.go | 83 | ||||
| -rw-r--r-- | internal/proxy/proxy.go | 269 | ||||
| -rw-r--r-- | internal/proxy/tee.go | 92 | ||||
| -rw-r--r-- | internal/store/store.go | 186 |
6 files changed, 821 insertions, 71 deletions
diff --git a/internal/ipc/ipc.go b/internal/ipc/ipc.go new file mode 100644 index 0000000..c9d92b4 --- /dev/null +++ b/internal/ipc/ipc.go @@ -0,0 +1,126 @@ +// Package ipc is the protocol between mitmuxd (which owns the proxy and +// the history database) and a client such as the TUI, spoken as +// newline-agnostic JSON messages over a Unix domain socket. This keeps +// the proxy engine running independently of any UI attached to it. +package ipc + +import ( + "encoding/json" + "errors" + "fmt" + "net" + + "mitmux/internal/store" +) + +// Request is sent by a client to the daemon. +type Request struct { + Type string `json:"type"` // "list", "get", or "subscribe" + Limit int `json:"limit,omitempty"` + BeforeID int64 `json:"before_id,omitempty"` + ID int64 `json:"id,omitempty"` +} + +// Response is sent by the daemon to a client. +type Response struct { + Type string `json:"type"` // "list", "get", "new", or "error" + Entries []store.Summary `json:"entries,omitempty"` // for "list" + Detail *EntryDetail `json:"detail,omitempty"` // for "get" + New *store.Summary `json:"new,omitempty"` // for "new" (subscribe push) + Error string `json:"error,omitempty"` +} + +// EntryDetail is a full history entry, raw bytes included. +type EntryDetail struct { + store.Summary + RequestRaw []byte `json:"request_raw"` + ResponseRaw []byte `json:"response_raw"` + RequestExact bool `json:"request_exact"` + ResponseExact bool `json:"response_exact"` +} + +// Client talks to a mitmuxd instance for request/response queries +// (list, get). Use Subscribe separately for the live-update stream. +type Client struct { + conn net.Conn + dec *json.Decoder + enc *json.Encoder +} + +// Dial connects to the daemon's control socket at path. +func Dial(path string) (*Client, error) { + conn, err := net.Dial("unix", path) + if err != nil { + return nil, fmt.Errorf("dial %s: %w", path, err) + } + return &Client{conn: conn, dec: json.NewDecoder(conn), enc: json.NewEncoder(conn)}, nil +} + +// Close closes the connection to the daemon. +func (c *Client) Close() error { + return c.conn.Close() +} + +// List returns up to limit history summaries older than beforeID (0 for +// the most recent), newest first. +func (c *Client) List(limit int, beforeID int64) ([]store.Summary, error) { + if err := c.enc.Encode(Request{Type: "list", Limit: limit, BeforeID: beforeID}); 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.Entries, nil +} + +// Get returns the full entry (raw bytes included) for id. +func (c *Client) Get(id int64) (*EntryDetail, error) { + if err := c.enc.Encode(Request{Type: "get", ID: id}); 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.Detail, nil +} + +// Subscribe opens a dedicated connection that streams newly captured +// history entries as they happen. The returned channel is closed when +// the connection ends; call the returned close func to stop early. +func Subscribe(path string) (<-chan store.Summary, func() error, error) { + conn, err := net.Dial("unix", path) + if err != nil { + return nil, nil, fmt.Errorf("dial %s: %w", path, err) + } + if err := json.NewEncoder(conn).Encode(Request{Type: "subscribe"}); err != nil { + conn.Close() + return nil, nil, err + } + + ch := make(chan store.Summary, 64) + go func() { + defer close(ch) + dec := json.NewDecoder(conn) + for { + var resp Response + if err := dec.Decode(&resp); err != nil { + return + } + if resp.Type == "new" && resp.New != nil { + select { + case ch <- *resp.New: + default: + } + } + } + }() + return ch, conn.Close, nil +} diff --git a/internal/ipc/server.go b/internal/ipc/server.go new file mode 100644 index 0000000..11ba033 --- /dev/null +++ b/internal/ipc/server.go @@ -0,0 +1,136 @@ +package ipc + +import ( + "encoding/json" + "log" + "net" + "sync" + + "mitmux/internal/store" +) + +// Hub fans out newly captured history entries to subscribed clients. +type Hub struct { + mu sync.Mutex + subs map[chan store.Summary]struct{} +} + +// NewHub creates an empty Hub. +func NewHub() *Hub { + return &Hub{subs: make(map[chan store.Summary]struct{})} +} + +// Broadcast notifies all current subscribers of e. Slow subscribers +// drop entries rather than blocking the proxy. +func (h *Hub) Broadcast(e store.Summary) { + h.mu.Lock() + defer h.mu.Unlock() + for ch := range h.subs { + select { + case ch <- e: + default: + } + } +} + +func (h *Hub) subscribe() chan store.Summary { + ch := make(chan store.Summary, 64) + h.mu.Lock() + h.subs[ch] = struct{}{} + h.mu.Unlock() + return ch +} + +func (h *Hub) unsubscribe(ch chan store.Summary) { + h.mu.Lock() + delete(h.subs, ch) + h.mu.Unlock() + close(ch) +} + +// Server serves the daemon side of the mitmux control protocol. +type Server struct { + db *store.Store + hub *Hub +} + +// NewServer creates a control-protocol Server backed by db, broadcasting +// through hub. +func NewServer(db *store.Store, hub *Hub) *Server { + return &Server{db: db, hub: hub} +} + +// Serve accepts connections on ln until it returns an error (e.g. the +// listener is closed). +func (s *Server) Serve(ln net.Listener) error { + for { + conn, err := ln.Accept() + if err != nil { + return err + } + go s.handleConn(conn) + } +} + +func (s *Server) handleConn(conn net.Conn) { + defer conn.Close() + dec := json.NewDecoder(conn) + enc := json.NewEncoder(conn) + + for { + var req Request + if err := dec.Decode(&req); err != nil { + return + } + + switch req.Type { + case "list": + entries, err := s.db.List(req.Limit, req.BeforeID) + if err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "list", Entries: entries}) + + case "get": + e, err := s.db.Get(req.ID) + if err != nil { + enc.Encode(Response{Type: "error", Error: err.Error()}) + continue + } + enc.Encode(Response{Type: "get", Detail: &EntryDetail{ + Summary: store.Summary{ + ID: e.ID, StartedAt: e.StartedAt, Duration: e.Duration, + Method: e.Method, Scheme: e.Scheme, Host: e.Host, Path: e.Path, + StatusCode: e.StatusCode, ReqSize: len(e.RequestRaw), RespSize: len(e.ResponseRaw), + Error: e.Error, + }, + RequestRaw: e.RequestRaw, ResponseRaw: e.ResponseRaw, + RequestExact: e.RequestExact, ResponseExact: e.ResponseExact, + }}) + + case "subscribe": + sub := s.hub.subscribe() + defer s.hub.unsubscribe(sub) + for e := range sub { + e := e + if err := enc.Encode(Response{Type: "new", New: &e}); err != nil { + return + } + } + return + + default: + enc.Encode(Response{Type: "error", Error: "unknown request type: " + req.Type}) + } + } +} + +// LogAndBroadcast is a convenience OnEntry callback: logs the entry and +// broadcasts it through hub. +func LogAndBroadcast(hub *Hub) func(store.Summary) { + return func(sum store.Summary) { + log.Printf("%s %s%s -> %d (%s)", sum.Method, sum.Host, sum.Path, sum.StatusCode, sum.Duration) + hub.Broadcast(sum) + } +} diff --git a/internal/proxy/capture.go b/internal/proxy/capture.go new file mode 100644 index 0000000..ccc95e9 --- /dev/null +++ b/internal/proxy/capture.go @@ -0,0 +1,83 @@ +package proxy + +import ( + "bytes" + "io" + "net/http" +) + +// cappedTee wraps an io.Reader, copying up to maxCaptureBytes of what +// passes through into an internal buffer while still passing everything +// through unmodified and unbounded to the real reader. Used to capture a +// bounded sample of a body for reconstruction when exact wire capture +// isn't available (the HTTP/2 leg - see below). +type cappedTee struct { + r io.Reader + buf bytes.Buffer +} + +func newCappedTee(r io.Reader) *cappedTee { + return &cappedTee{r: r} +} + +func (c *cappedTee) Read(p []byte) (int, error) { + n, err := c.r.Read(p) + if n > 0 { + if room := maxCaptureBytes - c.buf.Len(); room > 0 { + end := n + if end > room { + end = room + } + c.buf.Write(p[:end]) + } + } + return n, err +} + +// captureRequest returns the raw bytes of r for storage. When tee is +// non-nil (an HTTP/1.1 client connection), the bytes are exactly what +// was read off the wire. Otherwise (HTTP/2, which has no single "raw +// bytes" representation - it's multiplexed, HPACK-compressed framing) +// it's a reconstruction from the parsed request, exact=false. +func captureRequest(r *http.Request, tee *teeConn, bodyCap *cappedTee) (raw []byte, exact bool) { + if tee != nil { + return tee.Take(), true + } + + dump := r.Clone(r.Context()) + if bodyCap != nil { + dump.Body = io.NopCloser(bytes.NewReader(bodyCap.buf.Bytes())) + dump.ContentLength = int64(bodyCap.buf.Len()) + } else { + dump.Body = http.NoBody + dump.ContentLength = 0 + } + var buf bytes.Buffer + if err := dump.Write(&buf); err != nil { + return nil, false + } + return buf.Bytes(), false +} + +// captureResponse mirrors captureRequest for the upstream leg: exact +// wire bytes when tee is non-nil (upstream negotiated HTTP/1.1), +// otherwise a reconstruction. +func captureResponse(resp *http.Response, tee *teeConn, bodyCap *cappedTee) (raw []byte, exact bool) { + if tee != nil { + return tee.Take(), true + } + + dump := *resp + if bodyCap != nil { + dump.Body = io.NopCloser(bytes.NewReader(bodyCap.buf.Bytes())) + dump.ContentLength = int64(bodyCap.buf.Len()) + } else { + dump.Body = http.NoBody + dump.ContentLength = 0 + } + var buf bytes.Buffer + if err := dump.Write(&buf); err != nil { + return nil, false + } + return buf.Bytes(), false +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 4fe9a4f..0a20382 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -3,11 +3,25 @@ // are intercepted: mitmux terminates TLS with the client using a leaf // certificate signed by its own CA, and separately terminates TLS with // the real server, forwarding requests between the two. ALPN is -// negotiated with the real server first and mirrored to the client so -// HTTP/2 connections stay HTTP/2 end to end rather than being downgraded. +// negotiated independently on each side (see handleConnect) so HTTP/2 +// stays HTTP/2 end to end without one side being forced to match the +// other. Every request/response pair is captured to the history store - +// exactly, byte for byte, on HTTP/1.1 legs; reconstructed on HTTP/2 legs, +// which have no meaningful "raw bytes" of their own (see capture.go). +// +// Upstream requests are round-tripped manually (write the request, +// read the response off the same connection) rather than through +// http.Transport: Transport's automatic HTTP/2 dispatch keys off a +// literal *tls.Conn type assertion on the connection it dials, which a +// capturing wrapper around that connection defeats - the request would +// silently be parsed as HTTP/1.1 over what is actually HTTP/2 framing. +// Handling both protocols explicitly here, per request, avoids that and +// also removes any ambiguity about which connection served which +// request, since each request gets its own connection either way. package proxy import ( + "bufio" "context" "crypto/tls" "errors" @@ -21,6 +35,7 @@ import ( "golang.org/x/net/http2" "mitmux/internal/ca" + "mitmux/internal/store" ) // hopByHopHeaders are stripped before forwarding a request or response, @@ -42,40 +57,36 @@ var hopByHopHeaders = []string{ type Server struct { Addr string - ca *ca.CA - transport *http.Transport - server *http.Server + // OnEntry, if set, is called after each request/response pair is + // stored, so a daemon can broadcast it to live TUI subscribers. + OnEntry func(store.Summary) + + ca *ca.CA + store *store.Store + server *http.Server } // New creates a proxy Server bound to addr (e.g. "127.0.0.1:8080"), -// signing intercepted TLS connections with root. -func New(addr string, root *ca.CA) *Server { - s := &Server{ - Addr: addr, - ca: root, - transport: &http.Transport{ - Proxy: nil, - DialContext: (&net.Dialer{ - Timeout: 10 * time.Second, - }).DialContext, - ForceAttemptHTTP2: false, - MaxIdleConns: 100, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 10 * time.Second, - ExpectContinueTimeout: 1 * time.Second, - }, - } +// signing intercepted TLS connections with root and recording history to +// db. +func New(addr string, root *ca.CA, db *store.Store) *Server { + s := &Server{Addr: addr, ca: root, store: db} s.server = &http.Server{ - Addr: addr, - Handler: http.HandlerFunc(s.handle), + Addr: addr, + Handler: http.HandlerFunc(s.handle), + ConnContext: withClientTee, } return s } // ListenAndServe starts the proxy and blocks until it stops. func (s *Server) ListenAndServe() error { + ln, err := net.Listen("tcp", s.Addr) + if err != nil { + return err + } log.Printf("proxy listening on %s", s.Addr) - return s.server.ListenAndServe() + return s.server.Serve(&teeListener{Listener: ln}) } // Shutdown gracefully stops the proxy. @@ -91,15 +102,18 @@ func (s *Server) handle(w http.ResponseWriter, r *http.Request) { s.handleHTTP(w, r) } +// dialer resolves a fresh upstream connection for one request, along +// with the ALPN protocol negotiated for it ("http/1.1", "h2", or "" if +// not applicable/negotiated). +type dialer func(ctx context.Context) (conn net.Conn, negotiated string, err error) + // handleConnect intercepts a CONNECT request: it terminates TLS with the // client using a leaf certificate signed by mitmux's CA, then forwards // each request upstream over its own independently negotiated TLS // connection. Client-side and upstream-side ALPN are negotiated // separately (each offering both HTTP/2 and HTTP/1.1) rather than one // being forced to match the other, so e.g. an HTTP/1.1-only client -// reaching an HTTP/2-only-preferring server doesn't fail to connect - -// http.Transport (via http2.ConfigureTransport) bridges the two sides -// independently per request. +// reaching an HTTP/2-preferring server doesn't fail to connect. func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) { hostPort := r.Host hostname, _, err := net.SplitHostPort(hostPort) @@ -141,18 +155,11 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) { return } - tr := &http.Transport{ - DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - return dialUpstreamTLS(ctx, hostPort, hostname) - }, - } - if err := http2.ConfigureTransport(tr); err != nil { - log.Printf("configure h2 transport for %s: %v", hostname, err) + dial := func(ctx context.Context) (net.Conn, string, error) { + return dialUpstreamTLS(ctx, hostPort, hostname) } - defer tr.CloseIdleConnections() - handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - s.forward(tr, "https", hostname, w, r) + s.forward(dial, "https", hostname, w, r) }) if clientTLS.ConnectionState().NegotiatedProtocol == http2.NextProtoTLS { @@ -160,7 +167,8 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) { return } - err = http.Serve(newSingleConnListener(clientTLS), handler) + h1 := &http.Server{Handler: handler, ConnContext: withClientTee} + err = h1.Serve(newSingleConnListener(clientTLS)) if err != nil && !errors.Is(err, io.EOF) { log.Printf("h1 serve for %s: %v", hostname, err) } @@ -168,11 +176,11 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) { // dialUpstreamTLS connects to the real server, offering both HTTP/2 and // HTTP/1.1 over ALPN and letting the server pick. -func dialUpstreamTLS(ctx context.Context, hostPort, sni string) (*tls.Conn, error) { - dialer := &net.Dialer{Timeout: 10 * time.Second} - raw, err := dialer.DialContext(ctx, "tcp", hostPort) +func dialUpstreamTLS(ctx context.Context, hostPort, sni string) (net.Conn, string, error) { + nd := &net.Dialer{Timeout: 10 * time.Second} + raw, err := nd.DialContext(ctx, "tcp", hostPort) if err != nil { - return nil, err + return nil, "", err } conn := tls.Client(raw, &tls.Config{ ServerName: sni, @@ -180,28 +188,103 @@ func dialUpstreamTLS(ctx context.Context, hostPort, sni string) (*tls.Conn, erro }) if err := conn.HandshakeContext(ctx); err != nil { raw.Close() + return nil, "", err + } + return conn, conn.ConnectionState().NegotiatedProtocol, nil +} + +// dialUpstreamPlain connects to a plain (non-TLS) upstream for the +// non-CONNECT proxy path, which is always HTTP/1.1. +func dialUpstreamPlain(ctx context.Context, host string) (net.Conn, string, error) { + if _, _, err := net.SplitHostPort(host); err != nil { + host = net.JoinHostPort(host, "80") + } + nd := &net.Dialer{Timeout: 10 * time.Second} + conn, err := nd.DialContext(ctx, "tcp", host) + return conn, "http/1.1", err +} + +// roundTripH1 writes outReq directly to conn and reads the response back +// off the same connection, wrapping conn in a teeConn so the exact wire +// bytes of both can be captured. +func roundTripH1(conn net.Conn, outReq *http.Request) (*http.Response, *teeConn, error) { + tee := newTeeConn(conn) + if err := outReq.Write(tee); err != nil { + return nil, nil, err + } + resp, err := http.ReadResponse(bufio.NewReader(tee), outReq) + if err != nil { + return nil, nil, err + } + return resp, tee, nil +} + +// roundTripH2 sends outReq over a new single-connection HTTP/2 client. +func roundTripH2(conn net.Conn, outReq *http.Request) (*http.Response, error) { + cc, err := (&http2.Transport{}).NewClientConn(conn) + if err != nil { return nil, err } - return conn, nil + return cc.RoundTrip(outReq) } -// forward sends r upstream via rt and copies the response back to w, -// rewriting r's URL from origin-form (as read off the terminated TLS -// connection) to absolute-form for the round trip. -func (s *Server) forward(rt http.RoundTripper, scheme, hostname string, w http.ResponseWriter, r *http.Request) { +// forward dials upstream, sends r, copies the response back to w, and +// records the exchange to history. r's URL is rewritten from +// origin-form (as read off the terminated connection) to absolute-form +// for the round trip. +func (s *Server) forward(dial dialer, scheme, hostname string, w http.ResponseWriter, r *http.Request) { + clientTee := teeConnFromContext(r.Context()) + outReq := r.Clone(r.Context()) outReq.URL.Scheme = scheme outReq.URL.Host = hostname outReq.RequestURI = "" stripHopByHop(outReq.Header) - resp, err := rt.RoundTrip(outReq) + // Only needed when the client leg isn't tee-captured (HTTP/2): tee + // the body as it streams through so the reconstructed capture isn't + // missing it. + var reqBodyCap *cappedTee + if clientTee == nil && outReq.Body != nil { + reqBodyCap = newCappedTee(outReq.Body) + outReq.Body = io.NopCloser(reqBodyCap) + } + + started := time.Now() + conn, negotiated, dialErr := dial(r.Context()) + if dialErr != nil { + reqRaw, reqExact := captureRequest(r, clientTee, reqBodyCap) + s.record(started, time.Since(started), scheme, hostname, r, reqRaw, reqExact, nil, false, 0, dialErr.Error()) + http.Error(w, dialErr.Error(), http.StatusBadGateway) + return + } + defer conn.Close() + + var resp *http.Response + var upstreamTee *teeConn + var err error + if negotiated == http2.NextProtoTLS { + resp, err = roundTripH2(conn, outReq) + } else { + resp, upstreamTee, err = roundTripH1(conn, outReq) + } + duration := time.Since(started) + + reqRaw, reqExact := captureRequest(r, clientTee, reqBodyCap) + if err != nil { + s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, nil, false, 0, err.Error()) http.Error(w, err.Error(), http.StatusBadGateway) return } defer resp.Body.Close() + var respBodyCap *cappedTee + if upstreamTee == nil { + respBodyCap = newCappedTee(resp.Body) + resp.Body = io.NopCloser(respBodyCap) + } + stripHopByHop(resp.Header) for k, vv := range resp.Header { for _, v := range vv { @@ -210,6 +293,60 @@ func (s *Server) forward(rt http.RoundTripper, scheme, hostname string, w http.R } w.WriteHeader(resp.StatusCode) io.Copy(w, resp.Body) + + var respRaw []byte + var respExact bool + if upstreamTee != nil { + respRaw, respExact = upstreamTee.Take(), true + } else { + respRaw, respExact = captureResponse(resp, nil, respBodyCap) + } + + s.record(started, duration, scheme, hostname, r, reqRaw, reqExact, respRaw, respExact, resp.StatusCode, "") +} + +// record stores one history entry and notifies OnEntry. +func (s *Server) record(started time.Time, duration time.Duration, scheme, host string, r *http.Request, + reqRaw []byte, reqExact bool, respRaw []byte, respExact bool, status int, errMsg string) { + if s.store == nil { + return + } + + e := &store.Entry{ + StartedAt: started, + Duration: duration, + Method: r.Method, + Scheme: scheme, + Host: host, + Path: r.URL.Path, + StatusCode: status, + RequestRaw: reqRaw, + ResponseRaw: respRaw, + RequestExact: reqExact, + ResponseExact: respExact, + Error: errMsg, + } + id, err := s.store.Insert(e) + if err != nil { + log.Printf("store history entry: %v", err) + return + } + + if s.OnEntry != nil { + s.OnEntry(store.Summary{ + ID: id, + StartedAt: e.StartedAt, + Duration: e.Duration, + Method: e.Method, + Scheme: e.Scheme, + Host: e.Host, + Path: e.Path, + StatusCode: e.StatusCode, + ReqSize: len(reqRaw), + RespSize: len(respRaw), + Error: errMsg, + }) + } } // singleConnListener adapts one already-accepted net.Conn into a @@ -220,9 +357,14 @@ type singleConnListener struct { addr net.Addr } +// newSingleConnListener wraps c for one Accept, teeConn on the outside +// so a *teeConn is what ConnContext sees (see withClientTee) - wrapping +// it the other way around lets closeSignalConn's concrete type mask the +// teeConn from that type assertion, silently disabling capture. func newSingleConnListener(c net.Conn) *singleConnListener { ch := make(chan net.Conn, 1) - ch <- &closeSignalConn{Conn: c, onClose: sync.OnceFunc(func() { close(ch) })} + signaled := &closeSignalConn{Conn: c, onClose: sync.OnceFunc(func() { close(ch) })} + ch <- newTeeConn(signaled) return &singleConnListener{ch: ch, addr: c.LocalAddr()} } @@ -248,33 +390,18 @@ func (c *closeSignalConn) Close() error { return err } -// handleHTTP forwards a plain (non-CONNECT) proxy request and copies the -// response back unmodified. +// handleHTTP forwards a plain (non-CONNECT) proxy request, copies the +// response back, and records it to history. func (s *Server) handleHTTP(w http.ResponseWriter, r *http.Request) { if !r.URL.IsAbs() { http.Error(w, "mitmux: request must use absolute-form URI (configure as a proxy, not a target)", http.StatusBadRequest) return } - - outReq := r.Clone(r.Context()) - outReq.RequestURI = "" - stripHopByHop(outReq.Header) - - resp, err := s.transport.RoundTrip(outReq) - if err != nil { - http.Error(w, err.Error(), http.StatusBadGateway) - return - } - defer resp.Body.Close() - - stripHopByHop(resp.Header) - for k, vv := range resp.Header { - for _, v := range vv { - w.Header().Add(k, v) - } + host := r.URL.Host + dial := func(ctx context.Context) (net.Conn, string, error) { + return dialUpstreamPlain(ctx, host) } - w.WriteHeader(resp.StatusCode) - io.Copy(w, resp.Body) + s.forward(dial, r.URL.Scheme, r.URL.Host, w, r) } func stripHopByHop(h http.Header) { diff --git a/internal/proxy/tee.go b/internal/proxy/tee.go new file mode 100644 index 0000000..f6a385a --- /dev/null +++ b/internal/proxy/tee.go @@ -0,0 +1,92 @@ +package proxy + +import ( + "context" + "net" + "sync" +) + +// maxCaptureBytes bounds how much of any single request or response +// mitmux buffers for history storage, independent of how much data +// actually flows through the proxy. Proxying itself always streams the +// full body regardless of this limit - only what gets stored is capped, +// so a multi-gigabyte download can't be turned into a memory exhaustion +// vector just because the history view wants to remember it. +const maxCaptureBytes = 10 << 20 // 10 MiB + +// teeConn wraps a net.Conn, recording every byte read off the wire (up +// to maxCaptureBytes) so it can be attributed to a specific request or +// response later. Take returns everything recorded since the last call +// and resets the buffer, so callers must take exactly once per message +// they want attributed correctly - see forward() for why that's safe +// here (call sites synchronize on the request/response boundary itself). +type teeConn struct { + net.Conn + mu sync.Mutex + buf []byte +} + +func newTeeConn(c net.Conn) *teeConn { + return &teeConn{Conn: c} +} + +func (c *teeConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + if n > 0 { + c.mu.Lock() + if room := maxCaptureBytes - len(c.buf); room > 0 { + end := n + if end > room { + end = room + } + c.buf = append(c.buf, p[:end]...) + } + c.mu.Unlock() + } + return n, err +} + +// Take returns the bytes read since the last Take call (or since the +// connection was created) and resets the buffer. +func (c *teeConn) Take() []byte { + c.mu.Lock() + defer c.mu.Unlock() + out := c.buf + c.buf = nil + return out +} + +// teeListener wraps a net.Listener so every accepted connection is +// tee-captured. +type teeListener struct { + net.Listener +} + +func (l *teeListener) Accept() (net.Conn, error) { + c, err := l.Listener.Accept() + if err != nil { + return nil, err + } + return newTeeConn(c), nil +} + +type contextKey int + +const clientTeeKey contextKey = iota + +// teeConnFromContext returns the teeConn wrapping the client connection +// the current request was read from, as attached via http.Server's +// ConnContext hook. Returns nil for HTTP/2 client connections, which +// aren't tee-captured (see capture.go). +func teeConnFromContext(ctx context.Context) *teeConn { + tc, _ := ctx.Value(clientTeeKey).(*teeConn) + return tc +} + +func withClientTee(ctx context.Context, c net.Conn) context.Context { + tc, ok := c.(*teeConn) + if !ok { + return ctx + } + return context.WithValue(ctx, clientTeeKey, tc) +} diff --git a/internal/store/store.go b/internal/store/store.go new file mode 100644 index 0000000..c0219e4 --- /dev/null +++ b/internal/store/store.go @@ -0,0 +1,186 @@ +// Package store persists proxy history to SQLite in WAL mode. Request +// and response bytes are stored as-received where possible (see the +// Exact fields) rather than re-serialized from a parsed representation. +package store + +import ( + "database/sql" + "fmt" + "time" + + _ "modernc.org/sqlite" +) + +const schema = ` +CREATE TABLE IF NOT EXISTS history ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + started_at INTEGER NOT NULL, + duration_ms INTEGER NOT NULL, + method TEXT NOT NULL, + scheme TEXT NOT NULL, + host TEXT NOT NULL, + path TEXT NOT NULL, + status_code INTEGER, + request_raw BLOB NOT NULL, + response_raw BLOB, + request_exact INTEGER NOT NULL, + response_exact INTEGER NOT NULL, + error TEXT NOT NULL DEFAULT '' +); +` + +// Store is a handle to the history database. Safe for concurrent use. +type Store struct { + db *sql.DB +} + +// Open opens (creating if needed) the SQLite database at path in WAL mode. +func Open(path string) (*Store, error) { + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, fmt.Errorf("open db: %w", err) + } + // modernc.org/sqlite has no real connection pooling benefit here and + // SQLite only supports one writer at a time; serializing access + // through a single connection avoids SQLITE_BUSY entirely. + db.SetMaxOpenConns(1) + + for _, pragma := range []string{ + "PRAGMA journal_mode = WAL", + "PRAGMA synchronous = NORMAL", + "PRAGMA foreign_keys = ON", + } { + if _, err := db.Exec(pragma); err != nil { + db.Close() + return nil, fmt.Errorf("%s: %w", pragma, err) + } + } + if _, err := db.Exec(schema); err != nil { + db.Close() + return nil, fmt.Errorf("create schema: %w", err) + } + return &Store{db: db}, nil +} + +// Close closes the underlying database. +func (s *Store) Close() error { + return s.db.Close() +} + +// Entry is one captured request/response pair. +type Entry struct { + ID int64 + StartedAt time.Time + Duration time.Duration + Method string + Scheme string + Host string + Path string + StatusCode int // 0 if no response was received + RequestRaw []byte + ResponseRaw []byte // nil if no response was received + RequestExact bool // true if RequestRaw is wire-exact, false if reconstructed (e.g. HTTP/2) + ResponseExact bool + Error string // network/transport error, if the request never got a response +} + +// Summary is the lightweight metadata used for the history list view - +// no request/response bodies. +type Summary struct { + ID int64 + StartedAt time.Time + Duration time.Duration + Method string + Scheme string + Host string + Path string + StatusCode int + ReqSize int + RespSize int + Error string +} + +// Insert stores e and returns its assigned ID. +func (s *Store) Insert(e *Entry) (int64, error) { + var statusCode any + if e.StatusCode != 0 { + statusCode = e.StatusCode + } + res, err := s.db.Exec( + `INSERT INTO history + (started_at, duration_ms, method, scheme, host, path, status_code, + request_raw, response_raw, request_exact, response_exact, error) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + e.StartedAt.UnixMilli(), e.Duration.Milliseconds(), e.Method, e.Scheme, e.Host, e.Path, + statusCode, e.RequestRaw, e.ResponseRaw, boolToInt(e.RequestExact), boolToInt(e.ResponseExact), e.Error, + ) + if err != nil { + return 0, fmt.Errorf("insert history entry: %w", err) + } + return res.LastInsertId() +} + +// List returns up to limit history summaries older than beforeID (or the +// most recent if beforeID is 0), newest first. +func (s *Store) List(limit int, beforeID int64) ([]Summary, error) { + if limit <= 0 || limit > 1000 { + limit = 200 + } + if beforeID <= 0 { + beforeID = 1<<63 - 1 + } + rows, err := s.db.Query( + `SELECT id, started_at, duration_ms, method, scheme, host, path, + COALESCE(status_code, 0), length(request_raw), COALESCE(length(response_raw), 0), error + FROM history WHERE id < ? ORDER BY id DESC LIMIT ?`, + beforeID, limit, + ) + if err != nil { + return nil, fmt.Errorf("list history: %w", err) + } + defer rows.Close() + + var out []Summary + for rows.Next() { + var sum Summary + var startedAt, durationMs int64 + if err := rows.Scan(&sum.ID, &startedAt, &durationMs, &sum.Method, &sum.Scheme, &sum.Host, &sum.Path, + &sum.StatusCode, &sum.ReqSize, &sum.RespSize, &sum.Error); err != nil { + return nil, fmt.Errorf("scan history row: %w", err) + } + sum.StartedAt = time.UnixMilli(startedAt) + sum.Duration = time.Duration(durationMs) * time.Millisecond + out = append(out, sum) + } + return out, rows.Err() +} + +// Get returns the full entry (including raw bytes) for id. +func (s *Store) Get(id int64) (*Entry, error) { + row := s.db.QueryRow( + `SELECT id, started_at, duration_ms, method, scheme, host, path, + COALESCE(status_code, 0), request_raw, response_raw, + request_exact, response_exact, error + FROM history WHERE id = ?`, + id, + ) + var e Entry + var startedAt, durationMs int64 + var reqExact, respExact int + if err := row.Scan(&e.ID, &startedAt, &durationMs, &e.Method, &e.Scheme, &e.Host, &e.Path, + &e.StatusCode, &e.RequestRaw, &e.ResponseRaw, &reqExact, &respExact, &e.Error); err != nil { + return nil, fmt.Errorf("get history entry %d: %w", id, err) + } + e.StartedAt = time.UnixMilli(startedAt) + e.Duration = time.Duration(durationMs) * time.Millisecond + e.RequestExact = reqExact != 0 + e.ResponseExact = respExact != 0 + return &e, nil +} + +func boolToInt(b bool) int { + if b { + return 1 + } + return 0 +} |