package safehttp import ( "context" "errors" "fmt" "net" "net/http" "strings" "sync" "testing" "time" ) type fakeResolver struct { mu sync.Mutex addresses map[string][][]net.IPAddr } func (r *fakeResolver) LookupIPAddr(_ context.Context, host string) ([]net.IPAddr, error) { r.mu.Lock() defer r.mu.Unlock() answers := r.addresses[host] if len(answers) == 0 { return nil, fmt.Errorf("no answer") } answer := answers[0] if len(answers) > 1 { r.addresses[host] = answers[1:] } return answer, nil } func ips(values ...string) []net.IPAddr { result := make([]net.IPAddr, 0, len(values)) for _, value := range values { result = append(result, net.IPAddr{IP: net.ParseIP(value)}) } return result } func TestBlockedIPv4AndIPv6(t *testing.T) { blocked := []string{"127.0.0.1", "10.0.0.1", "169.254.169.254", "100.64.0.1", "0.0.0.0", "::1", "fc00::1", "fe80::1", "ff02::1", "::ffff:192.168.1.1"} for _, value := range blocked { t.Run(value, func(t *testing.T) { if err := validateAddresses(ips(value)); !errors.Is(err, ErrBlockedAddress) { t.Fatalf("validateAddresses(%s) error = %v", value, err) } }) } if err := validateAddresses(ips("93.184.216.34", "2606:4700:4700::1111")); err != nil { t.Fatalf("public addresses rejected: %v", err) } } func TestDialUsesValidatedIPAndRejectsDNSRebinding(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "public.test": {ips("93.184.216.34")}, "rebind.test": {ips("93.184.216.34"), ips("127.0.0.1")}, }} var dialed string fakeDial := func(_ context.Context, _, address string) (net.Conn, error) { dialed = address return nil, errors.New("synthetic dial stop") } client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2}) if err != nil { t.Fatal(err) } _, _ = client.Fetch(context.Background(), "http://public.test/data", 10) if dialed != "93.184.216.34:80" { t.Fatalf("dialed address = %q", dialed) } dialed = "" _, err = client.Fetch(context.Background(), "http://rebind.test/data", 10) if !errors.Is(err, ErrBlockedAddress) || dialed != "" { t.Fatalf("rebind result: dialed=%q error=%v", dialed, err) } } func TestFetchPublicMockAndSizeLimit(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "public.test": {ips("93.184.216.34"), ips("93.184.216.34"), ips("93.184.216.34"), ips("93.184.216.34")}, }} fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { client, server := net.Pipe() go func() { defer server.Close() buffer := make([]byte, 4096) _, _ = server.Read(buffer) _, _ = server.Write([]byte("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 11\r\nConnection: close\r\n\r\nhello world")) }() return client, nil } client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 1}) if err != nil { t.Fatal(err) } response, err := client.Fetch(context.Background(), "http://public.test/data", 11) if err != nil || string(response.Body) != "hello world" { t.Fatalf("Fetch() = %#v, %v", response, err) } _, err = client.Fetch(context.Background(), "http://public.test/data", 5) if !errors.Is(err, ErrResponseTooLarge) { t.Fatalf("size limit error = %v", err) } } func TestDoPostUsesSafeTransportAndForwardsHeaders(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "provider.test": {ips("93.184.216.34"), ips("93.184.216.34")}, }} requestText := make(chan string, 1) fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { client, server := net.Pipe() go func() { defer server.Close() buffer := make([]byte, 4096) count, _ := server.Read(buffer) requestText <- string(buffer[:count]) _, _ = server.Write([]byte("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 11\r\nConnection: close\r\n\r\n{\"ok\":true}")) }() return client, nil } client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 0}) if err != nil { t.Fatal(err) } response, err := client.Do(context.Background(), Request{ Method: http.MethodPost, URL: "http://provider.test/v1/chat/completions", Header: http.Header{"Authorization": []string{"Bearer test-key"}, "Content-Type": []string{"application/json"}}, Body: strings.NewReader(`{"model":"mock"}`), MaxBytes: 100, }) if err != nil || response.StatusCode != http.StatusOK { t.Fatalf("Do() = %#v, %v", response, err) } received := <-requestText if !strings.Contains(received, "POST /v1/chat/completions") || !strings.Contains(received, "Authorization: Bearer test-key") { t.Fatalf("unexpected outbound request: %q", received) } } func TestRedirectAndProxyBypassAreRejected(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "public.test": {ips("93.184.216.34"), ips("93.184.216.34")}, "private.test": {ips("10.0.0.8")}, }} fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { client, server := net.Pipe() go func() { defer server.Close() buffer := make([]byte, 4096) _, _ = server.Read(buffer) _, _ = server.Write([]byte("HTTP/1.1 302 Found\r\nLocation: http://private.test/result\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")) }() return client, nil } t.Setenv("HTTP_PROXY", "http://127.0.0.1:9999") client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2}) if err != nil { t.Fatal(err) } _, err = client.Fetch(context.Background(), "http://public.test/start", 100) if !errors.Is(err, ErrBlockedAddress) { t.Fatalf("redirect error = %v", err) } transport := client.httpClient.Transport.(*http.Transport) if transport.Proxy != nil { t.Fatal("safe transport unexpectedly configured a proxy") } } func TestAuthenticatedCrossOriginRedirectIsRejected(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "provider.test": {ips("93.184.216.34"), ips("93.184.216.34")}, "other.test": {ips("93.184.216.35")}, }} fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { client, server := net.Pipe() go func() { defer server.Close() buffer := make([]byte, 4096) _, _ = server.Read(buffer) _, _ = server.Write([]byte("HTTP/1.1 302 Found\r\nLocation: http://other.test/result\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")) }() return client, nil } client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2}) if err != nil { t.Fatal(err) } _, err = client.Do(context.Background(), Request{Method: http.MethodPost, URL: "http://provider.test/request", Header: http.Header{"Authorization": []string{"Bearer secret"}}, Body: strings.NewReader("{}"), MaxBytes: 100}) if !errors.Is(err, ErrCrossOriginAuth) { t.Fatalf("cross-origin redirect error = %v", err) } } func TestURLRestrictionsAndRedaction(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{"public.test": {ips("93.184.216.34")}}} client, err := New(Config{Resolver: resolver, Timeout: time.Second, MaxRedirects: 0}) if err != nil { t.Fatal(err) } for _, target := range []string{"file:///etc/passwd", "http://user:password@public.test/", "http://public.test:22/"} { if _, err := client.Fetch(context.Background(), target, 10); !errors.Is(err, ErrInvalidTarget) { t.Errorf("Fetch(%q) error = %v", target, err) } } redacted := RedactedURL("https://user:secret@Public.Test/path?q=secret") if strings.Contains(redacted, "secret") || redacted != "https://Public.Test" { t.Fatalf("RedactedURL() = %q", redacted) } } func TestCustomAllowedPortStillUsesValidatedDial(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "provider.test": {ips("93.184.216.34"), ips("93.184.216.34")}, }} var dialed string fakeDial := func(_ context.Context, _, address string) (net.Conn, error) { dialed = address return nil, errors.New("synthetic dial stop") } defaultClient, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second}) if err != nil { t.Fatal(err) } if _, err := defaultClient.Fetch(context.Background(), "http://provider.test:8080/data", 10); !errors.Is(err, ErrInvalidTarget) || dialed != "" { t.Fatalf("default port policy: dialed=%q error=%v", dialed, err) } customClient, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, AllowedPorts: []uint16{80, 443, 8080}}) if err != nil { t.Fatal(err) } _, _ = customClient.Fetch(context.Background(), "http://provider.test:8080/data", 10) if dialed != "93.184.216.34:8080" { t.Fatalf("custom port dialed address = %q", dialed) } } func TestMaliciousResultURLAndTimeout(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "slow.test": {ips("93.184.216.34"), ips("93.184.216.34")}, }} client, err := New(Config{ Resolver: resolver, DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { client, _ := net.Pipe() return client, nil }, Timeout: 20 * time.Millisecond, MaxRedirects: 0, }) if err != nil { t.Fatal(err) } if _, err := client.Fetch(context.Background(), "http://169.254.169.254/latest/meta-data", 100); !errors.Is(err, ErrBlockedAddress) { t.Fatalf("malicious result URL error = %v", err) } _, err = client.Fetch(context.Background(), "http://slow.test/result?signature=do-not-leak", 100) if !errors.Is(err, ErrRequestFailed) || strings.Contains(err.Error(), "do-not-leak") { t.Fatalf("timeout error = %v", err) } }