feat(02-02): capture named sessions through a loopback proxy
Record Nuxt/MCP traffic as ordered flows via parity:proxy, pin loopback upstream, and refuse oversized or credential-shaped fixtures. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
544
tide/proxy.go
Normal file
544
tide/proxy.go
Normal file
@@ -0,0 +1,544 @@
|
||||
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
|
||||
|
||||
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 cfg.VarsPath != "" {
|
||||
if err := prepareVarsFile(cfg.VarsPath, cfg.Fixtures); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
p := &Proxy{
|
||||
cfg: cfg,
|
||||
upstream: upstream,
|
||||
limit: maxBody(cfg.MaxBody),
|
||||
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 := rejectUnclassifiedCredentials(step); err != nil {
|
||||
return fmt.Errorf("tide: session %q: %w", state.session, 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 prepareVarsFile(path, fixtures string) error {
|
||||
absVars, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: vars path: %w", err)
|
||||
}
|
||||
absFix, err := filepath.Abs(fixtures)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: fixtures path: %w", err)
|
||||
}
|
||||
if absVars == absFix || strings.HasPrefix(absVars, absFix+string(os.PathSeparator)) {
|
||||
return fmt.Errorf("tide: vars file %q must be outside fixtures %q", path, fixtures)
|
||||
}
|
||||
if st, err := os.Stat(absVars); err == nil {
|
||||
if st.IsDir() {
|
||||
return fmt.Errorf("tide: vars %q is a directory", path)
|
||||
}
|
||||
if err := os.Chmod(absVars, 0o600); err != nil {
|
||||
return fmt.Errorf("tide: chmod vars: %w", err)
|
||||
}
|
||||
return nil
|
||||
} else if !os.IsNotExist(err) {
|
||||
return fmt.Errorf("tide: stat vars: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(absVars), 0o755); err != nil {
|
||||
return fmt.Errorf("tide: create vars dir: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(absVars, []byte("{}\n"), 0o600); err != nil {
|
||||
return fmt.Errorf("tide: create vars: %w", err)
|
||||
}
|
||||
return os.Chmod(absVars, 0o600)
|
||||
}
|
||||
|
||||
func isTruncated(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "exceeds")
|
||||
}
|
||||
|
||||
func rejectUnclassifiedCredentials(step Step) error {
|
||||
var parts []string
|
||||
for _, v := range step.Request.Headers {
|
||||
parts = append(parts, v)
|
||||
}
|
||||
parts = append(parts, step.Request.Query, string(step.Request.Body))
|
||||
for _, v := range step.Response.Headers {
|
||||
parts = append(parts, v)
|
||||
}
|
||||
parts = append(parts, string(step.Response.Body))
|
||||
classified := classifiedNames(step.Capture)
|
||||
for _, part := range parts {
|
||||
if hit := firstCredential(part); hit != "" && !classified[hit] {
|
||||
return fmt.Errorf("unclassified credential-shaped value in step %s", step.ID)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func classifiedNames(rules []CaptureRule) map[string]bool {
|
||||
out := make(map[string]bool)
|
||||
for _, rule := range rules {
|
||||
if rule.Category != "" {
|
||||
out[rule.Category] = true
|
||||
}
|
||||
out[rule.As] = true
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func firstCredential(s string) string {
|
||||
if s == "" {
|
||||
return ""
|
||||
}
|
||||
if jwtRe.MatchString(s) {
|
||||
return "jwt"
|
||||
}
|
||||
if invRe.MatchString(s) {
|
||||
return "token"
|
||||
}
|
||||
lower := strings.ToLower(s)
|
||||
if strings.Contains(lower, "auth_token=") {
|
||||
return "cookie"
|
||||
}
|
||||
if strings.Contains(lower, "client_secret=") {
|
||||
return "oauth_secret"
|
||||
}
|
||||
if strings.Contains(lower, "code_verifier=") {
|
||||
return "pkce"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var (
|
||||
jwtRe = mustCompileJWT()
|
||||
invRe = mustCompileInv()
|
||||
)
|
||||
|
||||
func mustCompileJWT() *regexpJWT {
|
||||
return ®expJWT{}
|
||||
}
|
||||
|
||||
func mustCompileInv() *regexpInv {
|
||||
return ®expInv{}
|
||||
}
|
||||
|
||||
// tiny wrappers keep the credential regexes local without extra files in task 1.
|
||||
type regexpJWT struct{}
|
||||
|
||||
func (regexpJWT) MatchString(s string) bool {
|
||||
return jwtLooksLike(s)
|
||||
}
|
||||
|
||||
type regexpInv struct{}
|
||||
|
||||
func (regexpInv) MatchString(s string) bool {
|
||||
return strings.Contains(s, "inv_") && invLooksLike(s)
|
||||
}
|
||||
|
||||
func jwtLooksLike(s string) bool {
|
||||
const prefix = "eyJ"
|
||||
for i := 0; i < len(s); i++ {
|
||||
j := strings.Index(s[i:], prefix)
|
||||
if j < 0 {
|
||||
return false
|
||||
}
|
||||
i += j
|
||||
if token := jwtAt(s[i:]); token != "" {
|
||||
return true
|
||||
}
|
||||
i++
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func jwtAt(s string) string {
|
||||
parts := 0
|
||||
n := 0
|
||||
for n < len(s) {
|
||||
c := s[n]
|
||||
if isJWTByte(c) {
|
||||
n++
|
||||
continue
|
||||
}
|
||||
if c == '.' {
|
||||
parts++
|
||||
n++
|
||||
if parts > 2 {
|
||||
return ""
|
||||
}
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if parts == 2 && n >= 20 {
|
||||
return s[:n]
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func isJWTByte(c byte) bool {
|
||||
return (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '-' || c == '_'
|
||||
}
|
||||
|
||||
func invLooksLike(s string) bool {
|
||||
for {
|
||||
i := strings.Index(s, "inv_")
|
||||
if i < 0 {
|
||||
return false
|
||||
}
|
||||
rest := s[i+4:]
|
||||
n := 0
|
||||
for n < len(rest) && isJWTByte(rest[n]) {
|
||||
n++
|
||||
}
|
||||
if n >= 8 {
|
||||
return true
|
||||
}
|
||||
s = rest
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user