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