238 lines
8.4 KiB
Go
238 lines
8.4 KiB
Go
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 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)
|
|
}
|
|
}
|