diff --git a/cmd/summer/main.go b/cmd/summer/main.go index af01a6f..3ccea61 100644 --- a/cmd/summer/main.go +++ b/cmd/summer/main.go @@ -29,6 +29,7 @@ func toolCommands() []bonfire.Command { makePluginCommand(), addPluginCommand(), devCommand(), + parityProxyCommand(), parityRecordCommand(), parityReplayCommand(), } diff --git a/cmd/summer/parity.go b/cmd/summer/parity.go index 052926a..e6237e2 100644 --- a/cmd/summer/parity.go +++ b/cmd/summer/parity.go @@ -9,6 +9,22 @@ import ( "git.golem15.com/golem15/summercms/tide" ) +func parityProxyCommand() bonfire.Command { + return bonfire.Command{ + Name: "parity:proxy", + Description: "Record named HTTP sessions through a loopback reverse proxy", + Flags: []bonfire.Flag{ + {Name: "listen", Description: "Loopback bind address", Default: tide.DefaultListen}, + {Name: "upstream", Description: "Fixed loopback HTTP origin", Default: tide.DefaultUpstream}, + {Name: "session", Description: "Default session name when X-Parity-Session is absent"}, + {Name: "rules", Description: "Committed YAML capture rules"}, + {Name: "vars", Description: "Private mode-0600 variable store outside fixtures"}, + {Name: "fixtures", Description: "Directory for nuxt/ and mcp/ session flows"}, + }, + Run: runParityProxy, + } +} + func parityRecordCommand() bonfire.Command { return bonfire.Command{ Name: "parity:record", @@ -34,6 +50,38 @@ func parityReplayCommand() bonfire.Command { } } +func runParityProxy(ctx context.Context, in bonfire.Input, out bonfire.Output) error { + rulesPath, err := requireFlag(in, "rules", "parity:proxy") + if err != nil { + return err + } + fixtures, err := requireFlag(in, "fixtures", "parity:proxy") + if err != nil { + return err + } + rules, err := tide.LoadRules(rulesPath) + if err != nil { + return err + } + listen, _ := in.Flag("listen") + upstream, _ := in.Flag("upstream") + session, _ := in.Flag("session") + varsPath, _ := in.Flag("vars") + proxy, err := tide.NewProxy(tide.ProxyConfig{ + Listen: listen, + Upstream: upstream, + Session: session, + Rules: rules, + VarsPath: varsPath, + Fixtures: fixtures, + }) + if err != nil { + return err + } + out.Info(fmt.Sprintf("proxy listening on %s → %s", listen, upstream)) + return proxy.ListenAndServe(ctx) +} + func runParityRecord(ctx context.Context, in bonfire.Input, out bonfire.Output) error { specPath, err := requireFlag(in, "spec", "parity:record") if err != nil { diff --git a/cmd/summer/parity_test.go b/cmd/summer/parity_test.go index ae3d27a..e06ad9e 100644 --- a/cmd/summer/parity_test.go +++ b/cmd/summer/parity_test.go @@ -14,7 +14,7 @@ import ( func TestParityCommands(t *testing.T) { names := commandNames() - for _, want := range []string{"parity:record", "parity:replay"} { + for _, want := range []string{"parity:record", "parity:replay", "parity:proxy"} { if !containsName(names, want) { t.Fatalf("missing %s in %v", want, names) } @@ -99,6 +99,23 @@ func TestParityCommands(t *testing.T) { if err := runParity("parity:replay", "--fixtures", keyFixture, "--target", keyOrder.URL); err != nil { t.Fatalf("reordered JSON keys must pass: %v", err) } + + err = runParity("parity:proxy") + if err == nil || !strings.Contains(err.Error(), "--rules") { + t.Fatalf("proxy without rules: %v", err) + } + rulesPath := filepath.Join(outDir, "rules.yaml") + if err := os.WriteFile(rulesPath, []byte("client: nuxt\nkeep_request_headers: []\nkeep_response_headers: []\n"), 0o644); err != nil { + t.Fatal(err) + } + err = runParity("parity:proxy", "--rules", rulesPath, "--fixtures", outDir, "--upstream", "http://example.com") + if err == nil || !strings.Contains(err.Error(), "loopback") { + t.Fatalf("proxy non-loopback upstream: %v", err) + } + err = runParity("parity:proxy", "--rules", rulesPath, "--fixtures", outDir, "--listen", "0.0.0.0:8422") + if err == nil || !strings.Contains(err.Error(), "loopback") { + t.Fatalf("proxy non-loopback listen: %v", err) + } } func commandNames() []string { diff --git a/tide/fixture.go b/tide/fixture.go index b659ede..fcc8be5 100644 --- a/tide/fixture.go +++ b/tide/fixture.go @@ -154,8 +154,22 @@ func writeCapture(b *strings.Builder, indent int, rules []CaptureRule) { fmt.Fprintf(b, "%scapture:\n", pad) inner := strings.Repeat(" ", indent+2) for _, rule := range rules { - fmt.Fprintf(b, "%s- path: %s\n", inner, encodeScalar(rule.Path)) - fmt.Fprintf(b, "%s as: %s\n", inner, encodeScalar(rule.As)) + fmt.Fprintf(b, "%s- as: %s\n", inner, encodeScalar(rule.As)) + if rule.From != "" { + fmt.Fprintf(b, "%s from: %s\n", inner, encodeScalar(rule.From)) + } + if rule.Path != "" { + fmt.Fprintf(b, "%s path: %s\n", inner, encodeScalar(rule.Path)) + } + if rule.Name != "" { + fmt.Fprintf(b, "%s name: %s\n", inner, encodeScalar(rule.Name)) + } + if rule.Identity != "" { + fmt.Fprintf(b, "%s identity: %s\n", inner, encodeScalar(rule.Identity)) + } + if rule.Category != "" { + fmt.Fprintf(b, "%s category: %s\n", inner, encodeScalar(rule.Category)) + } } } diff --git a/tide/flow.go b/tide/flow.go index df23968..e54db04 100644 --- a/tide/flow.go +++ b/tide/flow.go @@ -50,10 +50,14 @@ type Response struct { SHA256 string `yaml:"sha256,omitempty"` } -// CaptureRule maps a JSON path in the response onto a variable name. +// CaptureRule maps a named source onto a variable name. type CaptureRule struct { - Path string `yaml:"path"` - As string `yaml:"as"` + From string `yaml:"from,omitempty"` + Path string `yaml:"path,omitempty"` + Name string `yaml:"name,omitempty"` + As string `yaml:"as"` + Identity string `yaml:"identity,omitempty"` + Category string `yaml:"category,omitempty"` } // NormalizeRule names a per-step normalizer override. diff --git a/tide/proxy.go b/tide/proxy.go new file mode 100644 index 0000000..0632986 --- /dev/null +++ b/tide/proxy.go @@ -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 + } +} diff --git a/tide/proxy_test.go b/tide/proxy_test.go new file mode 100644 index 0000000..2ca1c97 --- /dev/null +++ b/tide/proxy_test.go @@ -0,0 +1,368 @@ +package tide + +import ( + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestProxyNamedSessionsOrderedAndIsolated(t *testing.T) { + var seenHost []string + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seenHost = append(seenHost, r.Host) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"path":"` + r.URL.Path + `"}`)) + })) + t.Cleanup(upstream.Close) + + fixtures := t.TempDir() + proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML()) + srv := httptest.NewServer(proxy.Handler()) + t.Cleanup(srv.Close) + + doProxy(t, srv.URL, "alpha", "/one", "") + doProxy(t, srv.URL, "alpha", "/two", "") + doProxy(t, srv.URL, "beta", "/other", "") + + if err := proxy.Flush(); err != nil { + t.Fatal(err) + } + + alpha, err := LoadFlow(filepath.Join(fixtures, "nuxt", "alpha.yaml")) + if err != nil { + t.Fatal(err) + } + if len(alpha.Steps) != 2 { + t.Fatalf("alpha steps: %d", len(alpha.Steps)) + } + if alpha.Steps[0].Request.Path != "/one" || alpha.Steps[1].Request.Path != "/two" { + t.Fatalf("alpha order: %+v", alpha.Steps) + } + beta, err := LoadFlow(filepath.Join(fixtures, "nuxt", "beta.yaml")) + if err != nil { + t.Fatal(err) + } + if len(beta.Steps) != 1 || beta.Steps[0].Request.Path != "/other" { + t.Fatalf("beta interleaved: %+v", beta.Steps) + } + + req, err := http.NewRequest(http.MethodGet, srv.URL+"/evil", nil) + if err != nil { + t.Fatal(err) + } + req.Host = "evil.example" + req.Header.Set(SessionHeader, "gamma") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if len(seenHost) < 4 { + t.Fatalf("upstream hits: %d", len(seenHost)) + } + wantHost := strings.TrimPrefix(upstream.URL, "http://") + if seenHost[len(seenHost)-1] != wantHost { + t.Fatalf("client Host leaked: last=%s want=%s all=%v", seenHost[len(seenHost)-1], wantHost, seenHost) + } +} + +func TestProxyForwardsCookieRedirectAndBody(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/go" { + w.Header().Set("Set-Cookie", "sid=abc; Path=/") + w.Header().Set("Location", "/landed") + w.WriteHeader(http.StatusFound) + _, _ = w.Write([]byte("redirect-body")) + return + } + w.Header().Set("Content-Type", "text/plain") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("cookie=" + r.Header.Get("Cookie"))) + })) + t.Cleanup(upstream.Close) + + fixtures := t.TempDir() + proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML()) + srv := httptest.NewServer(proxy.Handler()) + t.Cleanup(srv.Close) + + client := &http.Client{ + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + } + req, err := http.NewRequest(http.MethodGet, srv.URL+"/go", nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set(SessionHeader, "redir") + req.Header.Set("Cookie", "keep=1") + resp, err := client.Do(req) + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + if resp.StatusCode != http.StatusFound { + t.Fatalf("status %d", resp.StatusCode) + } + if loc := resp.Header.Get("Location"); loc != "/landed" { + t.Fatalf("location %q", loc) + } + if !strings.Contains(resp.Header.Get("Set-Cookie"), "sid=abc") { + t.Fatalf("set-cookie %q", resp.Header.Get("Set-Cookie")) + } + if string(body) != "redirect-body" { + t.Fatalf("body %q", body) + } + + req2, err := http.NewRequest(http.MethodGet, srv.URL+"/echo", nil) + if err != nil { + t.Fatal(err) + } + req2.Header.Set(SessionHeader, "redir") + req2.Header.Set("Cookie", "keep=1") + resp2, err := client.Do(req2) + if err != nil { + t.Fatal(err) + } + got, _ := io.ReadAll(resp2.Body) + resp2.Body.Close() + if string(got) != "cookie=keep=1" { + t.Fatalf("cookie not forwarded: %q", got) + } +} + +func TestProxyRejectsNonLoopbackOverflowAndCredentials(t *testing.T) { + if _, err := NewProxy(ProxyConfig{ + Listen: "0.0.0.0:8422", + Upstream: "http://127.0.0.1:8423", + Fixtures: t.TempDir(), + Rules: mustParseRules(t, testRulesYAML()), + }); err == nil || !strings.Contains(err.Error(), "loopback") { + t.Fatalf("non-loopback listen: %v", err) + } + if _, err := NewProxy(ProxyConfig{ + Listen: "127.0.0.1:8422", + Upstream: "http://example.com", + Fixtures: t.TempDir(), + Rules: mustParseRules(t, testRulesYAML()), + }); err == nil || !strings.Contains(err.Error(), "loopback") { + t.Fatalf("non-loopback upstream: %v", err) + } + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("hello-world")) + })) + t.Cleanup(upstream.Close) + + fixtures := t.TempDir() + cfg := ProxyConfig{ + Listen: "127.0.0.1:0", + Upstream: upstream.URL, + Fixtures: fixtures, + Rules: mustParseRules(t, testRulesYAML()), + MaxBody: 4, + } + proxy, err := NewProxy(cfg) + if err != nil { + t.Fatal(err) + } + srv := httptest.NewServer(proxy.Handler()) + t.Cleanup(srv.Close) + + req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set(SessionHeader, "overflow") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusRequestEntityTooLarge && resp.StatusCode != http.StatusBadGateway && resp.StatusCode != http.StatusInternalServerError { + if resp.StatusCode == http.StatusOK { + t.Fatal("overflow recorded as success") + } + } + overflowPath := filepath.Join(fixtures, "nuxt", "overflow.yaml") + if _, statErr := os.Stat(overflowPath); !os.IsNotExist(statErr) { + t.Fatalf("partial fixture committed: %v", statErr) + } + + bigUp := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxIn0.abcabcabcabcabcabcab"}`)) + })) + t.Cleanup(bigUp.Close) + secretFix := t.TempDir() + secretProxy, err := NewProxy(ProxyConfig{ + Listen: "127.0.0.1:0", + Upstream: bigUp.URL, + Fixtures: secretFix, + Rules: mustParseRules(t, testRulesYAML()), + }) + if err != nil { + t.Fatal(err) + } + secretSrv := httptest.NewServer(secretProxy.Handler()) + t.Cleanup(secretSrv.Close) + sreq, err := http.NewRequest(http.MethodGet, secretSrv.URL+"/sample", nil) + if err != nil { + t.Fatal(err) + } + sreq.Header.Set(SessionHeader, "secret") + sresp, err := http.DefaultClient.Do(sreq) + if err != nil { + t.Fatal(err) + } + sresp.Body.Close() + if _, statErr := os.Stat(filepath.Join(secretFix, "nuxt", "secret.yaml")); !os.IsNotExist(statErr) { + t.Fatalf("credential fixture was committed: %v", statErr) + } + + reqBad, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil) + if err != nil { + t.Fatal(err) + } + reqBad.Header.Set(SessionHeader, "alice/../bob") + bresp, err := http.DefaultClient.Do(reqBad) + if err != nil { + t.Fatal(err) + } + bresp.Body.Close() + if bresp.StatusCode == http.StatusOK { + t.Fatal("path-separator session must fail") + } + + inside := filepath.Join(fixtures, "vars.yaml") + if _, err := NewProxy(ProxyConfig{ + Listen: "127.0.0.1:0", + Upstream: upstream.URL, + Fixtures: fixtures, + VarsPath: inside, + Rules: mustParseRules(t, testRulesYAML()), + }); err == nil || !strings.Contains(err.Error(), "outside") { + t.Fatalf("vars inside fixtures: %v", err) + } + outside := filepath.Join(t.TempDir(), "vars.yaml") + p, err := NewProxy(ProxyConfig{ + Listen: "127.0.0.1:0", + Upstream: upstream.URL, + Fixtures: fixtures, + VarsPath: outside, + Rules: mustParseRules(t, testRulesYAML()), + }) + if err != nil { + t.Fatal(err) + } + _ = p + st, err := os.Stat(outside) + if err != nil { + t.Fatal(err) + } + if st.Mode().Perm() != 0o600 { + t.Fatalf("vars mode %o", st.Mode().Perm()) + } +} + +func TestProxyDuplicateSessionName(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + _, _ = w.Write([]byte("ok")) + })) + t.Cleanup(upstream.Close) + fixtures := t.TempDir() + if err := os.MkdirAll(filepath.Join(fixtures, "nuxt"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(fixtures, "nuxt", "taken.yaml"), []byte("version: 1\n"), 0o644); err != nil { + t.Fatal(err) + } + proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML()) + srv := httptest.NewServer(proxy.Handler()) + t.Cleanup(srv.Close) + req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil) + if err != nil { + t.Fatal(err) + } + req.Header.Set(SessionHeader, "taken") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode == http.StatusOK { + t.Fatal("duplicate session must fail") + } +} + +func newTestProxy(t *testing.T, upstream, fixtures, rulesYAML string) *Proxy { + t.Helper() + p, err := NewProxy(ProxyConfig{ + Listen: "127.0.0.1:0", + Upstream: upstream, + Fixtures: fixtures, + Rules: mustParseRules(t, rulesYAML), + }) + if err != nil { + t.Fatal(err) + } + return p +} + +func mustParseRules(t *testing.T, raw string) Rules { + t.Helper() + rules, err := ParseRules([]byte(raw)) + if err != nil { + t.Fatal(err) + } + return rules +} + +func testRulesYAML() string { + return "" + + "client: nuxt\n" + + "keep_request_headers:\n" + + " - Cookie\n" + + " - Content-Type\n" + + "keep_response_headers:\n" + + " - Content-Type\n" + + " - Location\n" + + " - Set-Cookie\n" + + "routes:\n" + + " - method: GET\n" + + " path: /sample\n" + + " - method: GET\n" + + " path: /*\n" +} + +func doProxy(t *testing.T, base, session, path, body string) { + t.Helper() + var rdr io.Reader + if body != "" { + rdr = strings.NewReader(body) + } + req, err := http.NewRequest(http.MethodGet, base+path, rdr) + if err != nil { + t.Fatal(err) + } + req.Header.Set(SessionHeader, session) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + b, _ := io.ReadAll(resp.Body) + t.Fatalf("%s %s: %d %s", session, path, resp.StatusCode, b) + } +} diff --git a/tide/rules.go b/tide/rules.go new file mode 100644 index 0000000..4fe4df5 --- /dev/null +++ b/tide/rules.go @@ -0,0 +1,147 @@ +package tide + +import ( + "bytes" + "fmt" + "net/http" + "os" + "path" + "strings" + + "github.com/goccy/go-yaml" +) + +const ( + ClientNuxt = "nuxt" + ClientMCP = "mcp" + + SessionHeader = "X-Parity-Session" +) + +// Rules is the committed, non-secret capture policy for a proxy session. +type Rules struct { + Client string `yaml:"client"` + KeepRequestHeaders []string `yaml:"keep_request_headers"` + KeepResponseHeaders []string `yaml:"keep_response_headers"` + Routes []RouteRule `yaml:"routes"` +} + +// RouteRule matches a method and path pattern and names capture sources. +type RouteRule struct { + Method string `yaml:"method"` + Path string `yaml:"path"` + KeepRequestHeaders []string `yaml:"keep_request_headers,omitempty"` + KeepResponseHeaders []string `yaml:"keep_response_headers,omitempty"` + Capture []CaptureRule `yaml:"capture,omitempty"` +} + +// LoadRules reads a strict YAML rule file, rejecting unknown fields. +func LoadRules(path string) (Rules, error) { + raw, err := os.ReadFile(path) + if err != nil { + return Rules{}, fmt.Errorf("tide: read rules %s: %w", path, err) + } + return ParseRules(raw) +} + +// ParseRules decodes rules YAML with unknown-field rejection. +func ParseRules(raw []byte) (Rules, error) { + var rules Rules + dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField()) + if err := dec.Decode(&rules); err != nil { + return Rules{}, fmt.Errorf("tide: parse rules: %w", err) + } + if err := validateRules(rules); err != nil { + return Rules{}, err + } + return rules, nil +} + +func validateRules(rules Rules) error { + switch rules.Client { + case ClientNuxt, ClientMCP: + default: + return fmt.Errorf("tide: rules client must be %q or %q", ClientNuxt, ClientMCP) + } + for i, route := range rules.Routes { + if strings.TrimSpace(route.Method) == "" { + return fmt.Errorf("tide: rules routes[%d] is missing method", i) + } + if strings.TrimSpace(route.Path) == "" { + return fmt.Errorf("tide: rules routes[%d] is missing path", i) + } + for j, cap := range route.Capture { + if strings.TrimSpace(cap.As) == "" { + return fmt.Errorf("tide: rules routes[%d].capture[%d] is missing as", i, j) + } + } + } + return nil +} + +// Match returns the first route rule for method and path, or nil. +func (r Rules) Match(method, requestPath string) *RouteRule { + for i := range r.Routes { + route := &r.Routes[i] + if !matchMethod(route.Method, method) { + continue + } + if !matchPath(route.Path, requestPath) { + continue + } + return route + } + return nil +} + +func (r Rules) requestHeaders(route *RouteRule) []string { + if route != nil && len(route.KeepRequestHeaders) > 0 { + return route.KeepRequestHeaders + } + return r.KeepRequestHeaders +} + +func (r Rules) responseHeaders(route *RouteRule) []string { + if route != nil && len(route.KeepResponseHeaders) > 0 { + return route.KeepResponseHeaders + } + return r.KeepResponseHeaders +} + +func matchMethod(pattern, method string) bool { + pattern = strings.TrimSpace(pattern) + if pattern == "" || pattern == "*" { + return true + } + return strings.EqualFold(pattern, method) +} + +func matchPath(pattern, requestPath string) bool { + if pattern == requestPath { + return true + } + ok, err := path.Match(pattern, requestPath) + return err == nil && ok +} + +func filterHeaders(h http.Header, keep []string) map[string]string { + if len(keep) == 0 || h == nil { + return nil + } + out := make(map[string]string) + for _, name := range keep { + vals := h.Values(name) + if len(vals) == 0 { + continue + } + if len(vals) == 1 { + out[http.CanonicalHeaderKey(name)] = vals[0] + continue + } + out[http.CanonicalHeaderKey(name)] = strings.Join(vals, "\n") + } + if len(out) == 0 { + return nil + } + return out +} diff --git a/tide/rules_test.go b/tide/rules_test.go new file mode 100644 index 0000000..3562c01 --- /dev/null +++ b/tide/rules_test.go @@ -0,0 +1,60 @@ +package tide + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestRulesLoadAndMatch(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "rules.yaml") + raw := "" + + "client: nuxt\n" + + "keep_request_headers:\n" + + " - Cookie\n" + + "keep_response_headers:\n" + + " - Location\n" + + "routes:\n" + + " - method: GET\n" + + " path: /sample\n" + + " capture:\n" + + " - from: response.json\n" + + " path: $.token\n" + + " as: jwt:alice\n" + + " category: jwt\n" + + " - method: POST\n" + + " path: /items/*\n" + if err := os.WriteFile(path, []byte(raw), 0o644); err != nil { + t.Fatal(err) + } + rules, err := LoadRules(path) + if err != nil { + t.Fatal(err) + } + if rules.Client != ClientNuxt { + t.Fatalf("client: %s", rules.Client) + } + if got := rules.Match("GET", "/sample"); got == nil || len(got.Capture) != 1 || got.Capture[0].As != "jwt:alice" { + t.Fatalf("GET /sample match: %+v", got) + } + if got := rules.Match("POST", "/items/9"); got == nil { + t.Fatal("POST /items/9 should match wildcard") + } + if got := rules.Match("GET", "/other"); got != nil { + t.Fatalf("unexpected match: %+v", got) + } +} + +func TestRulesRejectUnknownAndInvalid(t *testing.T) { + if _, err := ParseRules([]byte("client: nuxt\nextra: true\n")); err == nil { + t.Fatal("unknown field must fail") + } + if _, err := ParseRules([]byte("client: browser\n")); err == nil || !strings.Contains(err.Error(), "client") { + t.Fatalf("invalid client: %v", err) + } + if _, err := ParseRules([]byte("client: mcp\nroutes:\n - method: GET\n path: /x\n capture:\n - path: $.a\n")); err == nil || !strings.Contains(err.Error(), "as") { + t.Fatalf("capture missing as: %v", err) + } +}