package tide import ( "bytes" "crypto/tls" "crypto/x509" "io" "net/http" "net/http/httptest" "net/url" "os" "path/filepath" "strings" "testing" "time" ) type upstreamProxyEnv struct { cfg UpstreamProxyConfig caPEM []byte proxy *UpstreamProxy server *httptest.Server client *http.Client } func newUpstreamProxyEnv(t *testing.T, mode, script string) *upstreamProxyEnv { t.Helper() root := t.TempDir() cfg := UpstreamProxyConfig{ Listen: DefaultUpstreamProxyListen, CADir: filepath.Join(root, "ca"), Out: filepath.Join(root, "fixtures", "routes", "POST_things__ok.upstream.yaml"), Mode: mode, VarsPath: filepath.Join(root, "private", "vars.yaml"), } if script != "" { cfg.Script = filepath.Join(root, "script.yaml") if err := os.WriteFile(cfg.Script, []byte(script), 0o644); err != nil { t.Fatal(err) } } store, err := OpenStore(cfg.VarsPath) if err != nil { t.Fatal(err) } store.Set("secret:example-token", "example-token-value") if err := store.Save(); err != nil { t.Fatal(err) } proxy, err := NewUpstreamProxy(cfg) if err != nil { t.Fatalf("NewUpstreamProxy: %v", err) } srv := httptest.NewServer(proxy) t.Cleanup(srv.Close) caPEM, err := os.ReadFile(filepath.Join(cfg.CADir, parityCAFile)) if err != nil { t.Fatal(err) } pool := x509.NewCertPool() if !pool.AppendCertsFromPEM(caPEM) { t.Fatal("parity CA did not parse") } proxyURL, _ := url.Parse(srv.URL) client := &http.Client{ Timeout: 10 * time.Second, Transport: &http.Transport{ Proxy: http.ProxyURL(proxyURL), TLSClientConfig: tlsConfigWithRoots(pool), }, } return &upstreamProxyEnv{cfg: cfg, caPEM: caPEM, proxy: proxy, server: srv, client: client} } const exampleScript = `responses: - method: POST host: api.example.test path: /v1/things response: status: 201 headers: Content-Type: application/json X-Ratelimit-Remaining: "59" body: '{"id":7,"owner":"example-token-value"}' ` func TestUpstreamProxyScriptMode(t *testing.T) { env := newUpstreamProxyEnv(t, "script", exampleScript) req, _ := http.NewRequest(http.MethodPost, "https://api.example.test/v1/things?lang=en", strings.NewReader(`{"name":"widget"}`)) req.Header.Set("Authorization", "Bearer example-token-value") req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", "example-client/1.0") resp, err := env.client.Do(req) if err != nil { t.Fatalf("request through proxy: %v", err) } body, _ := io.ReadAll(resp.Body) _ = resp.Body.Close() if resp.StatusCode != http.StatusCreated || string(body) != `{"id":7,"owner":"example-token-value"}` { t.Fatalf("scripted response = %d %s", resp.StatusCode, body) } if resp.Header.Get("X-Ratelimit-Remaining") != "59" { t.Fatalf("scripted header missing: %v", resp.Header) } if err := env.proxy.Flush(); err != nil { t.Fatalf("Flush: %v", err) } raw, err := os.ReadFile(env.cfg.Out) if err != nil { t.Fatal(err) } if bytes.Contains(raw, []byte("example-token-value")) { t.Fatalf("sidecar leaks the credential:\n%s", raw) } st, _ := os.Stat(env.cfg.Out) if st.Mode().Perm() != 0o644 { t.Fatalf("sidecar mode = %v", st.Mode().Perm()) } s, err := LoadUpstream(env.cfg.Out) if err != nil { t.Fatalf("LoadUpstream: %v\n%s", err, raw) } if len(s.Exchanges) != 1 { t.Fatalf("exchanges = %d", len(s.Exchanges)) } ex := s.Exchanges[0] if got := ex.Request.Headers["Authorization"]; got != "Bearer {{secret:example-token}}" { t.Fatalf("Authorization = %q", got) } if ex.Request.URL != "https://api.example.test/v1/things?lang=en" || ex.Request.Method != "POST" { t.Fatalf("request = %+v", ex.Request) } if ex.Request.Body != `{"name":"widget"}` || ex.Request.Headers["User-Agent"] != "example-client/1.0" { t.Fatalf("request body/headers = %+v", ex.Request) } if ex.Response.Body != `{"id":7,"owner":"{{secret:example-token}}"}` { t.Fatalf("response body = %q", ex.Response.Body) } // The recorded sidecar replays offline. store, _ := OpenStore(env.cfg.VarsPath) fake := NewUpstreamFake(s, store) again, _ := http.NewRequest(http.MethodPost, "https://api.example.test/v1/things?lang=en", strings.NewReader(`{"name":"widget"}`)) again.Header = req.Header.Clone() got, err := fake.RoundTrip(again) if err != nil { t.Fatalf("replay: %v", err) } replayed, _ := io.ReadAll(got.Body) if string(replayed) != `{"id":7,"owner":"example-token-value"}` { t.Fatalf("replayed body = %s", replayed) } if err := fake.Verify(); err != nil { t.Fatal(err) } } func TestUpstreamProxyUnscriptedRequestFailsFlush(t *testing.T) { env := newUpstreamProxyEnv(t, "script", exampleScript) resp, err := env.client.Get("https://api.example.test/v1/unknown") if err != nil { t.Fatal(err) } _ = resp.Body.Close() if resp.StatusCode != upstreamProxyNoMatch { t.Fatalf("status = %d, want %d", resp.StatusCode, upstreamProxyNoMatch) } if err := env.proxy.Flush(); err == nil || !strings.Contains(err.Error(), "no scripted response") { t.Fatalf("Flush = %v, want the unscripted request reported", err) } if _, err := os.Stat(env.cfg.Out); !os.IsNotExist(err) { t.Fatalf("sidecar written despite failures: %v", err) } // Plain requests are refused: the proxy only tunnels. plain, err := http.Get(env.server.URL + "/x") if err != nil { t.Fatal(err) } _ = plain.Body.Close() if plain.StatusCode != http.StatusMethodNotAllowed { t.Fatalf("plain request status = %d", plain.StatusCode) } } func TestUpstreamProxyMultipartAndForwardGuard(t *testing.T) { script := `responses: - method: POST host: files.example.test path: /upload response: status: 202 ` env := newUpstreamProxyEnv(t, "script", script) var buf bytes.Buffer body := "--b1\r\nContent-Disposition: form-data; name=\"title\"\r\n\r\nHello\r\n" + "--b1\r\nContent-Disposition: form-data; name=\"file\"; filename=\"a.png\"\r\nContent-Type: image/png\r\n\r\nPNGDATA\r\n--b1--\r\n" buf.WriteString(body) req, _ := http.NewRequest(http.MethodPost, "https://files.example.test/upload", &buf) req.Header.Set("Content-Type", "multipart/form-data; boundary=b1") resp, err := env.client.Do(req) if err != nil { t.Fatal(err) } _ = resp.Body.Close() if err := env.proxy.Flush(); err != nil { t.Fatal(err) } s, err := LoadUpstream(env.cfg.Out) if err != nil { t.Fatal(err) } parts := s.Exchanges[0].Request.Parts if len(parts) != 2 || parts[0].Value != "Hello" || parts[1].Filename != "a.png" || parts[1].ContentType != "image/png" || len(parts[1].SHA256) != 64 { t.Fatalf("parts = %+v", parts) } // The fake accepts the same parts under another boundary. fake := NewUpstreamFake(s, nil) other := strings.ReplaceAll(body, "b1", "zz9") again, _ := http.NewRequest(http.MethodPost, "https://files.example.test/upload", strings.NewReader(other)) again.Header.Set("Content-Type", "multipart/form-data; boundary=zz9") again.Header.Set("User-Agent", "Go-http-client/1.1") if _, err := fake.RoundTrip(again); err != nil { t.Fatalf("replay multipart: %v", err) } if err := fake.Verify(); err != nil { t.Fatal(err) } fake = NewUpstreamFake(s, nil) changed, _ := http.NewRequest(http.MethodPost, "https://files.example.test/upload", strings.NewReader(strings.Replace(other, "PNGDATA", "PNGDATX", 1))) changed.Header.Set("Content-Type", "multipart/form-data; boundary=zz9") changed.Header.Set("User-Agent", "Go-http-client/1.1") if _, err := fake.RoundTrip(changed); err == nil || !strings.Contains(err.Error(), "sha256") { t.Fatalf("changed file bytes: %v", err) } // Forward mode sends through a PublicOnlyMode client: a loopback target // is refused at dial and the recording fails. target := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { t.Error("forward mode must not reach a loopback target") })) t.Cleanup(target.Close) fwd := newUpstreamProxyEnv(t, "forward", "") resp, err = fwd.client.Get(target.URL + "/secret") if err != nil { t.Fatal(err) } _ = resp.Body.Close() if resp.StatusCode != http.StatusBadGateway { t.Fatalf("forward to loopback = %d, want 502", resp.StatusCode) } if err := fwd.proxy.Flush(); err == nil || !strings.Contains(err.Error(), "private_ip") { t.Fatalf("Flush = %v, want private_ip", err) } } func TestUpstreamProxyRefusesNonLoopback(t *testing.T) { root := t.TempDir() base := UpstreamProxyConfig{ CADir: filepath.Join(root, "ca"), Out: filepath.Join(root, "fixtures", "x.upstream.yaml"), Script: filepath.Join(root, "script.yaml"), VarsPath: filepath.Join(root, "vars.yaml"), } if err := os.WriteFile(base.Script, []byte("responses: []\n"), 0o644); err != nil { t.Fatal(err) } for _, listen := range []string{"0.0.0.0:8425", "192.0.2.10:8425", ":8425"} { cfg := base cfg.Listen = listen if _, err := NewUpstreamProxy(cfg); err == nil || !strings.Contains(err.Error(), "loopback") { t.Fatalf("listen %s: err = %v, want loopback refusal", listen, err) } } cases := []struct { name string edit func(*UpstreamProxyConfig) want string }{ {"bad mode", func(c *UpstreamProxyConfig) { c.Mode = "replay" }, "script or forward"}, {"no out", func(c *UpstreamProxyConfig) { c.Out = "" }, "requires out"}, {"no ca", func(c *UpstreamProxyConfig) { c.CADir = "" }, "requires ca-dir"}, {"no vars", func(c *UpstreamProxyConfig) { c.VarsPath = "" }, "requires vars"}, {"no script", func(c *UpstreamProxyConfig) { c.Script = "" }, "requires script"}, {"ca inside out dir", func(c *UpstreamProxyConfig) { c.CADir = filepath.Join(root, "fixtures", "ca") }, "ca dir"}, {"vars inside out dir", func(c *UpstreamProxyConfig) { c.VarsPath = filepath.Join(root, "fixtures", "vars.yaml") }, "vars file"}, {"missing script", func(c *UpstreamProxyConfig) { c.Script = filepath.Join(root, "nope.yaml") }, "read upstream script"}, } for _, tc := range cases { cfg := base tc.edit(&cfg) if _, err := NewUpstreamProxy(cfg); err == nil || !strings.Contains(err.Error(), tc.want) { t.Fatalf("%s: err = %v, want %q", tc.name, err, tc.want) } } if _, err := os.Stat(filepath.Join(root, "fixtures", "ca", parityCAKeyFile)); err == nil { t.Fatal("a refused config must not create a CA key inside the fixtures tree") } bad := filepath.Join(root, "bad-script.yaml") _ = os.WriteFile(bad, []byte("responses:\n - method: GET\n host: a.test\n path: /x\n response: {}\n"), 0o644) cfg := base cfg.Script = bad if _, err := NewUpstreamProxy(cfg); err == nil || !strings.Contains(err.Error(), "response.status") { t.Fatalf("incomplete script entry: %v", err) } } func TestEnsureParityCA(t *testing.T) { dir := filepath.Join(t.TempDir(), "ca") certPath, err := EnsureParityCA(dir) if err != nil { t.Fatal(err) } keyPath := filepath.Join(dir, parityCAKeyFile) kst, err := os.Stat(keyPath) if err != nil { t.Fatal(err) } if kst.Mode().Perm() != 0o600 { t.Fatalf("key mode = %v, want 0600", kst.Mode().Perm()) } cst, _ := os.Stat(certPath) if cst.Mode().Perm() != 0o644 { t.Fatalf("cert mode = %v, want 0644", cst.Mode().Perm()) } first, _ := os.ReadFile(certPath) firstKey, _ := os.ReadFile(keyPath) _ = os.Chmod(keyPath, 0o644) again, err := EnsureParityCA(dir) if err != nil || again != certPath { t.Fatalf("second call = %q, %v", again, err) } second, _ := os.ReadFile(certPath) secondKey, _ := os.ReadFile(keyPath) if !bytes.Equal(first, second) || !bytes.Equal(firstKey, secondKey) { t.Fatal("second call must reuse the CA files") } if st, _ := os.Stat(keyPath); st.Mode().Perm() != 0o600 { t.Fatalf("reused key mode = %v, want 0600", st.Mode().Perm()) } cert, _, err := readParityCA(certPath, keyPath) if err != nil || !cert.IsCA { t.Fatalf("CA = %v, %v", cert, err) } if err := os.WriteFile(certPath, []byte("not pem"), 0o644); err != nil { t.Fatal(err) } if _, err := EnsureParityCA(dir); err == nil { t.Fatal("corrupt CA must be reported, not silently replaced") } if _, err := EnsureParityCA(""); err == nil { t.Fatal("empty dir must fail") } } func tlsConfigWithRoots(pool *x509.CertPool) *tls.Config { return &tls.Config{RootCAs: pool, MinVersion: tls.VersionTLS12} }