- WriteUpstream masks vars, hashes long base64 JSON strings and refuses unmasked Authorization/X-Api-Key - multipart requests recorded as ordered parts; the fake compares parts and hashed payloads - loopback CONNECT recording proxy with a local ECDSA parity CA, script and forward modes - parity:upstream command, README and parity docs
594 lines
18 KiB
Go
594 lines
18 KiB
Go
package tide
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"context"
|
|
"crypto/ecdsa"
|
|
"crypto/elliptic"
|
|
"crypto/rand"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"crypto/x509/pkix"
|
|
"encoding/pem"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"math/big"
|
|
"net"
|
|
"net/http"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/goccy/go-yaml"
|
|
|
|
"git.golem15.com/golem15/summercms/modules/fetchguard"
|
|
)
|
|
|
|
// DefaultUpstreamProxyListen is the loopback address the recording upstream
|
|
// proxy binds when no other address is given.
|
|
const DefaultUpstreamProxyListen = "127.0.0.1:8425"
|
|
|
|
const (
|
|
parityCAFile = "parity-ca.pem"
|
|
parityCAKeyFile = "parity-ca-key.pem"
|
|
// upstreamProxyNoMatch is the status a script-mode proxy answers when no
|
|
// scripted response matches a request.
|
|
upstreamProxyNoMatch = 599
|
|
)
|
|
|
|
// UpstreamScript is a hand-authored list of vendor responses the recording
|
|
// proxy answers in script mode, so recording needs no real vendor and no
|
|
// real credential.
|
|
type UpstreamScript struct {
|
|
Responses []UpstreamScriptResponse `yaml:"responses"`
|
|
}
|
|
|
|
// UpstreamScriptResponse answers the first request whose method, host and
|
|
// path match. Each entry answers once.
|
|
type UpstreamScriptResponse struct {
|
|
Method string `yaml:"method"`
|
|
Host string `yaml:"host"`
|
|
Path string `yaml:"path"`
|
|
Response UpstreamResponse `yaml:"response"`
|
|
}
|
|
|
|
// UpstreamProxyConfig configures NewUpstreamProxy.
|
|
type UpstreamProxyConfig struct {
|
|
// Listen is the loopback address to bind (DefaultUpstreamProxyListen).
|
|
Listen string
|
|
// CADir holds the locally generated parity CA (EnsureParityCA). It must
|
|
// be outside the directory that holds Out.
|
|
CADir string
|
|
// Out is the sidecar path Flush writes.
|
|
Out string
|
|
// Mode is "script" (default: answer from Script) or "forward" (send each
|
|
// request once to the real vendor through a guarded client).
|
|
Mode string
|
|
// Script is the UpstreamScript YAML path, required in script mode.
|
|
Script string
|
|
// VarsPath is the private variable store whose values are masked as
|
|
// {{name}} in the sidecar. It must be outside the directory that holds Out.
|
|
VarsPath string
|
|
}
|
|
|
|
// UpstreamProxy is a loopback HTTPS recording proxy. The reference backend
|
|
// sends its vendor calls through it (HTTPS_PROXY plus the parity CA); the
|
|
// proxy terminates TLS with a per-host certificate signed by the parity CA,
|
|
// answers each decrypted request from the script or the real vendor, and
|
|
// keeps every exchange for Flush.
|
|
type UpstreamProxy struct {
|
|
cfg UpstreamProxyConfig
|
|
store *Store
|
|
ca *x509.Certificate
|
|
caKey *ecdsa.PrivateKey
|
|
leafKey *ecdsa.PrivateKey
|
|
forward *fetchguard.Client
|
|
|
|
mu sync.Mutex
|
|
script []UpstreamScriptResponse
|
|
used []bool
|
|
exchanges []UpstreamExchange
|
|
errs []error
|
|
leaves map[string]*tls.Certificate
|
|
conns map[net.Conn]struct{}
|
|
}
|
|
|
|
// NewUpstreamProxy validates cfg, loads the script and the vars store and
|
|
// creates or reuses the parity CA.
|
|
func NewUpstreamProxy(cfg UpstreamProxyConfig) (*UpstreamProxy, error) {
|
|
if cfg.Listen == "" {
|
|
cfg.Listen = DefaultUpstreamProxyListen
|
|
}
|
|
if err := requireLoopbackAddr(cfg.Listen); err != nil {
|
|
return nil, err
|
|
}
|
|
if cfg.Mode == "" {
|
|
cfg.Mode = "script"
|
|
}
|
|
if cfg.Mode != "script" && cfg.Mode != "forward" {
|
|
return nil, fmt.Errorf("tide: upstream proxy mode %q must be script or forward", cfg.Mode)
|
|
}
|
|
for _, req := range []struct{ name, val string }{{"out", cfg.Out}, {"ca-dir", cfg.CADir}, {"vars", cfg.VarsPath}} {
|
|
if strings.TrimSpace(req.val) == "" {
|
|
return nil, fmt.Errorf("tide: upstream proxy requires %s", req.name)
|
|
}
|
|
}
|
|
outDir := filepath.Dir(cfg.Out)
|
|
if err := pathOutside("ca dir", cfg.CADir, outDir); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := pathOutside("vars file", cfg.VarsPath, outDir); err != nil {
|
|
return nil, err
|
|
}
|
|
p := &UpstreamProxy{cfg: cfg, leaves: map[string]*tls.Certificate{}, conns: map[net.Conn]struct{}{}}
|
|
if cfg.Mode == "script" {
|
|
if strings.TrimSpace(cfg.Script) == "" {
|
|
return nil, fmt.Errorf("tide: upstream proxy script mode requires script")
|
|
}
|
|
script, err := loadUpstreamScript(cfg.Script)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
p.script = script.Responses
|
|
p.used = make([]bool, len(script.Responses))
|
|
} else {
|
|
client, err := fetchguard.NewClient(fetchguard.Policy{
|
|
Mode: fetchguard.PublicOnlyMode,
|
|
MaxBytes: MaxUpstreamBody,
|
|
Timeout: 120 * time.Second,
|
|
}, nil)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
p.forward = client
|
|
}
|
|
certPath, err := EnsureParityCA(cfg.CADir)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if err := p.loadCA(certPath); err != nil {
|
|
return nil, err
|
|
}
|
|
p.leafKey, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("tide: leaf key: %w", err)
|
|
}
|
|
p.store, err = OpenStore(cfg.VarsPath)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return p, nil
|
|
}
|
|
|
|
func loadUpstreamScript(path string) (UpstreamScript, error) {
|
|
raw, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return UpstreamScript{}, fmt.Errorf("tide: read upstream script: %w", err)
|
|
}
|
|
var s UpstreamScript
|
|
dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField())
|
|
if err := dec.Decode(&s); err != nil {
|
|
return UpstreamScript{}, fmt.Errorf("tide: parse upstream script %s: %w", path, err)
|
|
}
|
|
for i, r := range s.Responses {
|
|
if r.Method == "" || r.Host == "" || r.Path == "" || r.Response.Status == 0 {
|
|
return UpstreamScript{}, fmt.Errorf("tide: upstream script %s: response %d needs method, host, path and response.status", path, i)
|
|
}
|
|
}
|
|
return s, nil
|
|
}
|
|
|
|
func pathOutside(label, path, dir string) error {
|
|
absPath, err := resolvePath(path)
|
|
if err != nil {
|
|
return fmt.Errorf("tide: %s: %w", label, err)
|
|
}
|
|
absDir, err := resolvePath(dir)
|
|
if err != nil {
|
|
return fmt.Errorf("tide: %s: %w", label, err)
|
|
}
|
|
if absPath == absDir || strings.HasPrefix(absPath, absDir+string(os.PathSeparator)) {
|
|
return fmt.Errorf("tide: %s %q must be outside the output directory %q", label, path, dir)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// EnsureParityCA creates, or reuses when both files exist, the ECDSA P-256
|
|
// self-signed parity CA in dir: parity-ca.pem (mode 0644) and
|
|
// parity-ca-key.pem (mode 0600). It returns the certificate path, which the
|
|
// reference backend trusts while recording. Keep dir outside the fixtures
|
|
// tree; the key never belongs in git.
|
|
func EnsureParityCA(dir string) (certPath string, err error) {
|
|
if strings.TrimSpace(dir) == "" {
|
|
return "", fmt.Errorf("tide: parity CA dir is required")
|
|
}
|
|
if err := os.MkdirAll(dir, 0o700); err != nil {
|
|
return "", fmt.Errorf("tide: parity CA dir: %w", err)
|
|
}
|
|
certPath = filepath.Join(dir, parityCAFile)
|
|
keyPath := filepath.Join(dir, parityCAKeyFile)
|
|
_, certErr := os.Stat(certPath)
|
|
_, keyErr := os.Stat(keyPath)
|
|
if certErr == nil && keyErr == nil {
|
|
if _, _, err := readParityCA(certPath, keyPath); err != nil {
|
|
return "", err
|
|
}
|
|
if err := os.Chmod(keyPath, 0o600); err != nil {
|
|
return "", fmt.Errorf("tide: chmod parity CA key: %w", err)
|
|
}
|
|
return certPath, nil
|
|
}
|
|
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
if err != nil {
|
|
return "", fmt.Errorf("tide: parity CA key: %w", err)
|
|
}
|
|
serial, err := randomSerial()
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
now := time.Now()
|
|
tmpl := &x509.Certificate{
|
|
SerialNumber: serial,
|
|
Subject: pkix.Name{CommonName: "SummerCMS parity upstream CA", Organization: []string{"SummerCMS parity"}},
|
|
NotBefore: now.Add(-time.Hour),
|
|
NotAfter: now.AddDate(2, 0, 0),
|
|
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign | x509.KeyUsageDigitalSignature,
|
|
BasicConstraintsValid: true,
|
|
IsCA: true,
|
|
MaxPathLenZero: true,
|
|
}
|
|
der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
|
|
if err != nil {
|
|
return "", fmt.Errorf("tide: parity CA certificate: %w", err)
|
|
}
|
|
keyDER, err := x509.MarshalECPrivateKey(key)
|
|
if err != nil {
|
|
return "", fmt.Errorf("tide: parity CA key: %w", err)
|
|
}
|
|
if err := writeFileMode(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}), 0o600); err != nil {
|
|
return "", err
|
|
}
|
|
if err := writeFileMode(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o644); err != nil {
|
|
return "", err
|
|
}
|
|
return certPath, nil
|
|
}
|
|
|
|
func writeFileMode(path string, data []byte, mode os.FileMode) error {
|
|
if err := os.WriteFile(path, data, mode); err != nil {
|
|
return fmt.Errorf("tide: write %s: %w", path, err)
|
|
}
|
|
if err := os.Chmod(path, mode); err != nil {
|
|
return fmt.Errorf("tide: chmod %s: %w", path, err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func readParityCA(certPath, keyPath string) (*x509.Certificate, *ecdsa.PrivateKey, error) {
|
|
certPEM, err := os.ReadFile(certPath)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("tide: read parity CA: %w", err)
|
|
}
|
|
keyPEM, err := os.ReadFile(keyPath)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("tide: read parity CA key: %w", err)
|
|
}
|
|
cb, _ := pem.Decode(certPEM)
|
|
kb, _ := pem.Decode(keyPEM)
|
|
if cb == nil || kb == nil {
|
|
return nil, nil, fmt.Errorf("tide: parity CA files in %s are not PEM", filepath.Dir(certPath))
|
|
}
|
|
cert, err := x509.ParseCertificate(cb.Bytes)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("tide: parse parity CA: %w", err)
|
|
}
|
|
key, err := x509.ParseECPrivateKey(kb.Bytes)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("tide: parse parity CA key: %w", err)
|
|
}
|
|
if !cert.IsCA {
|
|
return nil, nil, fmt.Errorf("tide: %s is not a CA certificate", certPath)
|
|
}
|
|
return cert, key, nil
|
|
}
|
|
|
|
func (p *UpstreamProxy) loadCA(certPath string) error {
|
|
cert, key, err := readParityCA(certPath, filepath.Join(filepath.Dir(certPath), parityCAKeyFile))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
p.ca, p.caKey = cert, key
|
|
return nil
|
|
}
|
|
|
|
func randomSerial() (*big.Int, error) {
|
|
n, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 126))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("tide: certificate serial: %w", err)
|
|
}
|
|
return n, nil
|
|
}
|
|
|
|
// leaf returns a certificate for host signed by the parity CA.
|
|
func (p *UpstreamProxy) leaf(host string) (*tls.Certificate, error) {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
if c, ok := p.leaves[host]; ok {
|
|
return c, nil
|
|
}
|
|
serial, err := randomSerial()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
now := time.Now()
|
|
tmpl := &x509.Certificate{
|
|
SerialNumber: serial,
|
|
Subject: pkix.Name{CommonName: host},
|
|
NotBefore: now.Add(-time.Hour),
|
|
NotAfter: now.Add(24 * time.Hour),
|
|
KeyUsage: x509.KeyUsageDigitalSignature,
|
|
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
|
}
|
|
if ip := net.ParseIP(host); ip != nil {
|
|
tmpl.IPAddresses = []net.IP{ip}
|
|
} else {
|
|
tmpl.DNSNames = []string{host}
|
|
}
|
|
der, err := x509.CreateCertificate(rand.Reader, tmpl, p.ca, &p.leafKey.PublicKey, p.caKey)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("tide: leaf certificate for %s: %w", host, err)
|
|
}
|
|
c := &tls.Certificate{Certificate: [][]byte{der, p.ca.Raw}, PrivateKey: p.leafKey}
|
|
p.leaves[host] = c
|
|
return c, nil
|
|
}
|
|
|
|
// ServeHTTP accepts CONNECT tunnels only; every other request is refused.
|
|
func (p *UpstreamProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodConnect {
|
|
http.Error(w, "tide upstream proxy accepts CONNECT only", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
host, port, err := net.SplitHostPort(r.Host)
|
|
if err != nil {
|
|
host, port = r.Host, "443"
|
|
}
|
|
hj, ok := w.(http.Hijacker)
|
|
if !ok {
|
|
http.Error(w, "hijacking unsupported", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
conn, rw, err := hj.Hijack()
|
|
if err != nil {
|
|
return
|
|
}
|
|
p.track(conn, true)
|
|
defer func() {
|
|
p.track(conn, false)
|
|
_ = conn.Close()
|
|
}()
|
|
if _, err := rw.WriteString("HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil {
|
|
return
|
|
}
|
|
if err := rw.Flush(); err != nil {
|
|
return
|
|
}
|
|
tlsConn := tls.Server(conn, &tls.Config{
|
|
MinVersion: tls.VersionTLS12,
|
|
NextProtos: []string{"http/1.1"},
|
|
GetCertificate: func(hello *tls.ClientHelloInfo) (*tls.Certificate, error) {
|
|
name := hello.ServerName
|
|
if name == "" {
|
|
name = host
|
|
}
|
|
return p.leaf(name)
|
|
},
|
|
})
|
|
if err := tlsConn.HandshakeContext(r.Context()); err != nil {
|
|
return
|
|
}
|
|
authority := host
|
|
if port != "443" {
|
|
authority = net.JoinHostPort(host, port)
|
|
}
|
|
br := bufio.NewReader(tlsConn)
|
|
for {
|
|
req, err := http.ReadRequest(br)
|
|
if err != nil {
|
|
return
|
|
}
|
|
resp := p.exchange(r.Context(), authority, req)
|
|
werr := resp.Write(tlsConn)
|
|
if werr != nil || req.Close || resp.Close {
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func (p *UpstreamProxy) track(c net.Conn, add bool) {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
if add {
|
|
p.conns[c] = struct{}{}
|
|
} else {
|
|
delete(p.conns, c)
|
|
}
|
|
}
|
|
|
|
// exchange answers one decrypted request and records it.
|
|
func (p *UpstreamProxy) exchange(ctx context.Context, authority string, req *http.Request) *http.Response {
|
|
body, err := io.ReadAll(io.LimitReader(req.Body, MaxUpstreamBody+1))
|
|
_ = req.Body.Close()
|
|
if err == nil && len(body) > MaxUpstreamBody {
|
|
err = fmt.Errorf("request body exceeds %d bytes", MaxUpstreamBody)
|
|
}
|
|
rawURL := "https://" + authority + req.URL.RequestURI()
|
|
if err != nil {
|
|
return p.failure(req, fmt.Errorf("tide: upstream proxy %s %s: %w", req.Method, redactURL(req.URL), err), http.StatusBadGateway)
|
|
}
|
|
recorded := UpstreamRequest{Method: req.Method, URL: rawURL, Headers: map[string]string{}}
|
|
for _, name := range UpstreamCompareHeaders {
|
|
if v := req.Header.Get(name); v != "" {
|
|
recorded.Headers[name] = v
|
|
}
|
|
}
|
|
if isMultipart(req.Header.Get("Content-Type")) {
|
|
parts, err := upstreamParts(req.Header.Get("Content-Type"), body)
|
|
if err != nil {
|
|
return p.failure(req, fmt.Errorf("tide: upstream proxy %s %s: %w", req.Method, rawURL, err), http.StatusBadGateway)
|
|
}
|
|
recorded.Parts = parts
|
|
recorded.Headers["Content-Type"] = "multipart/form-data"
|
|
} else {
|
|
recorded.Body = string(body)
|
|
}
|
|
|
|
var answer UpstreamResponse
|
|
if p.cfg.Mode == "forward" {
|
|
answer, err = p.forwardOnce(ctx, req, rawURL, body)
|
|
if err != nil {
|
|
return p.failure(req, fmt.Errorf("tide: upstream proxy forward %s %s: %w", req.Method, redactURL(req.URL), err), http.StatusBadGateway)
|
|
}
|
|
} else {
|
|
var ok bool
|
|
answer, ok = p.scripted(req.Method, strings.Split(authority, ":")[0], req.URL.Path)
|
|
if !ok {
|
|
return p.failure(req, fmt.Errorf("tide: upstream proxy: no scripted response for %s https://%s%s", req.Method, authority, req.URL.Path), upstreamProxyNoMatch)
|
|
}
|
|
}
|
|
p.mu.Lock()
|
|
p.exchanges = append(p.exchanges, UpstreamExchange{Request: recorded, Response: answer})
|
|
p.mu.Unlock()
|
|
return httpResponse(req, answer)
|
|
}
|
|
|
|
func (p *UpstreamProxy) scripted(method, host, path string) (UpstreamResponse, bool) {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
for i, s := range p.script {
|
|
if p.used[i] || !strings.EqualFold(s.Method, method) || !strings.EqualFold(s.Host, host) || s.Path != path {
|
|
continue
|
|
}
|
|
p.used[i] = true
|
|
return s.Response, true
|
|
}
|
|
return UpstreamResponse{}, false
|
|
}
|
|
|
|
func (p *UpstreamProxy) forwardOnce(ctx context.Context, in *http.Request, rawURL string, body []byte) (UpstreamResponse, error) {
|
|
out, err := http.NewRequestWithContext(ctx, in.Method, rawURL, bytes.NewReader(body))
|
|
if err != nil {
|
|
return UpstreamResponse{}, err
|
|
}
|
|
for k, vs := range in.Header {
|
|
switch http.CanonicalHeaderKey(k) {
|
|
case "Connection", "Proxy-Connection", "Proxy-Authorization", "Keep-Alive", "Te", "Trailer", "Transfer-Encoding", "Upgrade", "Accept-Encoding":
|
|
continue
|
|
}
|
|
out.Header[k] = append([]string(nil), vs...)
|
|
}
|
|
res, err := p.forward.Send(out)
|
|
if err != nil {
|
|
return UpstreamResponse{}, err
|
|
}
|
|
headers := map[string]string{}
|
|
for k := range res.Header {
|
|
if keepUpstreamResponseHeader(k) {
|
|
headers[http.CanonicalHeaderKey(k)] = res.Header.Get(k)
|
|
}
|
|
}
|
|
return UpstreamResponse{Status: res.StatusCode, Headers: headers, Body: string(res.Body)}, nil
|
|
}
|
|
|
|
// keepUpstreamResponseHeader keeps the response headers vendor clients read.
|
|
func keepUpstreamResponseHeader(name string) bool {
|
|
l := strings.ToLower(name)
|
|
return l == "content-type" || l == "retry-after" || l == "location" || strings.Contains(l, "ratelimit")
|
|
}
|
|
|
|
func (p *UpstreamProxy) failure(req *http.Request, err error, status int) *http.Response {
|
|
p.mu.Lock()
|
|
p.errs = append(p.errs, err)
|
|
p.mu.Unlock()
|
|
return httpResponse(req, UpstreamResponse{
|
|
Status: status,
|
|
Headers: map[string]string{"Content-Type": "text/plain; charset=utf-8"},
|
|
Body: "tide upstream proxy: " + err.Error() + "\n",
|
|
})
|
|
}
|
|
|
|
func httpResponse(req *http.Request, r UpstreamResponse) *http.Response {
|
|
h := http.Header{}
|
|
for k, v := range r.Headers {
|
|
h.Set(k, v)
|
|
}
|
|
return &http.Response{
|
|
Status: fmt.Sprintf("%d %s", r.Status, http.StatusText(r.Status)),
|
|
StatusCode: r.Status,
|
|
Proto: "HTTP/1.1",
|
|
ProtoMajor: 1,
|
|
ProtoMinor: 1,
|
|
Header: h,
|
|
Body: io.NopCloser(strings.NewReader(r.Body)),
|
|
ContentLength: int64(len(r.Body)),
|
|
Request: req,
|
|
}
|
|
}
|
|
|
|
// ListenAndServe binds the loopback listener and serves until ctx is
|
|
// cancelled. It does not write the sidecar; call Flush afterwards.
|
|
func (p *UpstreamProxy) ListenAndServe(ctx context.Context) error {
|
|
ln, err := net.Listen("tcp", p.cfg.Listen)
|
|
if err != nil {
|
|
return fmt.Errorf("tide: listen %s: %w", p.cfg.Listen, err)
|
|
}
|
|
if err := requireLoopbackAddr(ln.Addr().String()); err != nil {
|
|
_ = ln.Close()
|
|
return err
|
|
}
|
|
srv := &http.Server{Handler: p, ReadHeaderTimeout: 10 * time.Second}
|
|
errCh := make(chan error, 1)
|
|
go func() { errCh <- srv.Serve(ln) }()
|
|
select {
|
|
case <-ctx.Done():
|
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
_ = srv.Shutdown(shutdownCtx)
|
|
p.closeTunnels()
|
|
<-errCh
|
|
return nil
|
|
case err := <-errCh:
|
|
p.closeTunnels()
|
|
if errors.Is(err, http.ErrServerClosed) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
}
|
|
|
|
func (p *UpstreamProxy) closeTunnels() {
|
|
p.mu.Lock()
|
|
defer p.mu.Unlock()
|
|
for c := range p.conns {
|
|
_ = c.Close()
|
|
}
|
|
}
|
|
|
|
// Flush writes every recorded exchange to Out with WriteUpstream. It refuses
|
|
// to write when any request failed (an unscripted request, a forward error),
|
|
// so a partial recording never lands next to a fixture.
|
|
func (p *UpstreamProxy) Flush() error {
|
|
p.mu.Lock()
|
|
exchanges := append([]UpstreamExchange(nil), p.exchanges...)
|
|
errs := append([]error(nil), p.errs...)
|
|
p.mu.Unlock()
|
|
if len(errs) > 0 {
|
|
return fmt.Errorf("tide: upstream proxy recorded failures, %s not written: %w", p.cfg.Out, errors.Join(errs...))
|
|
}
|
|
return WriteUpstream(p.cfg.Out, UpstreamSidecar{Version: CurrentVersion, Exchanges: exchanges}, p.store)
|
|
}
|