Files
summercms/tide/proxy.go
Jakub Zych aa165fe3d0 feat(02-02): capture, scrub, and strictly diff stateful flows
Resolve named placeholders from a private variable store, mask
dates and ids after shape checks, and keep comparing independent steps.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-17 12:32:49 +02:00

381 lines
9.2 KiB
Go

package tide
import (
"bytes"
"context"
"fmt"
"io"
"net"
"net/http"
"net/http/httputil"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"time"
)
const (
DefaultListen = "127.0.0.1:8422"
DefaultUpstream = "http://127.0.0.1:8423"
)
type captureKey struct{}
type captureState struct {
session string
method string
path string
query string
reqBody []byte
reqHeaders http.Header
}
type sessionBuf struct {
name string
steps []Step
started bool
}
// ProxyConfig pins a loopback reverse proxy that records named sessions.
type ProxyConfig struct {
Listen string
Upstream string
Session string
Rules Rules
RulesPath string
VarsPath string
Fixtures string
MaxBody int64
}
// Proxy is a fixed-upstream recording reverse proxy.
type Proxy struct {
cfg ProxyConfig
upstream *url.URL
rp *httputil.ReverseProxy
limit int64
store *Store
mu sync.Mutex
sessions map[string]*sessionBuf
failed map[string]error
}
// NewProxy validates loopback bind/upstream and builds a recording handler.
func NewProxy(cfg ProxyConfig) (*Proxy, error) {
if strings.TrimSpace(cfg.Listen) == "" {
cfg.Listen = DefaultListen
}
if strings.TrimSpace(cfg.Upstream) == "" {
cfg.Upstream = DefaultUpstream
}
if err := requireLoopbackAddr(cfg.Listen); err != nil {
return nil, err
}
upstream, err := parseLoopbackUpstream(cfg.Upstream)
if err != nil {
return nil, err
}
if strings.TrimSpace(cfg.Fixtures) == "" {
return nil, fmt.Errorf("tide: proxy fixtures directory is required")
}
if err := validateRules(cfg.Rules); err != nil {
return nil, err
}
if err := varsOutsideFixtures(cfg.VarsPath, cfg.Fixtures); err != nil {
return nil, err
}
store, err := OpenStore(cfg.VarsPath)
if err != nil {
return nil, err
}
p := &Proxy{
cfg: cfg,
upstream: upstream,
limit: maxBody(cfg.MaxBody),
store: store,
sessions: make(map[string]*sessionBuf),
failed: make(map[string]error),
}
p.rp = &httputil.ReverseProxy{
Rewrite: func(pr *httputil.ProxyRequest) {
pr.SetURL(p.upstream)
pr.Out.Host = p.upstream.Host
pr.Out.Header.Del(SessionHeader)
},
ModifyResponse: p.modifyResponse,
ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) {
http.Error(w, "tide: upstream error", http.StatusBadGateway)
},
}
return p, nil
}
// Handler returns the recording reverse-proxy handler.
func (p *Proxy) Handler() http.Handler {
return p
}
func (p *Proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
state, err := p.beginCapture(r)
if err != nil {
status := http.StatusBadRequest
if isTruncated(err) {
status = http.StatusRequestEntityTooLarge
}
http.Error(w, err.Error(), status)
return
}
body := state.reqBody
r = r.WithContext(context.WithValue(r.Context(), captureKey{}, state))
r.Body = io.NopCloser(bytes.NewReader(body))
r.ContentLength = int64(len(body))
r.GetBody = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(body)), nil
}
p.rp.ServeHTTP(w, r)
}
// ListenAndServe binds the loopback listener until ctx is cancelled, then flushes.
func (p *Proxy) 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)
}
addr := ln.Addr().String()
if err := requireLoopbackAddr(addr); 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)
<-errCh
return p.Flush()
case err := <-errCh:
flushErr := p.Flush()
if err != nil && err != http.ErrServerClosed {
if flushErr != nil {
return fmt.Errorf("%v; flush: %w", err, flushErr)
}
return err
}
return flushErr
}
}
// Flush writes every complete captured session. Failed sessions leave no new fixture.
func (p *Proxy) Flush() error {
p.mu.Lock()
defer p.mu.Unlock()
var first error
for name, buf := range p.sessions {
if err := p.failed[name]; err != nil {
if first == nil {
first = err
}
continue
}
if err := p.writeSessionLocked(buf); err != nil {
if first == nil {
first = err
}
}
}
return first
}
func (p *Proxy) beginCapture(r *http.Request) (*captureState, error) {
session := strings.TrimSpace(r.Header.Get(SessionHeader))
if session == "" {
session = strings.TrimSpace(p.cfg.Session)
}
if err := validateSessionName(session); err != nil {
return nil, err
}
raw, err := readBounded(r.Body, p.limit)
if err != nil {
p.failSession(session, fmt.Errorf("tide: session %q: %w", session, err))
return nil, err
}
return &captureState{
session: session,
method: r.Method,
path: r.URL.Path,
query: r.URL.RawQuery,
reqBody: raw,
reqHeaders: r.Header.Clone(),
}, nil
}
func (p *Proxy) modifyResponse(resp *http.Response) error {
state, _ := resp.Request.Context().Value(captureKey{}).(*captureState)
if state == nil {
return fmt.Errorf("tide: missing capture state")
}
raw, err := readBounded(resp.Body, p.limit)
if err != nil {
resp.Body.Close()
p.failSession(state.session, fmt.Errorf("tide: session %q: %w", state.session, err))
return err
}
resp.Body = io.NopCloser(bytes.NewReader(raw))
resp.ContentLength = int64(len(raw))
resp.Header.Set("Content-Length", fmt.Sprintf("%d", len(raw)))
if err := p.recordStep(state, resp, raw); err != nil {
p.failSession(state.session, err)
return err
}
return nil
}
func (p *Proxy) recordStep(state *captureState, resp *http.Response, respBody []byte) error {
p.mu.Lock()
defer p.mu.Unlock()
if err := p.failed[state.session]; err != nil {
return err
}
buf, err := p.sessionLocked(state.session)
if err != nil {
return err
}
route := p.cfg.Rules.Match(state.method, state.path)
reqHeaders := filterHeaders(state.reqHeaders, p.cfg.Rules.requestHeaders(route))
respHeaders := filterHeaders(resp.Header, p.cfg.Rules.responseHeaders(route))
step := Step{
ID: fmt.Sprintf("%d", len(buf.steps)+1),
Request: Request{
Method: state.method,
Path: state.path,
Query: state.query,
Headers: reqHeaders,
Body: Body(state.reqBody),
},
Response: Response{
Status: resp.StatusCode,
Headers: respHeaders,
Body: Body(respBody),
},
}
if route != nil {
step.Capture = append([]CaptureRule(nil), route.Capture...)
}
if err := CaptureStep(p.store, &step); err != nil {
return fmt.Errorf("tide: session %q: %w", state.session, err)
}
if err := ScrubStep(p.store, &step); err != nil {
return fmt.Errorf("tide: session %q: %w", state.session, err)
}
if err := p.store.Save(); err != nil {
return err
}
buf.steps = append(buf.steps, step)
return p.writeSessionLocked(buf)
}
func (p *Proxy) sessionLocked(name string) (*sessionBuf, error) {
if buf, ok := p.sessions[name]; ok {
return buf, nil
}
dest := p.sessionPath(name)
if _, err := os.Stat(dest); err == nil {
return nil, fmt.Errorf("tide: duplicate session %q", name)
}
buf := &sessionBuf{name: name}
p.sessions[name] = buf
return buf, nil
}
func (p *Proxy) writeSessionLocked(buf *sessionBuf) error {
if buf == nil || len(buf.steps) == 0 {
return nil
}
if err := p.failed[buf.name]; err != nil {
return err
}
flow := Flow{
Version: CurrentVersion,
Name: buf.name,
Steps: append([]Step(nil), buf.steps...),
}
return SaveFlow(p.sessionPath(buf.name), flow)
}
func (p *Proxy) sessionPath(name string) string {
return filepath.Join(p.cfg.Fixtures, p.cfg.Rules.Client, name+".yaml")
}
func (p *Proxy) failSession(name string, err error) {
p.mu.Lock()
defer p.mu.Unlock()
if _, ok := p.failed[name]; !ok {
p.failed[name] = err
}
if buf := p.sessions[name]; buf != nil && len(buf.steps) == 0 {
_ = os.Remove(p.sessionPath(name))
}
}
func validateSessionName(name string) error {
name = strings.TrimSpace(name)
if name == "" {
return fmt.Errorf("tide: session name is required (X-Parity-Session or --session)")
}
if strings.ContainsAny(name, `/\`) || strings.Contains(name, "..") || name != filepath.Base(name) {
return fmt.Errorf("tide: session name %q must not contain path separators", name)
}
return nil
}
func requireLoopbackAddr(hostport string) error {
host, _, err := net.SplitHostPort(hostport)
if err != nil {
return fmt.Errorf("tide: listen %q: %w", hostport, err)
}
if !isLoopbackHost(host) {
return fmt.Errorf("tide: listen %q must be loopback", hostport)
}
return nil
}
func parseLoopbackUpstream(raw string) (*url.URL, error) {
u, err := url.Parse(raw)
if err != nil {
return nil, fmt.Errorf("tide: upstream: %w", err)
}
if u.Scheme != "http" {
return nil, fmt.Errorf("tide: upstream %q must be plain http on loopback", raw)
}
if u.Host == "" || !isLoopbackHost(u.Hostname()) {
return nil, fmt.Errorf("tide: upstream %q must be loopback", raw)
}
return u, nil
}
func isLoopbackHost(host string) bool {
if host == "" {
return false
}
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func isTruncated(err error) bool {
return err != nil && strings.Contains(err.Error(), "exceeds")
}