diff --git a/internal/config/concurrency_limits_internal_test.go b/internal/config/concurrency_limits_internal_test.go new file mode 100644 index 0000000..c5f4b20 --- /dev/null +++ b/internal/config/concurrency_limits_internal_test.go @@ -0,0 +1,162 @@ +package config + +import ( + "runtime" + "testing" +) + +// The variables that set the two concurrency limits. +const ( + testMaxConcurrentProcessingVar = "PIXA_MAX_CONCURRENT_PROCESSING" + testUpstreamConnectionsVar = "PIXA_UPSTREAM_CONNECTIONS" +) + +// TestOmittedConcurrencyLimitsUseDefaults checks that an omitted +// max_concurrent_processing is the number of CPUs Go uses and an omitted +// upstream_connections is 64. +func TestOmittedConcurrencyLimitsUseDefaults(t *testing.T) { + t.Parallel() + + c, err := configFromYAML(t, signingKeyLine) + if err != nil { + t.Fatalf("minimal config should be valid, got error: %v", err) + } + + if c.MaxConcurrentProcessing != runtime.GOMAXPROCS(0) { + t.Errorf("MaxConcurrentProcessing = %d, want %d, one per CPU", + c.MaxConcurrentProcessing, runtime.GOMAXPROCS(0)) + } + + if c.UpstreamConnections != 64 { + t.Errorf("UpstreamConnections = %d, want 64", c.UpstreamConnections) + } +} + +// TestExplicitConcurrencyLimitsAreUsed checks that valid values for the +// two limits are used as given. +func TestExplicitConcurrencyLimitsAreUsed(t *testing.T) { + t.Parallel() + + c, err := configFromYAML(t, signingKeyLine+ + "max_concurrent_processing: 3\nupstream_connections: 10\n") + if err != nil { + t.Fatalf("valid config should load, got error: %v", err) + } + + if c.MaxConcurrentProcessing != 3 { + t.Errorf("MaxConcurrentProcessing = %d, want 3", c.MaxConcurrentProcessing) + } + + if c.UpstreamConnections != 10 { + t.Errorf("UpstreamConnections = %d, want 10", c.UpstreamConnections) + } +} + +// TestInvalidConcurrencyLimitAbortsStartup checks that a limit that is +// not a whole number of at least 1, or is null, aborts startup naming the +// key and the value, and the variable too where the value could have come +// from it. +func TestInvalidConcurrencyLimitAbortsStartup(t *testing.T) { + t.Parallel() + + processing := keyMaxConcurrentProcessing + connections := keyUpstreamConnections + + runAbortCases(t, []abortCase{ + { + name: "max_concurrent_processing zero", + yaml: signingKeyLine + processing + ": 0\n", + wantErrSubstrings: []string{ + processing, testMaxConcurrentProcessingVar, "value 0", + }, + }, + { + name: "max_concurrent_processing negative", + yaml: signingKeyLine + processing + ": -2\n", + wantErrSubstrings: []string{ + processing, testMaxConcurrentProcessingVar, "value -2", + }, + }, + { + name: "max_concurrent_processing not a number", + yaml: signingKeyLine + processing + ": lots\n", + wantErrSubstrings: []string{ + processing, testMaxConcurrentProcessingVar, "lots", + }, + }, + { + name: "max_concurrent_processing fractional", + yaml: signingKeyLine + processing + ": 1.5\n", + wantErrSubstrings: []string{processing, "1.5"}, + }, + { + name: "max_concurrent_processing null", + yaml: signingKeyLine + processing + ": null\n", + wantErrSubstrings: []string{processing, nullValueText}, + }, + { + name: "upstream_connections zero", + yaml: signingKeyLine + connections + ": 0\n", + wantErrSubstrings: []string{ + connections, testUpstreamConnectionsVar, "value 0", + }, + }, + { + name: "upstream_connections negative", + yaml: signingKeyLine + connections + ": -5\n", + wantErrSubstrings: []string{ + connections, testUpstreamConnectionsVar, "value -5", + }, + }, + { + name: "upstream_connections not a number", + yaml: signingKeyLine + connections + ": many\n", + wantErrSubstrings: []string{ + connections, testUpstreamConnectionsVar, "many", + }, + }, + { + name: "upstream_connections null", + yaml: signingKeyLine + connections + ": null\n", + wantErrSubstrings: []string{connections, nullValueText}, + }, + }) +} + +// TestConcurrencyLimitsFromEnvironment checks that the two variables set +// the limits over the config file, and that an invalid value in either +// aborts startup naming the variable and the value. +func TestConcurrencyLimitsFromEnvironment(t *testing.T) { + t.Setenv(testMaxConcurrentProcessingVar, "3") + t.Setenv(testUpstreamConnectionsVar, "10") + + c, err := configFromYAML(t, signingKeyLine+ + "max_concurrent_processing: 5\nupstream_connections: 50\n") + if err != nil { + t.Fatalf("limits from the environment should load: %v", err) + } + + if c.MaxConcurrentProcessing != 3 || c.UpstreamConnections != 10 { + t.Errorf("limits = %d and %d, want 3 and 10 from the environment", + c.MaxConcurrentProcessing, c.UpstreamConnections) + } + + cases := []struct { + variable string + value string + }{ + {testMaxConcurrentProcessingVar, "lots"}, + {testMaxConcurrentProcessingVar, "0"}, + {testUpstreamConnectionsVar, "-1"}, + {testUpstreamConnectionsVar, "ten"}, + } + + for _, tc := range cases { + t.Run(tc.variable+"="+tc.value, func(t *testing.T) { + t.Setenv(tc.variable, tc.value) + + _, err := configFromYAML(t, signingKeyLine) + wantStartupError(t, err, tc.variable, tc.value) + }) + } +} diff --git a/internal/config/env_internal_test.go b/internal/config/env_internal_test.go index cdc8790..e61ffbf 100644 --- a/internal/config/env_internal_test.go +++ b/internal/config/env_internal_test.go @@ -67,6 +67,8 @@ func TestEnvironmentSetsEveryKey(t *testing.T) { t.Setenv("PIXA_ALLOWLIST_HOSTS", "s3.sneak.cloud,.example.com") t.Setenv("PIXA_ALLOW_HTTP", "true") t.Setenv("PIXA_UPSTREAM_CONNECTIONS_PER_HOST", "5") + t.Setenv("PIXA_UPSTREAM_CONNECTIONS", "10") + t.Setenv("PIXA_MAX_CONCURRENT_PROCESSING", "3") t.Setenv("PIXA_CACHE_MAX_BYTES", "1024") t.Setenv("PIXA_BLOCKED_NETWORKS", "203.0.113.0/24") t.Setenv("PIXA_TRUSTED_PROXIES", "192.0.2.0/24") @@ -93,6 +95,8 @@ func TestEnvironmentSetsEveryKey(t *testing.T) { AllowlistHosts: []string{testHostS3, ".example.com"}, AllowHTTP: true, UpstreamConnectionsPerHost: 5, + UpstreamConnections: 10, + MaxConcurrentProcessing: 3, CacheMaxBytes: 1024, cacheMaxBytesExplicit: true, BlockedNetworks: []netip.Prefix{netip.MustParsePrefix("203.0.113.0/24")}, diff --git a/internal/handlers/server_busy_internal_test.go b/internal/handlers/server_busy_internal_test.go new file mode 100644 index 0000000..2f3471c --- /dev/null +++ b/internal/handlers/server_busy_internal_test.go @@ -0,0 +1,49 @@ +package handlers + +import ( + "fmt" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "sneak.berlin/go/pixa/internal/httpfetcher" + "sneak.berlin/go/pixa/internal/imageprocessor" + "sneak.berlin/go/pixa/internal/imgcache" +) + +// TestServerBusyAnswers503 checks that both image routes answer 503 with +// a clear error when the image service gives up waiting for a free upstream +// connection or processing slot, wrapped as the service wraps them. +func TestServerBusyAnswers503(t *testing.T) { + t.Parallel() + + h := &Handlers{log: slog.New(slog.DiscardHandler)} + req := &imgcache.ImageRequest{SourceHost: "img.example.com", SourcePath: "/a.jpg"} + + for _, err := range []error{ + fmt.Errorf("upstream fetch failed: %w", httpfetcher.ErrTooManyConnections), + fmt.Errorf("image processing failed: %w", imageprocessor.ErrTooManyImages), + } { + plain := httptest.NewRecorder() + h.respondImageError(plain, req, err) + + encrypted := httptest.NewRecorder() + h.handleImageError(encrypted, err) + + for route, rec := range map[string]*httptest.ResponseRecorder{ + "/v1/image/": plain, "/v1/e/": encrypted, + } { + if rec.Code != http.StatusServiceUnavailable { + t.Errorf("%s for %v: status = %d, want %d", + route, err, rec.Code, http.StatusServiceUnavailable) + } + + if !strings.Contains(rec.Body.String(), "server busy, try again later") { + t.Errorf("%s for %v: body = %q, want the server busy error", + route, err, rec.Body.String()) + } + } + } +} diff --git a/internal/httpfetcher/max_connections_internal_test.go b/internal/httpfetcher/max_connections_internal_test.go new file mode 100644 index 0000000..235024b --- /dev/null +++ b/internal/httpfetcher/max_connections_internal_test.go @@ -0,0 +1,128 @@ +package httpfetcher + +import ( + "errors" + "net" + "strconv" + "testing" + "time" +) + +// imageURLOnPort is the fake upstream's image route on testPublicHost at +// port. Each port is a different host to the per-host limit, while the test +// dialer sends every port to the one test server. +func imageURLOnPort(port int) string { + return "http://" + net.JoinHostPort(testPublicHost, strconv.Itoa(port)) + + "/image" +} + +func TestDefaultConfigMaxConnections(t *testing.T) { + t.Parallel() + + if got := DefaultConfig().MaxConnections; got != DefaultMaxConnections { + t.Errorf("MaxConnections = %d, want %d", got, DefaultMaxConnections) + } +} + +// TestFetchLimitsConnectionsToAllHostsTogether checks that MaxConnections +// counts the fetches to every host together, apart from the per-host +// limit: with MaxConnections at 2 and two responses open from two hosts, a +// fetch from a third host, which has nothing open, waits the whole wait +// timeout and fails with ErrTooManyConnections. Closing one response lets +// it through. +func TestFetchLimitsConnectionsToAllHostsTogether(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxConnections = 2 + + f, _ := newServerFetcher(t, srv, cfg) + f.connectionWaitTimeout = 100 * time.Millisecond + + first, err := f.Fetch(testContext(t), imageURLOnPort(81)) + if err != nil { + t.Fatalf("first Fetch() error = %v", err) + } + + second, err := f.Fetch(testContext(t), imageURLOnPort(82)) + if err != nil { + t.Fatalf("second Fetch() error = %v", err) + } + + defer func() { _ = second.Content.Close() }() + + start := time.Now() + + _, err = f.Fetch(testContext(t), imageURLOnPort(83)) + if !errors.Is(err, ErrTooManyConnections) { + t.Fatalf("third Fetch() error = %v, want ErrTooManyConnections", err) + } + + if waited := time.Since(start); waited < f.connectionWaitTimeout { + t.Errorf("third Fetch() failed after %v, before waiting %v", + waited, f.connectionWaitTimeout) + } + + if held := semLen(f, testPublicHost+":83"); held != 0 { + t.Errorf("the refused fetch kept its host's slot: %d held", held) + } + + err = first.Content.Close() + if err != nil { + t.Fatalf("close first body: %v", err) + } + + third, err := f.Fetch(testContext(t), imageURLOnPort(83)) + if err != nil { + t.Fatalf("Fetch() after a response was closed: error = %v", err) + } + + _ = third.Content.Close() +} + +// TestFetchReleasesConnectionOnError checks that a fetch that fails after +// taking its connection gives it back: with MaxConnections at 1, the slot +// must be free after the failure and the next fetch must succeed. +func TestFetchReleasesConnectionOnError(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + url string + want error + }{ + {"upstream answers 500", upstreamURL("/status/500"), ErrUpstreamError}, + {"upstream sends HTML", upstreamURL("/html"), ErrInvalidContentType}, + // 198.51.100.7 (TEST-NET-2) passes the SSRF checks, and the test + // dialer refuses every host but testPublicHost. + {"connecting fails", "http://198.51.100.7/image", errUnexpectedDial}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxConnections = 1 + + f, _ := newServerFetcher(t, srv, cfg) + f.connectionWaitTimeout = 100 * time.Millisecond + + _, err := f.Fetch(testContext(t), tc.url) + if !errors.Is(err, tc.want) { + t.Fatalf("Fetch() error = %v, want %v", err, tc.want) + } + + if held := len(f.allHostsSemaphore); held != 0 { + t.Fatalf("connection still held after the error: %d held", held) + } + + res := fetchImage(t, f, "/image") + _ = res.Content.Close() + }) + } +} diff --git a/internal/imageprocessor/max_concurrent_processing_internal_test.go b/internal/imageprocessor/max_concurrent_processing_internal_test.go new file mode 100644 index 0000000..dc826c4 --- /dev/null +++ b/internal/imageprocessor/max_concurrent_processing_internal_test.go @@ -0,0 +1,302 @@ +package imageprocessor + +import ( + "bytes" + "context" + "errors" + "io" + "runtime" + "strings" + "sync" + "testing" + "testing/iotest" + "time" +) + +// errTestReadFailed is the error the unreadable test input returns. +var errTestReadFailed = errors.New("test input cannot be read") + +// readingCounter counts the Process calls reading their input at the same +// time and remembers the most there ever were. +type readingCounter struct { + mu sync.Mutex + reading int + most int +} + +func (c *readingCounter) start() { + c.mu.Lock() + defer c.mu.Unlock() + + c.reading++ + c.most = max(c.most, c.reading) +} + +func (c *readingCounter) stop() { + c.mu.Lock() + defer c.mu.Unlock() + + c.reading-- +} + +func (c *readingCounter) mostReading() int { + c.mu.Lock() + defer c.mu.Unlock() + + return c.most +} + +// gatedReader is a Process input. Its first Read counts the call in, +// reports it on entered and blocks until gate is closed; it counts the call +// out when it returns io.EOF. Process reads its input only while it holds a +// processing slot, so the count never goes above MaxConcurrentProcessing. +type gatedReader struct { + data *bytes.Reader + gate <-chan struct{} + entered chan<- struct{} + counter *readingCounter + started bool +} + +func (r *gatedReader) Read(p []byte) (int, error) { + if !r.started { + r.started = true + r.counter.start() + + r.entered <- struct{}{} + + <-r.gate + } + + n, err := r.data.Read(p) + if errors.Is(err, io.EOF) { + r.counter.stop() + } + + return n, err +} + +// smallJPEGRequest asks for a 5x5 JPEG. +func smallJPEGRequest() *Request { + return &Request{ + Size: Size{Width: 5, Height: 5}, + Format: FormatJPEG, + Quality: 85, + FitMode: FitCover, + } +} + +// processInBackground runs Process on reader in a new goroutine and sends +// its error on results. +func processInBackground( + proc *ImageProcessor, reader *gatedReader, results chan<- error, +) { + go func() { + result, err := proc.Process(context.Background(), reader, smallJPEGRequest()) + if err == nil { + _ = result.Content.Close() + } + + results <- err + }() +} + +// waitForEntries fails the test unless count Process calls report on +// entered within a few seconds. +func waitForEntries(t *testing.T, entered <-chan struct{}, count int) { + t.Helper() + + for range count { + select { + case <-entered: + case <-time.After(5 * time.Second): + t.Fatal("Process calls did not start reading their input") + } + } +} + +func TestNewDefaultsMaxConcurrentProcessingToCPUs(t *testing.T) { + t.Parallel() + + for _, limit := range []int{0, -1} { + proc := New(Params{MaxConcurrentProcessing: limit}) + if got := cap(proc.processingSemaphore); got != runtime.GOMAXPROCS(0) { + t.Errorf("MaxConcurrentProcessing %d: %d slots, want %d, one per CPU", + limit, got, runtime.GOMAXPROCS(0)) + } + } + + proc := New(Params{MaxConcurrentProcessing: 3}) + if got := cap(proc.processingSemaphore); got != 3 { + t.Errorf("MaxConcurrentProcessing 3: %d slots, want 3", got) + } +} + +// TestProcessNeverExceedsMaxConcurrentProcessing starts more Process calls +// than MaxConcurrentProcessing allows and holds the first ones inside +// Process until the test lets them go. No more than the limit may be +// working at once, and the calls held back must wait for a slot and then +// succeed. +func TestProcessNeverExceedsMaxConcurrentProcessing(t *testing.T) { + t.Parallel() + + const ( + limit = 2 + calls = 6 + ) + + proc := New(Params{MaxConcurrentProcessing: limit}) + input := createTestJPEG(t, 50, 50) + + counter := &readingCounter{} + gate := make(chan struct{}) + entered := make(chan struct{}, calls) + results := make(chan error, calls) + + openGate := sync.OnceFunc(func() { close(gate) }) + t.Cleanup(openGate) + + for range calls { + processInBackground(proc, &gatedReader{ + data: bytes.NewReader(input), gate: gate, entered: entered, + counter: counter, + }, results) + } + + waitForEntries(t, entered, limit) + + // A call beyond the limit would start reading its input now. + select { + case <-entered: + t.Fatalf("a Process call started while %d were already working", limit) + case <-time.After(100 * time.Millisecond): + } + + openGate() + + for range calls { + err := <-results + if err != nil { + t.Errorf("Process() error = %v, want nil once a slot is free", err) + } + } + + if most := counter.mostReading(); most > limit { + t.Errorf("%d Process calls worked at once, want at most %d", most, limit) + } +} + +// TestProcessWaitsThenFailsWhenNoSlotFrees holds the only slot and checks +// that another call waits the whole wait timeout, then fails with +// ErrTooManyImages instead of processing anyway. +func TestProcessWaitsThenFailsWhenNoSlotFrees(t *testing.T) { + t.Parallel() + + proc := New(Params{MaxConcurrentProcessing: 1}) + proc.processingWaitTimeout = 100 * time.Millisecond + + input := createTestJPEG(t, 10, 10) + + gate := make(chan struct{}) + entered := make(chan struct{}, 1) + held := make(chan error, 1) + + openGate := sync.OnceFunc(func() { close(gate) }) + t.Cleanup(openGate) + + processInBackground(proc, &gatedReader{ + data: bytes.NewReader(input), gate: gate, entered: entered, + counter: &readingCounter{}, + }, held) + waitForEntries(t, entered, 1) + + start := time.Now() + + _, err := proc.Process(context.Background(), bytes.NewReader(input), + smallJPEGRequest()) + if !errors.Is(err, ErrTooManyImages) { + t.Fatalf("Process() error = %v, want ErrTooManyImages", err) + } + + if waited := time.Since(start); waited < proc.processingWaitTimeout { + t.Errorf("Process() failed after %v, before waiting %v", + waited, proc.processingWaitTimeout) + } + + openGate() + + err = <-held + if err != nil { + t.Errorf("Process() holding the slot: error = %v, want nil", err) + } +} + +// TestProcessReleasesSlotOnError checks that Process gives its slot back +// when it fails, whether it fails early or late: with one slot, the slot +// must be free after the failure and the next call must succeed. +func TestProcessReleasesSlotOnError(t *testing.T) { + t.Parallel() + + valid := createTestJPEG(t, 10, 10) + + unsupported := smallJPEGRequest() + unsupported.Format = "bmp" + + cases := []struct { + name string + input io.Reader + req *Request + // want is the error Process must return; nil means any error. + want error + }{ + { + name: "input cannot be read", + input: iotest.ErrReader(errTestReadFailed), + req: smallJPEGRequest(), + want: errTestReadFailed, + }, + { + name: "input over the byte limit", + input: bytes.NewReader(createTestJPEG(t, 800, 600)), + req: smallJPEGRequest(), + want: ErrInputDataTooLarge, + }, + { + name: "input not an image", + input: strings.NewReader("not an image"), + req: smallJPEGRequest(), + }, + { + name: "output format not supported", + input: bytes.NewReader(valid), + req: unsupported, + want: ErrUnsupportedOutputFormat, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + proc := New(Params{MaxInputBytes: 4096, MaxConcurrentProcessing: 1}) + proc.processingWaitTimeout = 100 * time.Millisecond + + _, err := proc.Process(context.Background(), tc.input, tc.req) + if err == nil || (tc.want != nil && !errors.Is(err, tc.want)) { + t.Fatalf("Process() error = %v, want %v", err, tc.want) + } + + if held := len(proc.processingSemaphore); held != 0 { + t.Fatalf("slot still held after the error: %d held", held) + } + + result, err := proc.Process(context.Background(), bytes.NewReader(valid), + smallJPEGRequest()) + if err != nil { + t.Fatalf("Process() after the error = %v, want nil", err) + } + + _ = result.Content.Close() + }) + } +}