diff options
Diffstat (limited to 'internal/proxy/repeat.go')
| -rw-r--r-- | internal/proxy/repeat.go | 137 |
1 files changed, 137 insertions, 0 deletions
diff --git a/internal/proxy/repeat.go b/internal/proxy/repeat.go new file mode 100644 index 0000000..6373ee1 --- /dev/null +++ b/internal/proxy/repeat.go @@ -0,0 +1,137 @@ +package proxy + +import ( + "bufio" + "bytes" + "context" + "crypto/tls" + "fmt" + "io" + "net" + "net/http" + "time" + + "mitmux/internal/store" +) + +// Repeat sends raw exactly as given - no framing correction, no header +// injection - to scheme://host, and records the exchange to history +// with Source "repeater". This is the raw-byte send/resend primitive: +// unlike forward(), which round-trips a parsed *http.Request, Repeat +// exists specifically so an edited, possibly malformed request (the +// whole point of a Repeater tool) reaches the wire unmodified. +// +// Repeater only speaks HTTP/1.1: raw edited text has no equivalent in +// HTTP/2's binary framing, so the connection is negotiated HTTP/1.1-only +// rather than letting the server pick. +func (s *Server) Repeat(ctx context.Context, scheme, host string, raw []byte) (*store.Entry, error) { + started := time.Now() + method, path := parseRequestLine(raw) + + conn, err := dialForRepeat(ctx, scheme, host) + if err != nil { + return s.recordRepeat(started, time.Since(started), scheme, host, method, path, raw, nil, 0, err.Error()) + } + defer conn.Close() + + if _, err := conn.Write(raw); err != nil { + return s.recordRepeat(started, time.Since(started), scheme, host, method, path, raw, nil, 0, err.Error()) + } + + tee := newTeeConn(conn) + resp, err := http.ReadResponse(bufio.NewReader(tee), &http.Request{Method: method}) + duration := time.Since(started) + if err != nil { + return s.recordRepeat(started, duration, scheme, host, method, path, raw, nil, 0, err.Error()) + } + defer resp.Body.Close() + io.Copy(io.Discard, resp.Body) + + return s.recordRepeat(started, duration, scheme, host, method, path, raw, tee.Take(), resp.StatusCode, "") +} + +func (s *Server) recordRepeat(started time.Time, duration time.Duration, scheme, host, method, path string, + reqRaw, respRaw []byte, status int, errMsg string) (*store.Entry, error) { + e := &store.Entry{ + StartedAt: started, + Duration: duration, + Method: method, + Scheme: scheme, + Host: host, + Path: path, + StatusCode: status, + RequestRaw: reqRaw, + ResponseRaw: respRaw, + RequestExact: true, + ResponseExact: respRaw != nil, + Error: errMsg, + Source: "repeater", + } + if s.store != nil { + id, err := s.store.Insert(e) + if err != nil { + return nil, fmt.Errorf("store repeater entry: %w", err) + } + e.ID = id + 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, Source: e.Source, + }) + } + } + return e, nil +} + +// dialForRepeat connects to host for scheme, forcing HTTP/1.1 over ALPN +// when TLS is involved (see Repeat's doc comment for why). +func dialForRepeat(ctx context.Context, scheme, host string) (net.Conn, error) { + nd := &net.Dialer{Timeout: 10 * time.Second} + if scheme != "https" { + hostPort := host + if _, _, err := net.SplitHostPort(host); err != nil { + hostPort = net.JoinHostPort(host, "80") + } + return nd.DialContext(ctx, "tcp", hostPort) + } + + hostname, hostPort := host, host + if h, _, err := net.SplitHostPort(host); err == nil { + hostname = h + } else { + hostPort = net.JoinHostPort(host, "443") + } + raw, err := nd.DialContext(ctx, "tcp", hostPort) + if err != nil { + return nil, err + } + conn := tls.Client(raw, &tls.Config{ServerName: hostname, NextProtos: []string{"http/1.1"}}) + if err := conn.HandshakeContext(ctx); err != nil { + raw.Close() + return nil, err + } + return conn, nil +} + +// parseRequestLine extracts the method and request-target from the +// first line of a raw HTTP/1.1 request, without validating or parsing +// anything else - used only to label the stored entry and to tell +// http.ReadResponse whether this was a HEAD request (which changes +// response body framing rules). +func parseRequestLine(raw []byte) (method, path string) { + nl := bytes.IndexByte(raw, '\n') + if nl < 0 { + nl = len(raw) + } + line := bytes.TrimRight(raw[:nl], "\r\n") + fields := bytes.Fields(line) + if len(fields) > 0 { + method = string(fields[0]) + } + if len(fields) > 1 { + path = string(fields[1]) + } + return method, path +} |