diff --git a/README.md b/README.md index 7a3a013..001d863 100644 --- a/README.md +++ b/README.md @@ -235,6 +235,8 @@ variables set by the file's `env:` section are checked the same way. | `PIXA_TRUSTED_PROXIES` | `trusted_proxies` | CIDR ranges of proxies whose `X-Forwarded-For` is believed; default RFC 1918 | | `PIXA_ALLOW_HTTP` | `allow_http` | Allow plain-HTTP upstreams, for testing only; default `false` | | `PIXA_UPSTREAM_CONNECTIONS_PER_HOST` | `upstream_connections_per_host` | Concurrent connections per upstream host; default `20` | +| `PIXA_UPSTREAM_CONNECTIONS` | `upstream_connections` | Concurrent connections to all upstream hosts together; default `64` | +| `PIXA_MAX_CONCURRENT_PROCESSING` | `max_concurrent_processing` | Images processed at once; default the number of CPUs | | `PIXA_UPSTREAM_FETCH_TIMEOUT` | `upstream_fetch_timeout` | Time allowed for one fetch from an upstream host; default `30s` | | `PIXA_UPSTREAM_MAX_RESPONSE_SIZE` | `upstream_max_response_size` | Largest upstream response accepted, in bytes; default 50 MiB | | `PIXA_DOWNSTREAM_TIMEOUT` | `downstream_timeout` | Time allowed for answering one client request; default `60s` | @@ -286,12 +288,24 @@ Key settings in more detail: bytes; default `52428800` (50 MiB). It also limits the image data pixa decodes - `downstream_timeout` — time allowed for answering one client request, as a - duration; default `60s`. The upstream fetch counts toward it, so keep it - longer than `upstream_fetch_timeout` + duration; default `60s`. The upstream fetch counts toward it, and so do the + waits for an upstream connection and for a processing slot (up to 10 seconds + each), so keep it longer than `upstream_fetch_timeout` plus 20 seconds - `signing_key` — HMAC secret for URL signatures - `cache_max_bytes` — disk cache size limit in bytes; `0` disables the disk cache entirely; omitted defaults to 75% of the free space on the filesystem containing `/cache/` (minimum 500 MiB) +- `upstream_connections` — the most connections to upstream hosts at once, all + hosts together, on top of `upstream_connections_per_host`; default `64`. A + fetch holds its connection until its image has been processed. A fetch that + finds all of them in use waits up to 10 seconds for one to free up; if none + does, and `downstream_timeout` has not ended first, the request is answered + 503 with the error `server busy, try again later` +- `max_concurrent_processing` — the most images decoded and encoded at once; + default the number of CPUs pixa can use (`GOMAXPROCS`), which follows a + container's CPU limit. A request that finds all of them in use waits up to 10 + seconds for one to free up; if none does, and `downstream_timeout` has not + ended first, it is answered 503 the same way See `config.example.yml` for all options with defaults. diff --git a/TODO.md b/TODO.md index c252259..8276df8 100644 --- a/TODO.md +++ b/TODO.md @@ -25,11 +25,20 @@ The disk cache is now size-bounded with LRU eviction # Next Step -P1: rate limit global concurrent upstream fetches to prevent resource -exhaustion +P2: security: referer blacklist # Completed Steps +- 2026-09-29 bound concurrent image processing and upstream fetches (closes + #64): `max_concurrent_processing` (default the number of CPUs pixa can use) + limits the images decoded and encoded at once, and `upstream_connections` + (default 64) the connections to all upstream hosts together, on top of + `upstream_connections_per_host`; a fetch holds its connection until its image + has been processed, and a request whose source is cached reads it only once it + has a processing slot; a request that finds either limit reached waits up to + 10 seconds for a free one, then gets 503 `server busy, try again later`; + libvips runs one worker thread per image with its operation cache off; + documented in `README.md` and `config.example.yml`. - 2026-09-29 Dockerfiles install through `script/bootstrap` (closes #95): the `Dockerfile` lint and build stages and `Dockerfile.lint` copy `script/`, `go.mod` and `go.sum`, then run `script/bootstrap` in place of their own @@ -316,7 +325,6 @@ exhaustion # Future Steps - P2: security - - referer blacklist - per-IP rate limiting on the image routes - per-origin rate limiting - P2: HTTP response handling diff --git a/config.example.yml b/config.example.yml index 27c2422..d64d77e 100644 --- a/config.example.yml +++ b/config.example.yml @@ -71,6 +71,19 @@ allow_http: false # Maximum concurrent connections per upstream host (default: 20) upstream_connections_per_host: 20 +# Maximum concurrent connections to all upstream hosts together, on top of +# the per-host limit (default: 64). A fetch holds its connection until its +# image has been processed. A fetch that finds none free waits up to 10 +# seconds for one, and if none frees up the request is answered 503, unless +# downstream_timeout has ended first. +upstream_connections: 64 + +# Maximum number of images decoded and encoded at once (default: the +# number of CPUs pixa can use, which follows a container's CPU limit). A +# request that finds none free waits up to 10 seconds for one, and if none +# frees up it is answered 503, unless downstream_timeout has ended first. +# max_concurrent_processing: 4 + # Time allowed for one fetch from an upstream host (default: 30s) upstream_fetch_timeout: 30s @@ -78,8 +91,10 @@ upstream_fetch_timeout: 30s # (1 GiB) (default: 52428800, 50 MiB) upstream_max_response_size: 52428800 -# Time allowed for answering one client request, the upstream fetch -# included, so keep it longer than upstream_fetch_timeout (default: 60s) +# Time allowed for answering one client request (default: 60s). The +# upstream fetch counts toward it, and so do the waits for an upstream +# connection and for a processing slot (up to 10 seconds each), so keep it +# longer than upstream_fetch_timeout plus 20 seconds. downstream_timeout: 60s # The origin a browser lets read pixa's responses, sent as the CORS 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/config.go b/internal/config/config.go index 5c9504f..b8e041b 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -10,6 +10,7 @@ import ( "net/url" "os" "path/filepath" + "runtime" "sort" "strconv" "strings" @@ -25,6 +26,7 @@ const ( DefaultPort = 8080 DefaultStateDir = "/var/lib/pixa" DefaultUpstreamConnectionsPerHost = 20 + DefaultUpstreamConnections = 64 DefaultAccessControlAllowOrigin = "*" DefaultUpstreamFetchTimeout = 30 * time.Second DefaultUpstreamMaxResponseSize = 50 << 20 // 50 MiB @@ -46,6 +48,8 @@ const ( keyAllowlistHosts = "allowlist_hosts" keyAllowHTTP = "allow_http" keyUpstreamConnectionsPerHost = "upstream_connections_per_host" + keyUpstreamConnections = "upstream_connections" + keyMaxConcurrentProcessing = "max_concurrent_processing" keyCacheMaxBytes = "cache_max_bytes" keyBlockedNetworks = "blocked_networks" keyTrustedProxies = "trusted_proxies" @@ -79,7 +83,7 @@ var ( errNotAValidURL = errors.New("not a valid URL") errPortOutOfRange = errors.New("outside the valid port range") errSizeOutOfRange = errors.New("outside the accepted range") - errTooFewConnections = errors.New("must be at least 1") + errMustBeAtLeastOne = errors.New("must be at least 1") errValueTooShort = errors.New("value too short") errPlaceholderKey = errors.New( "is the placeholder from config.example.yml; " + @@ -126,6 +130,12 @@ type Config struct { AllowHTTP bool // Allow non-TLS upstream (testing only) UpstreamConnectionsPerHost int // Max concurrent connections per upstream host + // UpstreamConnections is the most concurrent connections to all + // upstream hosts together, on top of the per-host limit. + // MaxConcurrentProcessing is the most images processed at once. + UpstreamConnections int + MaxConcurrentProcessing int + // UpstreamFetchTimeout is the time allowed for one fetch from an // upstream host. UpstreamMaxResponseSize is the largest upstream // response accepted, in bytes, and also the image processor's input @@ -226,14 +236,12 @@ func New(_ fx.Lifecycle, params Params) (*Config, error) { // unparseable or invalid is an error: defaults apply only to omitted // keys, never to invalid explicit values. func newFromSmartConfig(sc *smartconfig.Config) (*Config, error) { - if sc != nil { - err := validateKnownKeys(sc) - if err != nil { - return nil, err - } + err := validateKnownKeys(sc) + if err != nil { + return nil, err } - err := validateAllowlistHostsValue(sc) + err = validateAllowlistHostsValue(sc) if err != nil { return nil, err } @@ -271,6 +279,12 @@ func newFromSmartConfig(sc *smartconfig.Config) (*Config, error) { AllowHTTP: loader.boolVal(keyAllowHTTP, false), UpstreamConnectionsPerHost: loader.intVal( keyUpstreamConnectionsPerHost, DefaultUpstreamConnectionsPerHost), + UpstreamConnections: loader.intVal( + keyUpstreamConnections, DefaultUpstreamConnections), + // Decoding and encoding are CPU-bound, so the default is one image + // per CPU Go uses, which follows a container's CPU limit. + MaxConcurrentProcessing: loader.intVal( + keyMaxConcurrentProcessing, runtime.GOMAXPROCS(0)), UpstreamFetchTimeout: loader.durationVal( keyUpstreamFetchTimeout, DefaultUpstreamFetchTimeout), UpstreamMaxResponseSize: loader.int64Val( @@ -321,8 +335,13 @@ func newFromSmartConfig(sc *smartconfig.Config) (*Config, error) { // being silently ignored, and rejects keys that are explicitly set to // null: a null is a SET value, never an omission, so it must not // silently take the default. The env section is permitted because -// smartconfig consumes it for environment variable injection. +// smartconfig consumes it for environment variable injection. A nil sc +// means no config file, which has no keys to check. func validateKnownKeys(sc *smartconfig.Config) error { + if sc == nil { + return nil + } + var unknown, nullKeys []string for key, value := range sc.Data() { @@ -392,7 +411,8 @@ func isKnownConfigKey(key string) bool { switch key { case keyDebug, keyMaintenanceMode, keyPort, keyStateDir, keySentryDSN, keyDBURL, keyMetrics, keySigningKey, keyAllowlistHosts, keyAllowHTTP, - keyUpstreamConnectionsPerHost, keyCacheMaxBytes, keyBlockedNetworks, + keyUpstreamConnectionsPerHost, keyUpstreamConnections, + keyMaxConcurrentProcessing, keyCacheMaxBytes, keyBlockedNetworks, keyTrustedProxies, keyAccessControlAllowOrigin, keyUpstreamFetchTimeout, keyUpstreamMaxResponseSize, keyDownstreamTimeout, "env": return true @@ -419,6 +439,8 @@ func envVarNames() map[string]string { keyAllowlistHosts: "PIXA_ALLOWLIST_HOSTS", keyAllowHTTP: "PIXA_ALLOW_HTTP", keyUpstreamConnectionsPerHost: "PIXA_UPSTREAM_CONNECTIONS_PER_HOST", + keyUpstreamConnections: "PIXA_UPSTREAM_CONNECTIONS", + keyMaxConcurrentProcessing: "PIXA_MAX_CONCURRENT_PROCESSING", keyCacheMaxBytes: "PIXA_CACHE_MAX_BYTES", keyBlockedNetworks: "PIXA_BLOCKED_NETWORKS", keyTrustedProxies: "PIXA_TRUSTED_PROXIES", @@ -562,10 +584,9 @@ func (c *Config) validate() error { settingName(keyPort), c.Port, errPortOutOfRange, maxPort) } - if c.UpstreamConnectionsPerHost < 1 { - return fmt.Errorf("%s: value %d %w", - settingName(keyUpstreamConnectionsPerHost), - c.UpstreamConnectionsPerHost, errTooFewConnections) + err = c.validateConcurrencyLimits() + if err != nil { + return err } if c.StateDir == "" { @@ -684,6 +705,30 @@ func (c *Config) validateAccessControlAllowOrigin() error { return nil } +// validateConcurrencyLimits checks that the two upstream connection limits +// and the image processing limit are at least 1. +func (c *Config) validateConcurrencyLimits() error { + if c.UpstreamConnectionsPerHost < 1 { + return fmt.Errorf("%s: value %d %w", + settingName(keyUpstreamConnectionsPerHost), + c.UpstreamConnectionsPerHost, errMustBeAtLeastOne) + } + + if c.UpstreamConnections < 1 { + return fmt.Errorf("%s: value %d %w", + settingName(keyUpstreamConnections), + c.UpstreamConnections, errMustBeAtLeastOne) + } + + if c.MaxConcurrentProcessing < 1 { + return fmt.Errorf("%s: value %d %w", + settingName(keyMaxConcurrentProcessing), + c.MaxConcurrentProcessing, errMustBeAtLeastOne) + } + + return nil +} + // validateAllowlistHost checks that an allowlist_hosts entry is a bare // hostname, optionally with a leading dot for suffix matching. URLs, // paths, and whitespace indicate a misconfigured entry. An entry with 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/handlers.go b/internal/handlers/handlers.go index 3da580d..b350595 100644 --- a/internal/handlers/handlers.go +++ b/internal/handlers/handlers.go @@ -113,15 +113,17 @@ func (s *Handlers) initImageService() error { fetcherCfg.MaxConnectionsPerHost = s.config.UpstreamConnectionsPerHost } + fetcherCfg.MaxConnections = s.config.UpstreamConnections fetcherCfg.BlockedNetworks = s.config.BlockedNetworks // Create the service svc, err := imgcache.NewService(&imgcache.ServiceConfig{ - Cache: cache, - FetcherConfig: fetcherCfg, - SigningKey: s.config.SigningKey, - Allowlist: s.config.AllowlistHosts, - Logger: s.log, + Cache: cache, + FetcherConfig: fetcherCfg, + SigningKey: s.config.SigningKey, + Allowlist: s.config.AllowlistHosts, + MaxConcurrentProcessing: s.config.MaxConcurrentProcessing, + Logger: s.log, }) if err != nil { return err diff --git a/internal/handlers/image.go b/internal/handlers/image.go index 561ca1d..3c94bc8 100644 --- a/internal/handlers/image.go +++ b/internal/handlers/image.go @@ -12,6 +12,7 @@ import ( "github.com/go-chi/chi/v5" "sneak.berlin/go/pixa/internal/encurl" "sneak.berlin/go/pixa/internal/httpfetcher" + "sneak.berlin/go/pixa/internal/imageprocessor" "sneak.berlin/go/pixa/internal/imgcache" ) @@ -217,6 +218,14 @@ func (s *Handlers) respondImageError( return } + if errors.Is(err, httpfetcher.ErrTooManyConnections) || + errors.Is(err, imageprocessor.ErrTooManyImages) { + s.respondError(w, "server busy, try again later", + http.StatusServiceUnavailable) + + return + } + s.respondError(w, "internal error", http.StatusInternalServerError) } diff --git a/internal/handlers/imageenc.go b/internal/handlers/imageenc.go index 5effd8b..56af31b 100644 --- a/internal/handlers/imageenc.go +++ b/internal/handlers/imageenc.go @@ -12,6 +12,7 @@ import ( "sneak.berlin/go/pixa/internal/encurl" "sneak.berlin/go/pixa/internal/httpfetcher" + "sneak.berlin/go/pixa/internal/imageprocessor" "sneak.berlin/go/pixa/internal/imgcache" ) @@ -124,6 +125,10 @@ func (s *Handlers) handleImageError(w http.ResponseWriter, err error) { s.respondError(w, "upstream error", http.StatusBadGateway) case errors.Is(err, httpfetcher.ErrUpstreamTimeout): s.respondError(w, "upstream timeout", http.StatusGatewayTimeout) + case errors.Is(err, httpfetcher.ErrTooManyConnections), + errors.Is(err, imageprocessor.ErrTooManyImages): + s.respondError(w, "server busy, try again later", + http.StatusServiceUnavailable) default: s.log.Error("image request failed", "error", err) s.respondError(w, "internal error", http.StatusInternalServerError) 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/httpfetcher.go b/internal/httpfetcher/httpfetcher.go index 64e27eb..d65f2da 100644 --- a/internal/httpfetcher/httpfetcher.go +++ b/internal/httpfetcher/httpfetcher.go @@ -1,5 +1,6 @@ // Package httpfetcher fetches content from upstream HTTP origins with SSRF -// protection, per-host connection limits, and content-type validation. +// protection, connection limits per host and for all hosts together, and +// content-type validation. package httpfetcher import ( @@ -28,8 +29,13 @@ const ( DefaultIdleConnTimeout = 90 * time.Second DefaultMaxRedirects = 10 DefaultMaxConnectionsPerHost = 20 + DefaultMaxConnections = 64 ) +// ConnectionWaitTimeout is how long Fetch waits for a free connection when +// MaxConnections fetches are already in progress. +const ConnectionWaitTimeout = 10 * time.Second + // MIME content types. const ( contentTypeJPEG = "image/jpeg" @@ -70,6 +76,7 @@ var ( ErrInvalidContentType = errors.New("invalid or unsupported content type") ErrUpstreamError = errors.New("upstream server error") ErrUpstreamTimeout = errors.New("upstream request timeout") + ErrTooManyConnections = errors.New("too many concurrent upstream connections") ) // Internal fetcher errors. @@ -122,6 +129,9 @@ type Config struct { AllowHTTP bool // MaxConnectionsPerHost limits concurrent connections to each upstream host. MaxConnectionsPerHost int + // MaxConnections limits concurrent connections to all upstream hosts + // together. + MaxConnections int // BlockedNetworks are operator-supplied CIDR ranges refused by the // dialer, in addition to the always-enforced built-in ranges. BlockedNetworks []netip.Prefix @@ -143,15 +153,22 @@ func DefaultConfig() *Config { }, AllowHTTP: false, MaxConnectionsPerHost: DefaultMaxConnectionsPerHost, + MaxConnections: DefaultMaxConnections, } } -// HTTPFetcher implements Fetcher with SSRF protection and per-host connection limits. +// HTTPFetcher implements Fetcher with SSRF protection and connection limits +// per host and for all hosts together. type HTTPFetcher struct { client *http.Client config *Config hostSems map[string]chan struct{} // per-host semaphores hostSemMu sync.Mutex // protects hostSems map + // allHostsSemaphore has one slot per connection allowed to all hosts + // together (config.MaxConnections). + allHostsSemaphore chan struct{} + // connectionWaitTimeout is ConnectionWaitTimeout; tests shorten it. + connectionWaitTimeout time.Duration } // New creates a new HTTPFetcher with SSRF protection. @@ -192,13 +209,18 @@ func New(config *Config) *HTTPFetcher { } return &HTTPFetcher{ - client: client, - config: config, - hostSems: make(map[string]chan struct{}), + client: client, + config: config, + hostSems: make(map[string]chan struct{}), + allHostsSemaphore: make(chan struct{}, config.MaxConnections), + connectionWaitTimeout: ConnectionWaitTimeout, } } -// Fetch retrieves content from the given URL with SSRF protection. +// Fetch retrieves content from the given URL with SSRF protection. When +// MaxConnections fetches are already in progress, it waits up to +// ConnectionWaitTimeout for one to finish, then fails with +// ErrTooManyConnections. func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, error) { // Validate URL before making request err := validateURL(ctx, url, f.config.AllowHTTP) @@ -206,24 +228,17 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro return nil, err } - // Extract host for rate limiting - host := extractHost(url) - - // Acquire semaphore slot for this host - sem := f.getHostSemaphore(host) - select { - case sem <- struct{}{}: - // Acquired slot - case <-ctx.Done(): - return nil, ctx.Err() + release, err := f.acquireConnection(ctx, extractHost(url)) + if err != nil { + return nil, err } - // If we fail before returning a result, release the slot + // If we fail before returning a result, release the connection success := false defer func() { if !success { - <-sem + release() } }() @@ -267,17 +282,52 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro return nil, fmt.Errorf("upstream request failed: %w", err) } - result, err := f.buildResult(resp, remoteAddr, fetchDuration, sem) + result, err := f.buildResult(resp, remoteAddr, fetchDuration, release) if err != nil { return nil, err } - // Mark success so defer doesn't release the semaphore + // Mark success so defer doesn't release the connection; closing the + // result's Content does success = true return result, nil } +// acquireConnection takes a slot for host, then one of the slots shared by +// all hosts, and returns the func that gives both back. The host's slot +// comes first, so fetches queued for one busy host hold no shared slot. +// Only the wait for a shared slot is bounded: after connectionWaitTimeout +// it fails with ErrTooManyConnections. +func (f *HTTPFetcher) acquireConnection( + ctx context.Context, host string, +) (func(), error) { + hostSem := f.getHostSemaphore(host) + + select { + case hostSem <- struct{}{}: + case <-ctx.Done(): + return nil, ctx.Err() + } + + select { + case f.allHostsSemaphore <- struct{}{}: + case <-time.After(f.connectionWaitTimeout): + <-hostSem + + return nil, ErrTooManyConnections + case <-ctx.Done(): + <-hostSem + + return nil, ctx.Err() + } + + return func() { + <-hostSem + <-f.allHostsSemaphore + }, nil +} + // getHostSemaphore returns the semaphore for a host, creating it if necessary. func (f *HTTPFetcher) getHostSemaphore(host string) chan struct{} { f.hostSemMu.Lock() @@ -293,12 +343,12 @@ func (f *HTTPFetcher) getHostSemaphore(host string) chan struct{} { } // buildResult validates the upstream response and assembles a FetchResult -// whose Content releases the host semaphore slot when closed. +// whose Content calls release when closed. func (f *HTTPFetcher) buildResult( resp *http.Response, remoteAddr string, fetchDuration time.Duration, - sem chan struct{}, + release func(), ) (*FetchResult, error) { // Extract HTTP version (strip "HTTP/" prefix) httpVersion := strings.TrimPrefix(resp.Proto, "HTTP/") @@ -333,7 +383,7 @@ func (f *HTTPFetcher) buildResult( } return &FetchResult{ - Content: &semaphoreReleasingReadCloser{limitedBody, resp.Body, sem}, + Content: &semaphoreReleasingReadCloser{limitedBody, resp.Body, release}, ContentLength: resp.ContentLength, ContentType: contentType, Headers: resp.Header, @@ -574,17 +624,18 @@ func (r *limitedReader) Read(p []byte) (int, error) { return n, err } -// semaphoreReleasingReadCloser releases a semaphore slot when closed. +// semaphoreReleasingReadCloser releases the fetch's connection slots when +// closed. type semaphoreReleasingReadCloser struct { *limitedReader - closer io.Closer - sem chan struct{} + closer io.Closer + release func() } func (r *semaphoreReleasingReadCloser) Close() error { err := r.closer.Close() - <-r.sem // Release semaphore slot + r.release() return err } diff --git a/internal/httpfetcher/max_connections_internal_test.go b/internal/httpfetcher/max_connections_internal_test.go new file mode 100644 index 0000000..d23816c --- /dev/null +++ b/internal/httpfetcher/max_connections_internal_test.go @@ -0,0 +1,165 @@ +package httpfetcher + +import ( + "context" + "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() +} + +// TestFetchFreesHostSlotWhenContextEndsWaitingForConnection checks that a +// fetch whose request context ends while it waits for a connection shared +// by all hosts gives its host's slot back. With MaxConnections at 1 and one +// response open, a fetch from another host takes that host's slot and waits; +// its context ends long before the 10 second wait timeout. +func TestFetchFreesHostSlotWhenContextEndsWaitingForConnection(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxConnections = 1 + + f, _ := newServerFetcher(t, srv, cfg) + + first, err := f.Fetch(testContext(t), imageURLOnPort(81)) + if err != nil { + t.Fatalf("first Fetch() error = %v", err) + } + + defer func() { _ = first.Content.Close() }() + + ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) + defer cancel() + + _, err = f.Fetch(ctx, imageURLOnPort(82)) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("second Fetch() error = %v, want context.DeadlineExceeded", err) + } + + if held := semLen(f, testPublicHost+":82"); held != 0 { + t.Errorf("the fetch kept its host's slot after its context ended: "+ + "%d held", held) + } +} + +// 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/imageprocessor.go b/internal/imageprocessor/imageprocessor.go index dd77ac1..6a9a92d 100644 --- a/internal/imageprocessor/imageprocessor.go +++ b/internal/imageprocessor/imageprocessor.go @@ -7,7 +7,9 @@ import ( "errors" "fmt" "io" + "runtime" "sync" + "time" "github.com/davidbyttow/govips/v2/vips" ) @@ -17,11 +19,21 @@ import ( //nolint:gochecknoglobals // package-level sync.Once for one-time vips init var vipsOnce sync.Once -// initVips initializes libvips with quiet logging. +// initVips initializes libvips with quiet logging, one worker thread per +// image and no operation cache. Process already works on one image per CPU +// by default, so more threads per image would only compete for the CPUs. +// Each request decodes different source bytes, so the operation cache +// would rarely be hit and would hold memory outside MaxConcurrentProcessing; +// repeated requests are served from pixa's disk cache instead. func initVips() { vipsOnce.Do(func() { vips.LoggingSettings(nil, vips.LogLevelError) - vips.Startup(nil) + vips.Startup(&vips.Config{ + ConcurrencyLevel: 1, + MaxCacheSize: 0, + MaxCacheMem: 0, + MaxCacheFiles: 0, + }) }) } @@ -106,9 +118,23 @@ var ErrInputDataTooLarge = errors.New("input data exceeds maximum allowed size") // not supported. var ErrUnsupportedOutputFormat = errors.New("unsupported output format") +// ErrTooManyImages is returned when MaxConcurrentProcessing images are being +// processed and none finishes within ProcessingWaitTimeout. +var ErrTooManyImages = errors.New("too many images being processed at once") + +// ProcessingWaitTimeout is how long Process waits for a free slot when +// MaxConcurrentProcessing images are already being processed. +const ProcessingWaitTimeout = 10 * time.Second + // ImageProcessor implements image transformation using libvips via govips. type ImageProcessor struct { maxInputBytes int64 + // processingSemaphore has one slot per image that may be processed at + // once. Process holds a slot from before it reads its input until it + // returns, so the input, the decoded image and the output all count. + processingSemaphore chan struct{} + // processingWaitTimeout is ProcessingWaitTimeout; tests shorten it. + processingWaitTimeout time.Duration } // Params holds configuration for creating an ImageProcessor. @@ -117,6 +143,9 @@ type Params struct { // MaxInputBytes is the maximum allowed input size in bytes. // If <= 0, DefaultMaxInputBytes is used. MaxInputBytes int64 + // MaxConcurrentProcessing is the most images processed at once. + // If <= 0, the number of CPUs Go uses (runtime.GOMAXPROCS(0)) is used. + MaxConcurrentProcessing int } // New creates a new image processor with the given parameters. @@ -129,17 +158,34 @@ func New(params Params) *ImageProcessor { maxInputBytes = DefaultMaxInputBytes } + maxConcurrentProcessing := params.MaxConcurrentProcessing + if maxConcurrentProcessing <= 0 { + maxConcurrentProcessing = runtime.GOMAXPROCS(0) + } + return &ImageProcessor{ - maxInputBytes: maxInputBytes, + maxInputBytes: maxInputBytes, + processingSemaphore: make(chan struct{}, maxConcurrentProcessing), + processingWaitTimeout: ProcessingWaitTimeout, } } -// Process transforms an image according to the request. +// Process transforms an image according to the request. When +// MaxConcurrentProcessing images are already being processed, it waits up +// to ProcessingWaitTimeout for one to finish, then fails with +// ErrTooManyImages. func (p *ImageProcessor) Process( - _ context.Context, + ctx context.Context, input io.Reader, req *Request, ) (*Result, error) { + release, err := p.acquireSlot(ctx) + if err != nil { + return nil, err + } + + defer release() + // Read input with a size limit to prevent unbounded memory consumption. // We read at most maxInputBytes+1 so we can detect if the input exceeds // the limit without consuming additional memory. @@ -285,6 +331,29 @@ func FormatToMIME(format Format) string { } } +// acquireSlot takes a slot in processingSemaphore, waiting at most +// processingWaitTimeout for one to free up, and returns the func that gives +// it back. A free slot is taken even when ctx has ended; only the wait for +// one stops when ctx ends, as the rest of Process does not check ctx. +func (p *ImageProcessor) acquireSlot(ctx context.Context) (func(), error) { + release := func() { <-p.processingSemaphore } + + select { + case p.processingSemaphore <- struct{}{}: + return release, nil + default: + } + + select { + case p.processingSemaphore <- struct{}{}: + return release, nil + case <-time.After(p.processingWaitTimeout): + return nil, ErrTooManyImages + case <-ctx.Done(): + return nil, ctx.Err() + } +} + // detectFormat returns the format string from a vips image. func (p *ImageProcessor) detectFormat(img *vips.ImageRef) string { format := img.Format() 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() + }) + } +} diff --git a/internal/imgcache/cache.go b/internal/imgcache/cache.go index e15b90a..ddac1b6 100644 --- a/internal/imgcache/cache.go +++ b/internal/imgcache/cache.go @@ -383,13 +383,16 @@ func (c *Cache) GetSourceMetadataID( return id, nil } -// GetSourceContent returns a reader for cached source content by its hash. -func (c *Cache) GetSourceContent(contentHash ContentHash) (io.ReadCloser, error) { +// GetSourceContent returns a reader for cached source content by its hash, +// and the content's size in bytes. +func (c *Cache) GetSourceContent( + contentHash ContentHash, +) (io.ReadCloser, int64, error) { if c.disabled { - return nil, ErrNotFound + return nil, 0, ErrNotFound } - return c.srcContent.Load(contentHash) + return c.srcContent.LoadWithSize(contentHash) } // CleanExpired removes expired entries from the cache. diff --git a/internal/imgcache/max_concurrent_processing_internal_test.go b/internal/imgcache/max_concurrent_processing_internal_test.go new file mode 100644 index 0000000..e4c310e --- /dev/null +++ b/internal/imgcache/max_concurrent_processing_internal_test.go @@ -0,0 +1,125 @@ +package imgcache + +import ( + "image/color" + "image/jpeg" + "io" + "os" + "testing" + "time" + + "sneak.berlin/go/pixa/internal/imageprocessor" +) + +// widthOnlyRequest asks for the test photo at width, its height scaled to +// keep the photo's aspect ratio. +func widthOnlyRequest(fixtures *TestFixtures, width int) *ImageRequest { + return &ImageRequest{ + SourceHost: fixtures.GoodHost, + SourcePath: testPathPhoto, + Size: Size{Width: width}, + Format: FormatJPEG, + Quality: 85, + FitMode: FitCover, + } +} + +// holdProcessingSlot takes one of proc's processing slots and returns the +// func that gives it back. Process takes its slot before it reads its input, +// so once it has read a byte from the pipe it holds the slot, until the pipe +// is closed. +func holdProcessingSlot( + t *testing.T, proc *imageprocessor.ImageProcessor, +) func() { + t.Helper() + + input, feed := io.Pipe() + + go func() { + _, _ = proc.Process(t.Context(), input, &imageprocessor.Request{}) + }() + + _, err := feed.Write([]byte{0}) + if err != nil { + t.Fatalf("Process call to hold the slot did not start: %v", err) + } + + release := func() { _ = feed.Close() } + t.Cleanup(release) + + return release +} + +// TestService_Get_WaitsForSlotBeforeReadingCachedSource checks that a +// request whose source is cached holds none of it while it waits for a +// processing slot: it reads the cached file only once it has a slot. With +// the only slot held, a request for a new width of the cached 100x100 photo +// waits; the cached file is then rewritten as a 100x50 image before the slot +// is freed, so the request must answer with that image scaled to 40x20. +func TestService_Get_WaitsForSlotBeforeReadingCachedSource(t *testing.T) { + t.Parallel() + + svc, fixtures := SetupTestService(t) + svc.processor = imageprocessor.New( + imageprocessor.Params{MaxConcurrentProcessing: 1}, + ) + + // A first request caches the photo as a source. + resp, err := svc.Get(t.Context(), widthOnlyRequest(fixtures, 50)) + if err != nil { + t.Fatalf("first Get() error = %v", err) + } + + _ = resp.Content.Close() + + contentHash, _, err := svc.cache.LookupSource(t.Context(), + widthOnlyRequest(fixtures, 50)) + if err != nil || contentHash == "" { + t.Fatalf("LookupSource() = %q, %v; want the cached source", + contentHash, err) + } + + release := holdProcessingSlot(t, svc.processor) + + var ( + waited *ImageResponse + waitedErr error + ) + + done := make(chan struct{}) + + go func() { + defer close(done) + + waited, waitedErr = svc.Get(t.Context(), widthOnlyRequest(fixtures, 40)) + }() + + // Give the request time to reach the slot: had it read the cached source + // before waiting, it would have read it by now. + time.Sleep(100 * time.Millisecond) + + err = os.WriteFile(svc.cache.srcContent.hashToPath(contentHash), + generateTestJPEG(t, 100, 50, color.RGBA{0, 0, 255, 255}), 0o600) + if err != nil { + t.Fatalf("failed to rewrite the cached source: %v", err) + } + + release() + <-done + + if waitedErr != nil { + t.Fatalf("Get() error = %v", waitedErr) + } + + defer func() { _ = waited.Content.Close() }() + + output, err := jpeg.DecodeConfig(waited.Content) + if err != nil { + t.Fatalf("failed to decode the response: %v", err) + } + + if output.Width != 40 || output.Height != 20 { + t.Errorf("response is %dx%d, want 40x20: the request read the cached "+ + "source before it had a processing slot", output.Width, output.Height) + } +} diff --git a/internal/imgcache/service.go b/internal/imgcache/service.go index d85dee8..e1152dc 100644 --- a/internal/imgcache/service.go +++ b/internal/imgcache/service.go @@ -43,6 +43,9 @@ type ServiceConfig struct { SigningKey string // Allowlist is the list of hosts that don't require signatures Allowlist []string + // MaxConcurrentProcessing is the most images processed at once; zero + // uses the image processor's default, one per CPU + MaxConcurrentProcessing int // Logger for logging Logger *slog.Logger } @@ -91,9 +94,10 @@ func NewService(cfg *ServiceConfig) (*Service, error) { } maxResponseSize := fetcherCfg.MaxResponseSize - processor := imageprocessor.New( - imageprocessor.Params{MaxInputBytes: maxResponseSize}, - ) + processor := imageprocessor.New(imageprocessor.Params{ + MaxInputBytes: maxResponseSize, + MaxConcurrentProcessing: cfg.MaxConcurrentProcessing, + }) return &Service{ cache: cfg.Cache, @@ -237,38 +241,37 @@ func (s *Service) GenerateSignedURL( baseURL, path, sig, exp, req.Quality, req.FitMode), nil } -// loadCachedSource attempts to load source content from cache, returning nil -// if the cached data is unavailable or exceeds maxResponseSize. -func (s *Service) loadCachedSource(contentHash ContentHash) []byte { - reader, err := s.cache.GetSourceContent(contentHash) +// loadCachedSource opens source content from cache, without reading it, and +// returns it with its size; nil if the cached data is unavailable, empty or +// exceeds maxResponseSize. +func (s *Service) loadCachedSource( + contentHash ContentHash, +) (io.ReadCloser, int64) { + reader, size, err := s.cache.GetSourceContent(contentHash) if err != nil { s.log.Warn("failed to load cached source, fetching", "error", err) - return nil + return nil, 0 } - // Bound the read to maxResponseSize to prevent unbounded memory use - // from unexpectedly large cached files. - limited := io.LimitReader(reader, s.maxResponseSize+1) - data, err := io.ReadAll(limited) - _ = reader.Close() + if size > s.maxResponseSize { + _ = reader.Close() - if err != nil { - s.log.Warn("failed to read cached source, fetching", "error", err) - - return nil - } - - if int64(len(data)) > s.maxResponseSize { s.log.Warn("cached source exceeds max response size, discarding", "hash", contentHash, "max_bytes", s.maxResponseSize, ) - return nil + return nil, 0 } - return data + if size == 0 { + _ = reader.Close() + + return nil, 0 + } + + return reader, size } // processFromSourceOrFetch processes an image, using cached source content @@ -285,22 +288,27 @@ func (s *Service) processFromSourceOrFetch( s.log.Warn("source lookup failed", "error", err) } - var sourceData []byte + var ( + source io.ReadCloser + sourceSize int64 + ) if contentHash != "" { s.log.Debug("using cached source", "hash", contentHash) - sourceData = s.loadCachedSource(contentHash) + source, sourceSize = s.loadCachedSource(contentHash) } // Fetch from upstream if we don't have source data or it's empty - if len(sourceData) == 0 { + if source == nil { return s.fetchAndProcess(ctx, req, cacheKey) } - // Process using cached source; nothing was fetched from upstream - resp, err := s.processAndStore( - ctx, req, cacheKey, sourceData, int64(len(sourceData)), - ) + defer func() { _ = source.Close() }() + + // Process using cached source; nothing was fetched from upstream. The + // image processor reads the source only once it has a processing slot, + // so a request waiting for one holds none of it in memory. + resp, err := s.processAndStore(ctx, req, cacheKey, source, sourceSize) return resp, 0, err } @@ -334,6 +342,10 @@ func (s *Service) fetchAndProcess( return nil, 0, fmt.Errorf("upstream fetch failed: %w", err) } + // Closing the body frees the upstream connection. It is closed only + // after processing, so the fetcher's connection limit also bounds the + // fetched sources held in memory while their requests wait for a + // processing slot. defer func() { _ = fetchResult.Content.Close() }() // Read and validate the source content @@ -379,17 +391,20 @@ func (s *Service) fetchAndProcess( // Continue even if caching fails } - resp, err := s.processAndStore(ctx, req, cacheKey, sourceData, fetchBytes) + resp, err := s.processAndStore( + ctx, req, cacheKey, bytes.NewReader(sourceData), fetchBytes, + ) return resp, fetchBytes, err } -// processAndStore processes an image and stores the result. +// processAndStore processes the image read from source and stores the +// result. func (s *Service) processAndStore( ctx context.Context, req *ImageRequest, cacheKey VariantKey, - sourceData []byte, + source io.Reader, fetchBytes int64, ) (*ImageResponse, error) { // Process the image @@ -402,7 +417,7 @@ func (s *Service) processAndStore( FitMode: imageprocessor.FitMode(req.FitMode), } - processResult, err := s.processor.Process(ctx, bytes.NewReader(sourceData), processReq) + processResult, err := s.processor.Process(ctx, source, processReq) if err != nil { return nil, fmt.Errorf("image processing failed: %w", err) }