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