Files
summercms/modules/tide/upstream_proxy.go
Jakub Zych ee0004fb65 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
2026-10-03 19:55:42 +02:00

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)
}