Files
chorus/internal/platform/http/safehttp_test.go
T

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)
}
}