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