feat(14-01): record vendor calls with summer parity:upstream and replay them offline
- 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
This commit is contained in:
593
modules/tide/upstream_proxy.go
Normal file
593
modules/tide/upstream_proxy.go
Normal file
@@ -0,0 +1,593 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user