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:
@@ -29,6 +29,7 @@ func toolCommands() []bonfire.Command {
|
|||||||
makePluginCommand(),
|
makePluginCommand(),
|
||||||
addPluginCommand(),
|
addPluginCommand(),
|
||||||
devCommand(),
|
devCommand(),
|
||||||
|
parityProxyCommand(),
|
||||||
parityRecordCommand(),
|
parityRecordCommand(),
|
||||||
parityReplayCommand(),
|
parityReplayCommand(),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,22 @@ import (
|
|||||||
"git.golem15.com/golem15/summercms/tide"
|
"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 {
|
func parityRecordCommand() bonfire.Command {
|
||||||
return bonfire.Command{
|
return bonfire.Command{
|
||||||
Name: "parity:record",
|
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 {
|
func runParityRecord(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
|
||||||
specPath, err := requireFlag(in, "spec", "parity:record")
|
specPath, err := requireFlag(in, "spec", "parity:record")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ import (
|
|||||||
|
|
||||||
func TestParityCommands(t *testing.T) {
|
func TestParityCommands(t *testing.T) {
|
||||||
names := commandNames()
|
names := commandNames()
|
||||||
for _, want := range []string{"parity:record", "parity:replay"} {
|
for _, want := range []string{"parity:record", "parity:replay", "parity:proxy"} {
|
||||||
if !containsName(names, want) {
|
if !containsName(names, want) {
|
||||||
t.Fatalf("missing %s in %v", want, names)
|
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 {
|
if err := runParity("parity:replay", "--fixtures", keyFixture, "--target", keyOrder.URL); err != nil {
|
||||||
t.Fatalf("reordered JSON keys must pass: %v", err)
|
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 {
|
func commandNames() []string {
|
||||||
|
|||||||
@@ -154,8 +154,22 @@ func writeCapture(b *strings.Builder, indent int, rules []CaptureRule) {
|
|||||||
fmt.Fprintf(b, "%scapture:\n", pad)
|
fmt.Fprintf(b, "%scapture:\n", pad)
|
||||||
inner := strings.Repeat(" ", indent+2)
|
inner := strings.Repeat(" ", indent+2)
|
||||||
for _, rule := range rules {
|
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))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
10
tide/flow.go
10
tide/flow.go
@@ -50,10 +50,14 @@ type Response struct {
|
|||||||
SHA256 string `yaml:"sha256,omitempty"`
|
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 {
|
type CaptureRule struct {
|
||||||
Path string `yaml:"path"`
|
From string `yaml:"from,omitempty"`
|
||||||
As string `yaml:"as"`
|
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.
|
// NormalizeRule names a per-step normalizer override.
|
||||||
|
|||||||
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
|
||||||
|
}
|
||||||
|
}
|
||||||
368
tide/proxy_test.go
Normal file
368
tide/proxy_test.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
147
tide/rules.go
Normal file
147
tide/rules.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
60
tide/rules_test.go
Normal file
60
tide/rules_test.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user