srdusr
aboutsummaryrefslogtreecommitdiffstats
path: root/internal/proxy/repeat.go
blob: 6373ee11466807ffe023e7ce7071756f6688d560 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
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
}