diff --git a/internal/config/config_validation_internal_test.go b/internal/config/config_validation_internal_test.go index c12d27d..cb6f9ca 100644 --- a/internal/config/config_validation_internal_test.go +++ b/internal/config/config_validation_internal_test.go @@ -6,6 +6,7 @@ import ( "path/filepath" "strings" "testing" + "time" "git.eeqj.de/sneak/smartconfig" ) @@ -599,3 +600,217 @@ func TestEnsureStateDirFailsOnUncreatablePath(t *testing.T) { t.Errorf("error %q does not name the offending key state_dir", err.Error()) } } + +// TestOmittedOriginTimeoutsAndSizeUseDefaults checks that the CORS +// origin, the upstream fetch timeout, the upstream response size limit +// and the downstream timeout default to the values pixa used before they +// could be configured. +func TestOmittedOriginTimeoutsAndSizeUseDefaults(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.AccessControlAllowOrigin != "*" { + t.Errorf("AccessControlAllowOrigin = %q, want *", c.AccessControlAllowOrigin) + } + + if c.UpstreamFetchTimeout != 30*time.Second { + t.Errorf("UpstreamFetchTimeout = %v, want 30s", c.UpstreamFetchTimeout) + } + + if c.UpstreamMaxResponseSize != 50<<20 { + t.Errorf("UpstreamMaxResponseSize = %d, want %d (50 MiB)", + c.UpstreamMaxResponseSize, 50<<20) + } + + if c.DownstreamTimeout != 60*time.Second { + t.Errorf("DownstreamTimeout = %v, want 60s", c.DownstreamTimeout) + } +} + +// TestExplicitOriginTimeoutsAndSizeAreUsed checks that valid values for +// the CORS origin, the two timeouts and the response size limit are used +// as given. +func TestExplicitOriginTimeoutsAndSizeAreUsed(t *testing.T) { + t.Parallel() + + c, err := configFromYAML(t, signingKeyLine+` +access_control_allow_origin: https://app.example.com +upstream_fetch_timeout: 10s +upstream_max_response_size: 1048576 +downstream_timeout: 2m +`) + if err != nil { + t.Fatalf("valid config should load, got error: %v", err) + } + + if c.AccessControlAllowOrigin != "https://app.example.com" { + t.Errorf("AccessControlAllowOrigin = %q, want https://app.example.com", + c.AccessControlAllowOrigin) + } + + if c.UpstreamFetchTimeout != 10*time.Second { + t.Errorf("UpstreamFetchTimeout = %v, want 10s", c.UpstreamFetchTimeout) + } + + if c.UpstreamMaxResponseSize != 1048576 { + t.Errorf("UpstreamMaxResponseSize = %d, want 1048576", + c.UpstreamMaxResponseSize) + } + + if c.DownstreamTimeout != 2*time.Minute { + t.Errorf("DownstreamTimeout = %v, want 2m", c.DownstreamTimeout) + } +} + +// TestOriginWithPortOrAnyOriginIsAccepted checks the other accepted forms +// of access_control_allow_origin: "*", and an origin with a port. +func TestOriginWithPortOrAnyOriginIsAccepted(t *testing.T) { + t.Parallel() + + for _, origin := range []string{"*", "http://localhost:3000"} { + c, err := configFromYAML(t, signingKeyLine+ + "access_control_allow_origin: \""+origin+"\"\n") + if err != nil { + t.Fatalf("origin %q should be accepted, got error: %v", origin, err) + } + + if c.AccessControlAllowOrigin != origin { + t.Errorf("AccessControlAllowOrigin = %q, want %q", + c.AccessControlAllowOrigin, origin) + } + } +} + +// invalidTimeoutCases are configs where upstream_fetch_timeout or +// downstream_timeout is not a positive Go duration string; each must +// abort startup naming the key and the value. +func invalidTimeoutCases() []abortCase { + return []abortCase{ + { + name: "upstream_fetch_timeout not a duration", + yaml: signingKeyLine + "upstream_fetch_timeout: soon\n", + wantErrSubstrings: []string{keyUpstreamFetchTimeout, "soon"}, + }, + { + name: "upstream_fetch_timeout number without a unit", + yaml: signingKeyLine + "upstream_fetch_timeout: 45\n", + wantErrSubstrings: []string{keyUpstreamFetchTimeout, "45"}, + }, + { + name: "upstream_fetch_timeout zero", + yaml: signingKeyLine + "upstream_fetch_timeout: 0s\n", + wantErrSubstrings: []string{keyUpstreamFetchTimeout, "0s"}, + }, + { + name: "upstream_fetch_timeout negative", + yaml: signingKeyLine + "upstream_fetch_timeout: -5s\n", + wantErrSubstrings: []string{keyUpstreamFetchTimeout, "-5s"}, + }, + { + name: "upstream_fetch_timeout null", + yaml: signingKeyLine + "upstream_fetch_timeout: null\n", + wantErrSubstrings: []string{keyUpstreamFetchTimeout, nullValueText}, + }, + { + name: "downstream_timeout not a duration", + yaml: signingKeyLine + "downstream_timeout: 1 minute\n", + wantErrSubstrings: []string{keyDownstreamTimeout, "1 minute"}, + }, + { + name: "downstream_timeout zero", + yaml: signingKeyLine + "downstream_timeout: 0s\n", + wantErrSubstrings: []string{keyDownstreamTimeout, "0s"}, + }, + { + name: "downstream_timeout negative", + yaml: signingKeyLine + "downstream_timeout: -1m\n", + wantErrSubstrings: []string{keyDownstreamTimeout, "-1m"}, + }, + { + name: "downstream_timeout null", + yaml: signingKeyLine + "downstream_timeout:\n", + wantErrSubstrings: []string{keyDownstreamTimeout, nullValueText}, + }, + } +} + +// invalidSizeAndOriginCases are configs where upstream_max_response_size +// is not a positive whole number of bytes, or access_control_allow_origin +// is neither "*" nor an origin; each must abort startup naming the key +// and the value. +func invalidSizeAndOriginCases() []abortCase { + return []abortCase{ + { + name: "upstream_max_response_size with a unit", + yaml: signingKeyLine + "upstream_max_response_size: 50MB\n", + wantErrSubstrings: []string{keyUpstreamMaxResponseSize, "50MB"}, + }, + { + name: "upstream_max_response_size fractional", + yaml: signingKeyLine + "upstream_max_response_size: 1.5\n", + wantErrSubstrings: []string{keyUpstreamMaxResponseSize, "1.5"}, + }, + { + name: "upstream_max_response_size zero", + yaml: signingKeyLine + "upstream_max_response_size: 0\n", + wantErrSubstrings: []string{keyUpstreamMaxResponseSize, "0"}, + }, + { + name: "upstream_max_response_size negative", + yaml: signingKeyLine + "upstream_max_response_size: -1\n", + wantErrSubstrings: []string{keyUpstreamMaxResponseSize, "-1"}, + }, + { + name: "upstream_max_response_size null", + yaml: signingKeyLine + "upstream_max_response_size: null\n", + wantErrSubstrings: []string{keyUpstreamMaxResponseSize, nullValueText}, + }, + { + name: "access_control_allow_origin bare hostname", + yaml: signingKeyLine + "access_control_allow_origin: example.com\n", + wantErrSubstrings: []string{ + keyAccessControlAllowOrigin, "example.com", + }, + }, + { + name: "access_control_allow_origin with a path", + yaml: signingKeyLine + + "access_control_allow_origin: https://example.com/images\n", + wantErrSubstrings: []string{ + keyAccessControlAllowOrigin, "https://example.com/images", + }, + }, + { + name: "access_control_allow_origin trailing slash", + yaml: signingKeyLine + + "access_control_allow_origin: https://example.com/\n", + wantErrSubstrings: []string{ + keyAccessControlAllowOrigin, "https://example.com/", + }, + }, + { + name: "access_control_allow_origin empty", + yaml: signingKeyLine + "access_control_allow_origin: \"\"\n", + wantErrSubstrings: []string{keyAccessControlAllowOrigin}, + }, + { + name: "access_control_allow_origin null", + yaml: signingKeyLine + "access_control_allow_origin: null\n", + wantErrSubstrings: []string{keyAccessControlAllowOrigin, nullValueText}, + }, + } +} + +// TestInvalidOriginTimeoutOrSizeAbortsStartup verifies the +// no-silent-fallback rule for the CORS origin, the two timeouts and the +// response size limit: a value that does not parse or is out of range +// aborts startup naming the key and the value. +func TestInvalidOriginTimeoutOrSizeAbortsStartup(t *testing.T) { + t.Parallel() + + runAbortCases(t, append(invalidTimeoutCases(), invalidSizeAndOriginCases()...)) +} diff --git a/internal/config/env_internal_test.go b/internal/config/env_internal_test.go index 57f1f15..cdc8790 100644 --- a/internal/config/env_internal_test.go +++ b/internal/config/env_internal_test.go @@ -8,6 +8,7 @@ import ( "slices" "strings" "testing" + "time" "sneak.berlin/go/pixa/internal/globals" "sneak.berlin/go/pixa/internal/logger" @@ -69,6 +70,10 @@ func TestEnvironmentSetsEveryKey(t *testing.T) { 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") + t.Setenv("PIXA_ACCESS_CONTROL_ALLOW_ORIGIN", "https://app.example.com") + t.Setenv("PIXA_UPSTREAM_FETCH_TIMEOUT", "10s") + t.Setenv("PIXA_UPSTREAM_MAX_RESPONSE_SIZE", "1048576") + t.Setenv("PIXA_DOWNSTREAM_TIMEOUT", "2m") c, err := newFromSmartConfig(nil) if err != nil { @@ -92,6 +97,10 @@ func TestEnvironmentSetsEveryKey(t *testing.T) { cacheMaxBytesExplicit: true, BlockedNetworks: []netip.Prefix{netip.MustParsePrefix("203.0.113.0/24")}, TrustedProxies: []netip.Prefix{netip.MustParsePrefix("192.0.2.0/24")}, + AccessControlAllowOrigin: "https://app.example.com", + UpstreamFetchTimeout: 10 * time.Second, + UpstreamMaxResponseSize: 1048576, + DownstreamTimeout: 2 * time.Minute, } if !reflect.DeepEqual(*c, want) { @@ -280,6 +289,30 @@ func TestInvalidDebugFromEnvironmentAbortsStartup(t *testing.T) { wantStartupError(t, err, "PIXA_DEBUG", "maybe") } +// TestInvalidOriginTimeoutOrSizeFromEnvironmentAbortsStartup checks that +// an invalid CORS origin, timeout or response size limit in its variable +// aborts startup naming the variable and the value. +func TestInvalidOriginTimeoutOrSizeFromEnvironmentAbortsStartup(t *testing.T) { + cases := []struct { + variable string + value string + }{ + {"PIXA_ACCESS_CONTROL_ALLOW_ORIGIN", "example.com"}, + {"PIXA_UPSTREAM_FETCH_TIMEOUT", "soon"}, + {"PIXA_UPSTREAM_MAX_RESPONSE_SIZE", "50MB"}, + {"PIXA_DOWNSTREAM_TIMEOUT", "0s"}, + } + + for _, tc := range cases { + t.Run(tc.variable, func(t *testing.T) { + t.Setenv(tc.variable, tc.value) + + _, err := configFromYAML(t, signingKeyLine) + wantStartupError(t, err, tc.variable, tc.value) + }) + } +} + // TestConfigFileAloneBehavesAsBefore checks that with no variables set // (TestMain unsets them) the config file's values are used and omitted // keys take their defaults. diff --git a/internal/middleware/middleware_internal_test.go b/internal/middleware/middleware_internal_test.go index 9992a60..3421784 100644 --- a/internal/middleware/middleware_internal_test.go +++ b/internal/middleware/middleware_internal_test.go @@ -9,6 +9,53 @@ import ( "sneak.berlin/go/pixa/internal/config" ) +// TestCORSAnswersWithConfiguredOrigin checks that the CORS middleware +// uses access_control_allow_origin: "*" lets any origin read responses, +// and a single origin lets that origin read them and no other. +func TestCORSAnswersWithConfiguredOrigin(t *testing.T) { + t.Parallel() + + const appOrigin = "https://app.example.com" + + cases := []struct { + configured string + requestOrigin string + want string + }{ + {"*", "https://any.example.com", "*"}, + {appOrigin, appOrigin, appOrigin}, + {appOrigin, "https://other.example.com", ""}, + } + + testHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + for _, tc := range cases { + mw := &Middleware{ + log: slog.Default(), + config: &config.Config{AccessControlAllowOrigin: tc.configured}, + } + + handler := mw.CORS()(testHandler) + + req := httptest.NewRequestWithContext( + t.Context(), http.MethodGet, "/v1/image/example.com/a.jpg/1x1.png", nil) + req.Header.Set("Origin", tc.requestOrigin) + + rec := httptest.NewRecorder() + + handler.ServeHTTP(rec, req) + + got := rec.Header().Get("Access-Control-Allow-Origin") + if got != tc.want { + t.Errorf("configured %q, request from %q: "+ + "Access-Control-Allow-Origin = %q, want %q", + tc.configured, tc.requestOrigin, got, tc.want) + } + } +} + func TestSecurityHeaders(t *testing.T) { t.Parallel() diff --git a/internal/server/http_internal_test.go b/internal/server/http_internal_test.go index d94a380..25b60fb 100644 --- a/internal/server/http_internal_test.go +++ b/internal/server/http_internal_test.go @@ -11,11 +11,15 @@ import ( // carries every hardening timeout wired onto it, including the slowloris // defense (ReadHeaderTimeout) and the keep-alive bound (IdleTimeout). This // guards against a field being defined but never set on the server, so -// each assertion compares the server field to its constant. +// each assertion compares the server field to its constant, or, for +// WriteTimeout, to downstream_timeout from the config. func TestNewHTTPServerTimeouts(t *testing.T) { t.Parallel() - s := &Server{config: &config.Config{Port: 8080}} + s := &Server{config: &config.Config{ + Port: 8080, + DownstreamTimeout: 45 * time.Second, + }} srv := s.newHTTPServer() @@ -26,7 +30,7 @@ func TestNewHTTPServerTimeouts(t *testing.T) { }{ {"ReadTimeout", srv.ReadTimeout, HTTPReadTimeout}, {"ReadHeaderTimeout", srv.ReadHeaderTimeout, HTTPReadHeaderTimeout}, - {"WriteTimeout", srv.WriteTimeout, HTTPWriteTimeout}, + {"WriteTimeout", srv.WriteTimeout, 45 * time.Second}, {"IdleTimeout", srv.IdleTimeout, HTTPIdleTimeout}, }