chore: update golangci-lint to v2.12.2 with canonical config
All checks were successful
check / check (push) Successful in 2m3s
All checks were successful
check / check (push) Successful in 2m3s
Replace .golangci.yml with the canonical v2-schema config (default: all minus six disabled linters, lll 88, tests included) and bump every golangci-lint pin to v2.12.2: - Dockerfile: golangci/golangci-lint:v2.12.2-alpine (hash-pinned) - script/bootstrap: GOLANGCI_LINT_VERSION 2.12.2 with new linux-amd64/arm64 release-archive sha256 pins Fix all 747 findings the stricter config surfaces, with no behavior changes: t.Parallel() throughout the test suite, static sentinel errors and errors.Is comparisons, checked error returns, context propagation (contextcheck/noctx), 88-column wrapping, extracted constants and helpers for goconst/dupl/funlen/cyclop, exhaustive switch cases replicating existing defaults, and white-box test files renamed to *_internal_test.go for testpackage. Three nolint:tagliatelle directives preserve the existing snake_case JSON wire and on-disk metadata formats.
This commit is contained in:
133
.golangci.yml
133
.golangci.yml
@@ -1,117 +1,34 @@
|
|||||||
version: "2"
|
version: "2"
|
||||||
|
|
||||||
|
# Config schema uses the golangci-lint v2 layout (settings live under
|
||||||
|
# linters.settings, not top-level linters-settings) so that the
|
||||||
|
# thresholds below are actually applied by golangci-lint >= v2.
|
||||||
|
|
||||||
run:
|
run:
|
||||||
go: "1.24"
|
timeout: 5m
|
||||||
tests: false
|
modules-download-mode: readonly
|
||||||
|
|
||||||
linters:
|
linters:
|
||||||
enable:
|
default: all
|
||||||
# Additional linters requested
|
disable:
|
||||||
- testifylint # Checks usage of github.com/stretchr/testify
|
# Genuinely incompatible with project patterns
|
||||||
- usetesting # usetesting is an analyzer that detects using os.Setenv instead of t.Setenv since Go 1.17
|
- exhaustruct # Requires all struct fields
|
||||||
# - tagliatelle # Disabled: we need snake_case for external API compatibility
|
- depguard # Dependency allow/block lists
|
||||||
- nlreturn # nlreturn checks for a new line before return and branch statements
|
- godot # Requires comments to end with periods
|
||||||
- nilnil # Checks that there is no simultaneous return of nil error and an invalid value
|
- wsl # Deprecated, replaced by wsl_v5
|
||||||
- nestif # Reports deeply nested if statements
|
- wrapcheck # Too verbose for internal packages
|
||||||
- mnd # An analyzer to detect magic numbers
|
- varnamelen # Short names like db, id are idiomatic Go
|
||||||
- lll # Reports long lines
|
settings:
|
||||||
- intrange # intrange is a linter to find places where for loops could make use of an integer range
|
lll:
|
||||||
- gochecknoglobals # Check that no global variables exist
|
line-length: 88
|
||||||
|
funlen:
|
||||||
# Default/existing linters that are commonly useful
|
lines: 80
|
||||||
- govet
|
statements: 50
|
||||||
- errcheck
|
cyclop:
|
||||||
- staticcheck
|
max-complexity: 15
|
||||||
- unused
|
dupl:
|
||||||
- ineffassign
|
threshold: 100
|
||||||
- misspell
|
|
||||||
- revive
|
|
||||||
- gosec
|
|
||||||
- unconvert
|
|
||||||
- unparam
|
|
||||||
|
|
||||||
linters-settings:
|
|
||||||
lll:
|
|
||||||
line-length: 120
|
|
||||||
|
|
||||||
nestif:
|
|
||||||
min-complexity: 4
|
|
||||||
|
|
||||||
nlreturn:
|
|
||||||
block-size: 2
|
|
||||||
|
|
||||||
revive:
|
|
||||||
rules:
|
|
||||||
- name: var-naming
|
|
||||||
arguments:
|
|
||||||
- []
|
|
||||||
- []
|
|
||||||
- "upperCaseConst=true"
|
|
||||||
|
|
||||||
tagliatelle:
|
|
||||||
case:
|
|
||||||
rules:
|
|
||||||
json: snake
|
|
||||||
yaml: snake
|
|
||||||
xml: snake
|
|
||||||
bson: snake
|
|
||||||
|
|
||||||
testifylint:
|
|
||||||
enable-all: true
|
|
||||||
|
|
||||||
usetesting: {}
|
|
||||||
|
|
||||||
issues:
|
issues:
|
||||||
max-issues-per-linter: 0
|
max-issues-per-linter: 0
|
||||||
max-same-issues: 0
|
max-same-issues: 0
|
||||||
exclude-rules:
|
|
||||||
# Exclude unused parameter warnings for cobra command signatures
|
|
||||||
- text: "parameter '(args|cmd)' seems to be unused"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|
||||||
# Allow ALL_CAPS constant names
|
|
||||||
- text: "don't use ALL_CAPS in Go names"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|
||||||
# Allow snake_case JSON tags for external API compatibility
|
|
||||||
- path: "internal/types/ris.go"
|
|
||||||
linters:
|
|
||||||
- tagliatelle
|
|
||||||
|
|
||||||
# Allow snake_case JSON tags for database models
|
|
||||||
- path: "internal/database/models.go"
|
|
||||||
linters:
|
|
||||||
- tagliatelle
|
|
||||||
|
|
||||||
# Allow generic package name for types that define data structures
|
|
||||||
- path: "internal/types/"
|
|
||||||
text: "avoid meaningless package names"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|
||||||
# Allow globals in the globals package (by design)
|
|
||||||
- path: "internal/globals/"
|
|
||||||
linters:
|
|
||||||
- gochecknoglobals
|
|
||||||
|
|
||||||
# Allow globals in main (Version/Buildarch set by ldflags)
|
|
||||||
- path: "cmd/"
|
|
||||||
linters:
|
|
||||||
- gochecknoglobals
|
|
||||||
|
|
||||||
# Allow blank imports for driver registration
|
|
||||||
- text: "blank-imports"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|
||||||
# Allow unused fx.Lifecycle parameters (required by fx signature)
|
|
||||||
- text: "parameter 'lc' seems to be unused"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|
||||||
# Allow unused context parameters in fx hooks
|
|
||||||
- text: "parameter 'ctx' seems to be unused"
|
|
||||||
linters:
|
|
||||||
- revive
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
# Lint stage
|
# Lint stage
|
||||||
# golangci/golangci-lint:v2.10.1-alpine, 2026-02-17
|
# golangci/golangci-lint:v2.12.2-alpine, 2026-08-07
|
||||||
FROM golangci/golangci-lint:v2.10.1-alpine@sha256:33bc6b6156d4c7da87175f187090019769903d04dd408833b83083ed214b0ddf AS lint
|
FROM golangci/golangci-lint:v2.12.2-alpine@sha256:91b27804074a0bacea298707f016911e60cf0cdbc6c7bf5ccacb5f0606d18d60 AS lint
|
||||||
|
|
||||||
RUN apk add --no-cache make build-base vips-dev libheif-dev pkgconfig
|
RUN apk add --no-cache make build-base vips-dev libheif-dev pkgconfig
|
||||||
|
|
||||||
|
|||||||
10
TODO.md
10
TODO.md
@@ -24,6 +24,16 @@ fill up
|
|||||||
|
|
||||||
# Completed Steps
|
# Completed Steps
|
||||||
|
|
||||||
|
- 2026-08-07 update golangci-lint to v2.12.2 with the canonical
|
||||||
|
`.golangci.yml` (v2 schema, `default: all` minus six disabled
|
||||||
|
linters, `lll` 88, tests included): bumped the pinned
|
||||||
|
`golangci/golangci-lint:v2.12.2-alpine` image in `Dockerfile` and the
|
||||||
|
release-archive sha256 pins in `script/bootstrap`; fixed all 747
|
||||||
|
findings the stricter config surfaced (notably `paralleltest`,
|
||||||
|
`wsl_v5`, `goconst`, `lll`, `noinlineerr`, `err113`, `errcheck`,
|
||||||
|
`testpackage` — white-box test files renamed to
|
||||||
|
`*_internal_test.go`); three `//nolint:tagliatelle` directives keep
|
||||||
|
the snake_case JSON wire/disk formats unchanged; `make check` green
|
||||||
- 2026-08-07 manual test pass of the auth and encrypted URL flows
|
- 2026-08-07 manual test pass of the auth and encrypted URL flows
|
||||||
against a locally built and running `pixad` (built from `main` at
|
against a locally built and running `pixad` (built from `main` at
|
||||||
`6573b9d`, port 18099, local throwaway config); all six checks
|
`6573b9d`, port 18099, local throwaway config); all six checks
|
||||||
|
|||||||
@@ -30,7 +30,8 @@ func main() {
|
|||||||
|
|
||||||
rootCmd.Flags().StringVarP(&configPath, "config", "c", "", "path to config file")
|
rootCmd.Flags().StringVarP(&configPath, "config", "c", "", "path to config file")
|
||||||
|
|
||||||
if err := rootCmd.Execute(); err != nil {
|
err := rootCmd.Execute()
|
||||||
|
if err != nil {
|
||||||
fmt.Fprintln(os.Stderr, err)
|
fmt.Fprintln(os.Stderr, err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,7 +10,8 @@ import (
|
|||||||
type HostAllowList struct {
|
type HostAllowList struct {
|
||||||
// exactHosts contains hosts that must match exactly (e.g., "cdn.example.com")
|
// exactHosts contains hosts that must match exactly (e.g., "cdn.example.com")
|
||||||
exactHosts map[string]struct{}
|
exactHosts map[string]struct{}
|
||||||
// suffixHosts contains domain suffixes to match (e.g., ".example.com" matches "cdn.example.com")
|
// suffixHosts contains domain suffixes to match
|
||||||
|
// (e.g., ".example.com" matches "cdn.example.com")
|
||||||
suffixHosts []string
|
suffixHosts []string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -7,104 +7,37 @@ import (
|
|||||||
"sneak.berlin/go/pixa/internal/allowlist"
|
"sneak.berlin/go/pixa/internal/allowlist"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestHostAllowList_IsAllowed(t *testing.T) {
|
const (
|
||||||
tests := []struct {
|
testExactHost = "cdn.example.com"
|
||||||
name string
|
testImageURL = "https://cdn.example.com/image.jpg"
|
||||||
patterns []string
|
testSuffix = ".example.com"
|
||||||
testURL string
|
)
|
||||||
want bool
|
|
||||||
}{
|
type isAllowedCase struct {
|
||||||
{
|
name string
|
||||||
name: "exact match",
|
patterns []string
|
||||||
patterns: []string{"cdn.example.com"},
|
testURL string
|
||||||
testURL: "https://cdn.example.com/image.jpg",
|
want bool
|
||||||
want: true,
|
}
|
||||||
},
|
|
||||||
{
|
func runIsAllowedCases(t *testing.T, tests []isAllowedCase) {
|
||||||
name: "exact match case insensitive",
|
t.Helper()
|
||||||
patterns: []string{"CDN.Example.COM"},
|
|
||||||
testURL: "https://cdn.example.com/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "exact match not found",
|
|
||||||
patterns: []string{"cdn.example.com"},
|
|
||||||
testURL: "https://other.example.com/image.jpg",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "suffix match",
|
|
||||||
patterns: []string{".example.com"},
|
|
||||||
testURL: "https://cdn.example.com/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "suffix match deep subdomain",
|
|
||||||
patterns: []string{".example.com"},
|
|
||||||
testURL: "https://cdn.images.example.com/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "suffix match apex domain",
|
|
||||||
patterns: []string{".example.com"},
|
|
||||||
testURL: "https://example.com/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "suffix match not found",
|
|
||||||
patterns: []string{".example.com"},
|
|
||||||
testURL: "https://notexample.com/image.jpg",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "suffix match partial not allowed",
|
|
||||||
patterns: []string{".example.com"},
|
|
||||||
testURL: "https://fakeexample.com/image.jpg",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "multiple patterns",
|
|
||||||
patterns: []string{"cdn.example.com", ".images.org", "static.test.net"},
|
|
||||||
testURL: "https://photos.images.org/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "empty allow list",
|
|
||||||
patterns: []string{},
|
|
||||||
testURL: "https://cdn.example.com/image.jpg",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "nil url",
|
|
||||||
patterns: []string{"cdn.example.com"},
|
|
||||||
testURL: "",
|
|
||||||
want: false,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "url with port",
|
|
||||||
patterns: []string{"cdn.example.com"},
|
|
||||||
testURL: "https://cdn.example.com:443/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "whitespace in patterns",
|
|
||||||
patterns: []string{" cdn.example.com ", " .other.com "},
|
|
||||||
testURL: "https://cdn.example.com/image.jpg",
|
|
||||||
want: true,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
w := allowlist.New(tt.patterns)
|
w := allowlist.New(tt.patterns)
|
||||||
|
|
||||||
var u *url.URL
|
var u *url.URL
|
||||||
|
|
||||||
if tt.testURL != "" {
|
if tt.testURL != "" {
|
||||||
var err error
|
parsed, err := url.Parse(tt.testURL)
|
||||||
u, err = url.Parse(tt.testURL)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to parse test URL: %v", err)
|
t.Fatalf("failed to parse test URL: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
u = parsed
|
||||||
}
|
}
|
||||||
|
|
||||||
got := w.IsAllowed(u)
|
got := w.IsAllowed(u)
|
||||||
@@ -115,7 +48,101 @@ func TestHostAllowList_IsAllowed(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHostAllowList_IsAllowed_ExactMatch(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
runIsAllowedCases(t, []isAllowedCase{
|
||||||
|
{
|
||||||
|
name: "exact match",
|
||||||
|
patterns: []string{testExactHost},
|
||||||
|
testURL: testImageURL,
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "exact match case insensitive",
|
||||||
|
patterns: []string{"CDN.Example.COM"},
|
||||||
|
testURL: testImageURL,
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "exact match not found",
|
||||||
|
patterns: []string{testExactHost},
|
||||||
|
testURL: "https://other.example.com/image.jpg",
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple patterns",
|
||||||
|
patterns: []string{testExactHost, ".images.org", "static.test.net"},
|
||||||
|
testURL: "https://photos.images.org/image.jpg",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty allow list",
|
||||||
|
patterns: []string{},
|
||||||
|
testURL: testImageURL,
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil url",
|
||||||
|
patterns: []string{testExactHost},
|
||||||
|
testURL: "",
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "url with port",
|
||||||
|
patterns: []string{testExactHost},
|
||||||
|
testURL: "https://cdn.example.com:443/image.jpg",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "whitespace in patterns",
|
||||||
|
patterns: []string{" cdn.example.com ", " .other.com "},
|
||||||
|
testURL: testImageURL,
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHostAllowList_IsAllowed_SuffixMatch(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
runIsAllowedCases(t, []isAllowedCase{
|
||||||
|
{
|
||||||
|
name: "suffix match",
|
||||||
|
patterns: []string{testSuffix},
|
||||||
|
testURL: testImageURL,
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "suffix match deep subdomain",
|
||||||
|
patterns: []string{testSuffix},
|
||||||
|
testURL: "https://cdn.images.example.com/image.jpg",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "suffix match apex domain",
|
||||||
|
patterns: []string{testSuffix},
|
||||||
|
testURL: "https://example.com/image.jpg",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "suffix match not found",
|
||||||
|
patterns: []string{testSuffix},
|
||||||
|
testURL: "https://notexample.com/image.jpg",
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "suffix match partial not allowed",
|
||||||
|
patterns: []string{testSuffix},
|
||||||
|
testURL: "https://fakeexample.com/image.jpg",
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
func TestHostAllowList_IsEmpty(t *testing.T) {
|
func TestHostAllowList_IsEmpty(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
patterns []string
|
patterns []string
|
||||||
@@ -145,6 +172,8 @@ func TestHostAllowList_IsEmpty(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
w := allowlist.New(tt.patterns)
|
w := allowlist.New(tt.patterns)
|
||||||
if got := w.IsEmpty(); got != tt.want {
|
if got := w.IsEmpty(); got != tt.want {
|
||||||
t.Errorf("IsEmpty() = %v, want %v", got, tt.want)
|
t.Errorf("IsEmpty() = %v, want %v", got, tt.want)
|
||||||
@@ -154,6 +183,8 @@ func TestHostAllowList_IsEmpty(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHostAllowList_Count(t *testing.T) {
|
func TestHostAllowList_Count(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
patterns []string
|
patterns []string
|
||||||
@@ -183,6 +214,8 @@ func TestHostAllowList_Count(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
w := allowlist.New(tt.patterns)
|
w := allowlist.New(tt.patterns)
|
||||||
if got := w.Count(); got != tt.want {
|
if got := w.Count(); got != tt.want {
|
||||||
t.Errorf("Count() = %v, want %v", got, tt.want)
|
t.Errorf("Count() = %v, want %v", got, tt.want)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"os"
|
"os"
|
||||||
@@ -24,10 +25,17 @@ const (
|
|||||||
// Params defines dependencies for Config.
|
// Params defines dependencies for Config.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Globals *globals.Globals
|
Globals *globals.Globals
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Static validation errors.
|
||||||
|
var (
|
||||||
|
errSigningKeyRequired = errors.New("signing_key is required")
|
||||||
|
errSigningKeyTooShort = errors.New("signing_key too short")
|
||||||
|
)
|
||||||
|
|
||||||
// Config holds application configuration values.
|
// Config holds application configuration values.
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Debug bool
|
Debug bool
|
||||||
@@ -61,17 +69,19 @@ func New(_ fx.Lifecycle, params Params) (*Config, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
c := &Config{
|
c := &Config{
|
||||||
Debug: getBool(sc, "debug", false),
|
Debug: getBool(sc, "debug", false),
|
||||||
MaintenanceMode: getBool(sc, "maintenance_mode", false),
|
MaintenanceMode: getBool(sc, "maintenance_mode", false),
|
||||||
Port: getInt(sc, "port", DefaultPort),
|
Port: getInt(sc, "port", DefaultPort),
|
||||||
StateDir: getString(sc, "state_dir", DefaultStateDir),
|
StateDir: getString(sc, "state_dir", DefaultStateDir),
|
||||||
SentryDSN: getString(sc, "sentry_dsn", ""),
|
SentryDSN: getString(sc, "sentry_dsn", ""),
|
||||||
MetricsUsername: getString(sc, "metrics.username", ""),
|
MetricsUsername: getString(sc, "metrics.username", ""),
|
||||||
MetricsPassword: getString(sc, "metrics.password", ""),
|
MetricsPassword: getString(sc, "metrics.password", ""),
|
||||||
SigningKey: getString(sc, "signing_key", ""),
|
SigningKey: getString(sc, "signing_key", ""),
|
||||||
AllowlistHosts: getStringSlice(sc, "allowlist_hosts"),
|
AllowlistHosts: getStringSlice(sc),
|
||||||
AllowHTTP: getBool(sc, "allow_http", false),
|
AllowHTTP: getBool(sc, "allow_http", false),
|
||||||
UpstreamConnectionsPerHost: getInt(sc, "upstream_connections_per_host", DefaultUpstreamConnectionsPerHost),
|
UpstreamConnectionsPerHost: getInt(
|
||||||
|
sc, "upstream_connections_per_host", DefaultUpstreamConnectionsPerHost,
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build DBURL from StateDir if not explicitly set
|
// Build DBURL from StateDir if not explicitly set
|
||||||
@@ -85,7 +95,8 @@ func New(_ fx.Lifecycle, params Params) (*Config, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Validate required configuration
|
// Validate required configuration
|
||||||
if err := c.validate(); err != nil {
|
err = c.validate()
|
||||||
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -95,19 +106,22 @@ func New(_ fx.Lifecycle, params Params) (*Config, error) {
|
|||||||
// validate checks that all required configuration values are set.
|
// validate checks that all required configuration values are set.
|
||||||
func (c *Config) validate() error {
|
func (c *Config) validate() error {
|
||||||
if c.SigningKey == "" {
|
if c.SigningKey == "" {
|
||||||
return fmt.Errorf("signing_key is required")
|
return errSigningKeyRequired
|
||||||
}
|
}
|
||||||
|
|
||||||
// Minimum key length for security (32 bytes = 256 bits)
|
// Minimum key length for security (32 bytes = 256 bits)
|
||||||
const minKeyLength = 32
|
const minKeyLength = 32
|
||||||
if len(c.SigningKey) < minKeyLength {
|
if len(c.SigningKey) < minKeyLength {
|
||||||
return fmt.Errorf("signing_key must be at least %d characters", minKeyLength)
|
return fmt.Errorf(
|
||||||
|
"%w: must be at least %d characters", errSigningKeyTooShort, minKeyLength,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// loadConfigFile loads configuration from PIXA_CONFIG_PATH env var or standard locations.
|
// loadConfigFile loads configuration from the PIXA_CONFIG_PATH env var
|
||||||
|
// or standard locations.
|
||||||
func loadConfigFile(log *slog.Logger, appName string) (*smartconfig.Config, error) {
|
func loadConfigFile(log *slog.Logger, appName string) (*smartconfig.Config, error) {
|
||||||
// Check for explicit config path from environment
|
// Check for explicit config path from environment
|
||||||
if envPath := os.Getenv("PIXA_CONFIG_PATH"); envPath != "" {
|
if envPath := os.Getenv("PIXA_CONFIG_PATH"); envPath != "" {
|
||||||
@@ -133,8 +147,9 @@ func loadConfigFile(log *slog.Logger, appName string) (*smartconfig.Config, erro
|
|||||||
|
|
||||||
for _, path := range configPaths {
|
for _, path := range configPaths {
|
||||||
cleanPath := filepath.Clean(path)
|
cleanPath := filepath.Clean(path)
|
||||||
//nolint:gosec // G703: paths are hardcoded config locations
|
|
||||||
if _, statErr := os.Stat(cleanPath); statErr == nil {
|
_, statErr := os.Stat(cleanPath)
|
||||||
|
if statErr == nil {
|
||||||
sc, err := smartconfig.NewFromConfigPath(path)
|
sc, err := smartconfig.NewFromConfigPath(path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warn("failed to parse config file", "path", path, "error", err)
|
log.Warn("failed to parse config file", "path", path, "error", err)
|
||||||
@@ -190,18 +205,18 @@ func getBool(sc *smartconfig.Config, key string, defaultVal bool) bool {
|
|||||||
return val
|
return val
|
||||||
}
|
}
|
||||||
|
|
||||||
func getStringSlice(sc *smartconfig.Config, key string) []string {
|
func getStringSlice(sc *smartconfig.Config) []string {
|
||||||
if sc == nil {
|
if sc == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
val, ok := sc.Get(key)
|
val, ok := sc.Get("allowlist_hosts")
|
||||||
if !ok || val == nil {
|
if !ok || val == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Handle YAML list format
|
// Handle YAML list format
|
||||||
if slice, ok := val.([]interface{}); ok {
|
if slice, ok := val.([]any); ok {
|
||||||
result := make([]string, 0, len(slice))
|
result := make([]string, 0, len(slice))
|
||||||
for _, item := range slice {
|
for _, item := range slice {
|
||||||
if str, ok := item.(string); ok {
|
if str, ok := item.(string); ok {
|
||||||
|
|||||||
98
internal/config/config_internal_test.go
Normal file
98
internal/config/config_internal_test.go
Normal file
@@ -0,0 +1,98 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.eeqj.de/sneak/smartconfig"
|
||||||
|
)
|
||||||
|
|
||||||
|
// writeTestConfig writes yamlContent to a temp config file and returns
|
||||||
|
// the file path.
|
||||||
|
func writeTestConfig(t *testing.T, yamlContent string) string {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
configPath := filepath.Join(tmpDir, "config.yml")
|
||||||
|
|
||||||
|
err := os.WriteFile(configPath, []byte(yamlContent), 0o600)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to write test config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return configPath
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkAllowlistHosts loads the config at configPath and asserts that
|
||||||
|
// getStringSlice returns the three expected hosts.
|
||||||
|
func checkAllowlistHosts(t *testing.T, configPath string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
sc, err := loadTestConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to load config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
hosts := getStringSlice(sc)
|
||||||
|
|
||||||
|
if len(hosts) != 3 {
|
||||||
|
t.Errorf("expected 3 hosts, got %d: %v", len(hosts), hosts)
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := []string{"static.sneak.cloud", "sneak.berlin", "s3.sneak.cloud"}
|
||||||
|
for i, want := range expected {
|
||||||
|
if i >= len(hosts) {
|
||||||
|
t.Errorf("missing host at index %d: want %q", i, want)
|
||||||
|
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if hosts[i] != want {
|
||||||
|
t.Errorf("host[%d] = %q, want %q", i, hosts[i], want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetStringSlice_YAMLList(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
yamlContent := `
|
||||||
|
allowlist_hosts:
|
||||||
|
- static.sneak.cloud
|
||||||
|
- sneak.berlin
|
||||||
|
- s3.sneak.cloud
|
||||||
|
`
|
||||||
|
|
||||||
|
checkAllowlistHosts(t, writeTestConfig(t, yamlContent))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetStringSlice_CommaSeparated(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Backwards compatibility with comma-separated string values.
|
||||||
|
yamlContent := `allowlist_hosts: "static.sneak.cloud, sneak.berlin, s3.sneak.cloud"`
|
||||||
|
|
||||||
|
checkAllowlistHosts(t, writeTestConfig(t, yamlContent))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetStringSlice_Empty(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
configPath := writeTestConfig(t, `port: 8080`)
|
||||||
|
|
||||||
|
sc, err := loadTestConfig(configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to load config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
hosts := getStringSlice(sc)
|
||||||
|
if len(hosts) != 0 {
|
||||||
|
t.Errorf("expected nil or empty slice, got %v", hosts)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// loadTestConfig is a helper to load a config file for testing.
|
||||||
|
func loadTestConfig(path string) (*smartconfig.Config, error) {
|
||||||
|
return smartconfig.NewFromConfigPath(path)
|
||||||
|
}
|
||||||
@@ -1,113 +0,0 @@
|
|||||||
package config
|
|
||||||
|
|
||||||
import (
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"git.eeqj.de/sneak/smartconfig"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestGetStringSlice_YAMLList(t *testing.T) {
|
|
||||||
// Create a temp config file with YAML list format
|
|
||||||
tmpDir := t.TempDir()
|
|
||||||
configPath := filepath.Join(tmpDir, "config.yml")
|
|
||||||
|
|
||||||
yamlContent := `
|
|
||||||
allowlist_hosts:
|
|
||||||
- static.sneak.cloud
|
|
||||||
- sneak.berlin
|
|
||||||
- s3.sneak.cloud
|
|
||||||
`
|
|
||||||
err := os.WriteFile(configPath, []byte(yamlContent), 0644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to write test config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Load config using smartconfig
|
|
||||||
sc, err := loadTestConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to load config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test that getStringSlice correctly parses YAML list
|
|
||||||
hosts := getStringSlice(sc, "allowlist_hosts")
|
|
||||||
|
|
||||||
if len(hosts) != 3 {
|
|
||||||
t.Errorf("expected 3 hosts, got %d: %v", len(hosts), hosts)
|
|
||||||
}
|
|
||||||
|
|
||||||
expected := []string{"static.sneak.cloud", "sneak.berlin", "s3.sneak.cloud"}
|
|
||||||
for i, want := range expected {
|
|
||||||
if i >= len(hosts) {
|
|
||||||
t.Errorf("missing host at index %d: want %q", i, want)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if hosts[i] != want {
|
|
||||||
t.Errorf("host[%d] = %q, want %q", i, hosts[i], want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetStringSlice_CommaSeparated(t *testing.T) {
|
|
||||||
// Test backwards compatibility with comma-separated string
|
|
||||||
tmpDir := t.TempDir()
|
|
||||||
configPath := filepath.Join(tmpDir, "config.yml")
|
|
||||||
|
|
||||||
yamlContent := `allowlist_hosts: "static.sneak.cloud, sneak.berlin, s3.sneak.cloud"`
|
|
||||||
|
|
||||||
err := os.WriteFile(configPath, []byte(yamlContent), 0644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to write test config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
sc, err := loadTestConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to load config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
hosts := getStringSlice(sc, "allowlist_hosts")
|
|
||||||
|
|
||||||
if len(hosts) != 3 {
|
|
||||||
t.Errorf("expected 3 hosts, got %d: %v", len(hosts), hosts)
|
|
||||||
}
|
|
||||||
|
|
||||||
expected := []string{"static.sneak.cloud", "sneak.berlin", "s3.sneak.cloud"}
|
|
||||||
for i, want := range expected {
|
|
||||||
if i >= len(hosts) {
|
|
||||||
t.Errorf("missing host at index %d: want %q", i, want)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if hosts[i] != want {
|
|
||||||
t.Errorf("host[%d] = %q, want %q", i, hosts[i], want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetStringSlice_Empty(t *testing.T) {
|
|
||||||
tmpDir := t.TempDir()
|
|
||||||
configPath := filepath.Join(tmpDir, "config.yml")
|
|
||||||
|
|
||||||
yamlContent := `port: 8080`
|
|
||||||
|
|
||||||
err := os.WriteFile(configPath, []byte(yamlContent), 0644)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to write test config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
sc, err := loadTestConfig(configPath)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to load config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
hosts := getStringSlice(sc, "allowlist_hosts")
|
|
||||||
|
|
||||||
if hosts != nil && len(hosts) != 0 {
|
|
||||||
t.Errorf("expected nil or empty slice, got %v", hosts)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// loadTestConfig is a helper to load a config file for testing
|
|
||||||
func loadTestConfig(path string) (*smartconfig.Config, error) {
|
|
||||||
return smartconfig.NewFromConfigPath(path)
|
|
||||||
}
|
|
||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"embed"
|
"embed"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -29,10 +30,15 @@ const bootstrapVersion = 0
|
|||||||
// Params defines dependencies for Database.
|
// Params defines dependencies for Database.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// errInvalidMigrationFilename is returned when a migration filename does
|
||||||
|
// not match the "<version>[_<description>].sql" pattern.
|
||||||
|
var errInvalidMigrationFilename = errors.New("invalid migration filename")
|
||||||
|
|
||||||
// Database wraps the SQL database connection.
|
// Database wraps the SQL database connection.
|
||||||
type Database struct {
|
type Database struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
@@ -48,33 +54,31 @@ type Database struct {
|
|||||||
func ParseMigrationVersion(filename string) (int, error) {
|
func ParseMigrationVersion(filename string) (int, error) {
|
||||||
name := strings.TrimSuffix(filename, filepath.Ext(filename))
|
name := strings.TrimSuffix(filename, filepath.Ext(filename))
|
||||||
if name == "" {
|
if name == "" {
|
||||||
return 0, fmt.Errorf("invalid migration filename %q: empty name", filename)
|
return 0, fmt.Errorf("%w %q: empty name", errInvalidMigrationFilename, filename)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Split on underscore to separate version from description.
|
// Split on underscore to separate version from description.
|
||||||
// If there's no underscore, the entire stem is the version.
|
// If there's no underscore, the entire stem is the version.
|
||||||
versionStr := name
|
versionStr, _, _ := strings.Cut(name, "_")
|
||||||
if idx := strings.IndexByte(name, '_'); idx >= 0 {
|
|
||||||
versionStr = name[:idx]
|
|
||||||
}
|
|
||||||
|
|
||||||
if versionStr == "" {
|
if versionStr == "" {
|
||||||
return 0, fmt.Errorf("invalid migration filename %q: empty version prefix", filename)
|
return 0, fmt.Errorf(
|
||||||
|
"%w %q: empty version prefix", errInvalidMigrationFilename, filename,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate the version is purely numeric.
|
// Validate the version is purely numeric.
|
||||||
for _, ch := range versionStr {
|
for _, ch := range versionStr {
|
||||||
if ch < '0' || ch > '9' {
|
if ch < '0' || ch > '9' {
|
||||||
return 0, fmt.Errorf(
|
return 0, fmt.Errorf(
|
||||||
"invalid migration filename %q: version %q contains non-numeric character %q",
|
"%w %q: version %q contains non-numeric character %q",
|
||||||
filename, versionStr, string(ch),
|
errInvalidMigrationFilename, filename, versionStr, string(ch),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
version, err := strconv.Atoi(versionStr)
|
version, err := strconv.Atoi(versionStr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("invalid migration filename %q: %w", filename, err)
|
return 0, fmt.Errorf("%w %q: %w", errInvalidMigrationFilename, filename, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return version, nil
|
return version, nil
|
||||||
@@ -97,6 +101,7 @@ func New(lc fx.Lifecycle, params Params) (*Database, error) {
|
|||||||
},
|
},
|
||||||
OnStop: func(_ context.Context) error {
|
OnStop: func(_ context.Context) error {
|
||||||
s.log.Info("Database OnStop Hook")
|
s.log.Info("Database OnStop Hook")
|
||||||
|
|
||||||
if s.db != nil {
|
if s.db != nil {
|
||||||
return s.db.Close()
|
return s.db.Close()
|
||||||
}
|
}
|
||||||
@@ -108,30 +113,6 @@ func New(lc fx.Lifecycle, params Params) (*Database, error) {
|
|||||||
return s, nil
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Database) connect(ctx context.Context) error {
|
|
||||||
dbURL := s.config.DBURL
|
|
||||||
|
|
||||||
s.log.Info("connecting to database", "url", dbURL)
|
|
||||||
|
|
||||||
db, err := sql.Open("sqlite", dbURL)
|
|
||||||
if err != nil {
|
|
||||||
s.log.Error("failed to open database", "error", err)
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := db.PingContext(ctx); err != nil {
|
|
||||||
s.log.Error("failed to ping database", "error", err)
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
s.db = db
|
|
||||||
s.log.Info("database connected")
|
|
||||||
|
|
||||||
return ApplyMigrations(ctx, s.db, s.log)
|
|
||||||
}
|
|
||||||
|
|
||||||
// collectMigrations reads the embedded schema directory and returns
|
// collectMigrations reads the embedded schema directory and returns
|
||||||
// migration filenames sorted lexicographically.
|
// migration filenames sorted lexicographically.
|
||||||
func collectMigrations() ([]string, error) {
|
func collectMigrations() ([]string, error) {
|
||||||
@@ -191,7 +172,8 @@ func bootstrapMigrationsTable(ctx context.Context, db *sql.DB, log *slog.Logger)
|
|||||||
// This is exported so tests can apply the real schema without the full fx
|
// This is exported so tests can apply the real schema without the full fx
|
||||||
// lifecycle.
|
// lifecycle.
|
||||||
func ApplyMigrations(ctx context.Context, db *sql.DB, log *slog.Logger) error {
|
func ApplyMigrations(ctx context.Context, db *sql.DB, log *slog.Logger) error {
|
||||||
if err := bootstrapMigrationsTable(ctx, db, log); err != nil {
|
err := bootstrapMigrationsTable(ctx, db, log)
|
||||||
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -261,3 +243,28 @@ func ApplyMigrations(ctx context.Context, db *sql.DB, log *slog.Logger) error {
|
|||||||
func (s *Database) DB() *sql.DB {
|
func (s *Database) DB() *sql.DB {
|
||||||
return s.db
|
return s.db
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Database) connect(ctx context.Context) error {
|
||||||
|
dbURL := s.config.DBURL
|
||||||
|
|
||||||
|
s.log.Info("connecting to database", "url", dbURL)
|
||||||
|
|
||||||
|
db, err := sql.Open("sqlite", dbURL)
|
||||||
|
if err != nil {
|
||||||
|
s.log.Error("failed to open database", "error", err)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = db.PingContext(ctx)
|
||||||
|
if err != nil {
|
||||||
|
s.log.Error("failed to ping database", "error", err)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.db = db
|
||||||
|
s.log.Info("database connected")
|
||||||
|
|
||||||
|
return ApplyMigrations(ctx, s.db, s.log)
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package database
|
package database
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -17,12 +16,14 @@ func openTestDB(t *testing.T) *sql.DB {
|
|||||||
t.Fatalf("failed to open test db: %v", err)
|
t.Fatalf("failed to open test db: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Cleanup(func() { db.Close() })
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
return db
|
return db
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParseMigrationVersion(t *testing.T) {
|
func TestParseMigrationVersion(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
filename string
|
filename string
|
||||||
@@ -78,6 +79,8 @@ func TestParseMigrationVersion(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got, err := ParseMigrationVersion(tt.filename)
|
got, err := ParseMigrationVersion(tt.filename)
|
||||||
if tt.wantErr {
|
if tt.wantErr {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -101,37 +104,50 @@ func TestParseMigrationVersion(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyMigrations_CreatesSchemaAndTables(t *testing.T) {
|
func TestApplyMigrations_CreatesSchemaAndTables(t *testing.T) {
|
||||||
db := openTestDB(t)
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := ApplyMigrations(ctx, db, nil); err != nil {
|
db := openTestDB(t)
|
||||||
|
ctx := t.Context()
|
||||||
|
|
||||||
|
err := ApplyMigrations(ctx, db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("ApplyMigrations failed: %v", err)
|
t.Fatalf("ApplyMigrations failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// The schema_migrations table must exist and contain at least
|
// The schema_migrations table must exist and contain at least
|
||||||
// version 0 (the bootstrap) and 1 (the initial schema).
|
// version 0 (the bootstrap) and 1 (the initial schema).
|
||||||
rows, err := db.Query("SELECT version FROM schema_migrations ORDER BY version")
|
rows, err := db.QueryContext(
|
||||||
|
ctx, "SELECT version FROM schema_migrations ORDER BY version",
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to query schema_migrations: %v", err)
|
t.Fatalf("failed to query schema_migrations: %v", err)
|
||||||
}
|
}
|
||||||
defer rows.Close()
|
|
||||||
|
defer func() { _ = rows.Close() }()
|
||||||
|
|
||||||
var versions []int
|
var versions []int
|
||||||
|
|
||||||
for rows.Next() {
|
for rows.Next() {
|
||||||
var v int
|
var v int
|
||||||
if err := rows.Scan(&v); err != nil {
|
|
||||||
t.Fatalf("failed to scan version: %v", err)
|
scanErr := rows.Scan(&v)
|
||||||
|
if scanErr != nil {
|
||||||
|
t.Fatalf("failed to scan version: %v", scanErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
versions = append(versions, v)
|
versions = append(versions, v)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := rows.Err(); err != nil {
|
err = rows.Err()
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("row iteration error: %v", err)
|
t.Fatalf("row iteration error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(versions) < 2 {
|
if len(versions) < 2 {
|
||||||
t.Fatalf("expected at least 2 migrations recorded, got %d: %v", len(versions), versions)
|
t.Fatalf(
|
||||||
|
"expected at least 2 migrations recorded, got %d: %v",
|
||||||
|
len(versions), versions,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
if versions[0] != 0 {
|
if versions[0] != 0 {
|
||||||
@@ -143,10 +159,15 @@ func TestApplyMigrations_CreatesSchemaAndTables(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Verify that the application tables created by 001.sql exist.
|
// Verify that the application tables created by 001.sql exist.
|
||||||
for _, table := range []string{"source_content", "source_metadata", "output_content", "request_cache", "negative_cache", "cache_stats"} {
|
tables := []string{
|
||||||
|
"source_content", "source_metadata", "output_content",
|
||||||
|
"request_cache", "negative_cache", "cache_stats",
|
||||||
|
}
|
||||||
|
for _, table := range tables {
|
||||||
var count int
|
var count int
|
||||||
|
|
||||||
err := db.QueryRow(
|
err := db.QueryRowContext(
|
||||||
|
ctx,
|
||||||
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?",
|
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?",
|
||||||
table,
|
table,
|
||||||
).Scan(&count)
|
).Scan(&count)
|
||||||
@@ -161,22 +182,28 @@ func TestApplyMigrations_CreatesSchemaAndTables(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyMigrations_Idempotent(t *testing.T) {
|
func TestApplyMigrations_Idempotent(t *testing.T) {
|
||||||
db := openTestDB(t)
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := ApplyMigrations(ctx, db, nil); err != nil {
|
db := openTestDB(t)
|
||||||
|
ctx := t.Context()
|
||||||
|
|
||||||
|
err := ApplyMigrations(ctx, db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("first ApplyMigrations failed: %v", err)
|
t.Fatalf("first ApplyMigrations failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Running a second time must succeed without errors.
|
// Running a second time must succeed without errors.
|
||||||
if err := ApplyMigrations(ctx, db, nil); err != nil {
|
err = ApplyMigrations(ctx, db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("second ApplyMigrations failed: %v", err)
|
t.Fatalf("second ApplyMigrations failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify no duplicate rows in schema_migrations.
|
// Verify no duplicate rows in schema_migrations.
|
||||||
var count int
|
var count int
|
||||||
|
|
||||||
err := db.QueryRow("SELECT COUNT(*) FROM schema_migrations WHERE version = 0").Scan(&count)
|
err = db.QueryRowContext(
|
||||||
|
ctx, "SELECT COUNT(*) FROM schema_migrations WHERE version = 0",
|
||||||
|
).Scan(&count)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to count version 0 rows: %v", err)
|
t.Fatalf("failed to count version 0 rows: %v", err)
|
||||||
}
|
}
|
||||||
@@ -187,17 +214,21 @@ func TestApplyMigrations_Idempotent(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestBootstrapMigrationsTable_FreshDatabase(t *testing.T) {
|
func TestBootstrapMigrationsTable_FreshDatabase(t *testing.T) {
|
||||||
db := openTestDB(t)
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
if err := bootstrapMigrationsTable(ctx, db, nil); err != nil {
|
db := openTestDB(t)
|
||||||
|
ctx := t.Context()
|
||||||
|
|
||||||
|
err := bootstrapMigrationsTable(ctx, db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("bootstrapMigrationsTable failed: %v", err)
|
t.Fatalf("bootstrapMigrationsTable failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// schema_migrations table must exist.
|
// schema_migrations table must exist.
|
||||||
var tableCount int
|
var tableCount int
|
||||||
|
|
||||||
err := db.QueryRow(
|
err = db.QueryRowContext(
|
||||||
|
ctx,
|
||||||
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='schema_migrations'",
|
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='schema_migrations'",
|
||||||
).Scan(&tableCount)
|
).Scan(&tableCount)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -211,8 +242,8 @@ func TestBootstrapMigrationsTable_FreshDatabase(t *testing.T) {
|
|||||||
// Version 0 must be recorded.
|
// Version 0 must be recorded.
|
||||||
var recorded int
|
var recorded int
|
||||||
|
|
||||||
err = db.QueryRow(
|
err = db.QueryRowContext(
|
||||||
"SELECT COUNT(*) FROM schema_migrations WHERE version = 0",
|
ctx, "SELECT COUNT(*) FROM schema_migrations WHERE version = 0",
|
||||||
).Scan(&recorded)
|
).Scan(&recorded)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to check version: %v", err)
|
t.Fatalf("failed to check version: %v", err)
|
||||||
@@ -48,7 +48,8 @@ type Generator struct {
|
|||||||
key [seal.KeySize]byte
|
key [seal.KeySize]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewGenerator creates an encrypted URL generator with a key derived from the signing key.
|
// NewGenerator creates an encrypted URL generator with a key derived
|
||||||
|
// from the signing key.
|
||||||
func NewGenerator(signingKey string) (*Generator, error) {
|
func NewGenerator(signingKey string) (*Generator, error) {
|
||||||
key, err := seal.DeriveKey([]byte(signingKey), urlKeySalt)
|
key, err := seal.DeriveKey([]byte(signingKey), urlKeySalt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -77,7 +78,8 @@ func (g *Generator) Parse(token string) (*Payload, error) {
|
|||||||
// Decrypt
|
// Decrypt
|
||||||
data, err := seal.Decrypt(g.key, token)
|
data, err := seal.Decrypt(g.key, token)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, seal.ErrDecryptionFailed) || errors.Is(err, seal.ErrInvalidPayload) {
|
if errors.Is(err, seal.ErrDecryptionFailed) ||
|
||||||
|
errors.Is(err, seal.ErrInvalidPayload) {
|
||||||
return nil, ErrDecryptFailed
|
return nil, ErrDecryptFailed
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -86,7 +88,9 @@ func (g *Generator) Parse(token string) (*Payload, error) {
|
|||||||
|
|
||||||
// CBOR decode
|
// CBOR decode
|
||||||
var p Payload
|
var p Payload
|
||||||
if err := cbor.Unmarshal(data, &p); err != nil {
|
|
||||||
|
err = cbor.Unmarshal(data, &p)
|
||||||
|
if err != nil {
|
||||||
return nil, ErrInvalidFormat
|
return nil, ErrInvalidFormat
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,22 +1,33 @@
|
|||||||
package encurl
|
package encurl_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/encurl"
|
||||||
"sneak.berlin/go/pixa/internal/imgcache"
|
"sneak.berlin/go/pixa/internal/imgcache"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Shared test fixture strings.
|
||||||
|
const (
|
||||||
|
testSourceHost = "cdn.example.com"
|
||||||
|
testSourcePath = "/images/photo.jpg"
|
||||||
|
testSourceQuery = "v=2"
|
||||||
|
)
|
||||||
|
|
||||||
func TestGenerator_GenerateAndParse(t *testing.T) {
|
func TestGenerator_GenerateAndParse(t *testing.T) {
|
||||||
gen, err := NewGenerator("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
gen, err := encurl.NewGenerator("test-signing-key-12345")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewGenerator() error = %v", err)
|
t.Fatalf("NewGenerator() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := &Payload{
|
payload := &encurl.Payload{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testSourceHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testSourcePath,
|
||||||
SourceQuery: "v=2",
|
SourceQuery: testSourceQuery,
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: imgcache.FormatWebP,
|
Format: imgcache.FormatWebP,
|
||||||
@@ -43,38 +54,48 @@ func TestGenerator_GenerateAndParse(t *testing.T) {
|
|||||||
if parsed.SourceHost != payload.SourceHost {
|
if parsed.SourceHost != payload.SourceHost {
|
||||||
t.Errorf("SourceHost = %q, want %q", parsed.SourceHost, payload.SourceHost)
|
t.Errorf("SourceHost = %q, want %q", parsed.SourceHost, payload.SourceHost)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.SourcePath != payload.SourcePath {
|
if parsed.SourcePath != payload.SourcePath {
|
||||||
t.Errorf("SourcePath = %q, want %q", parsed.SourcePath, payload.SourcePath)
|
t.Errorf("SourcePath = %q, want %q", parsed.SourcePath, payload.SourcePath)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.SourceQuery != payload.SourceQuery {
|
if parsed.SourceQuery != payload.SourceQuery {
|
||||||
t.Errorf("SourceQuery = %q, want %q", parsed.SourceQuery, payload.SourceQuery)
|
t.Errorf("SourceQuery = %q, want %q", parsed.SourceQuery, payload.SourceQuery)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.Width != payload.Width {
|
if parsed.Width != payload.Width {
|
||||||
t.Errorf("Width = %d, want %d", parsed.Width, payload.Width)
|
t.Errorf("Width = %d, want %d", parsed.Width, payload.Width)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.Height != payload.Height {
|
if parsed.Height != payload.Height {
|
||||||
t.Errorf("Height = %d, want %d", parsed.Height, payload.Height)
|
t.Errorf("Height = %d, want %d", parsed.Height, payload.Height)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.Format != payload.Format {
|
if parsed.Format != payload.Format {
|
||||||
t.Errorf("Format = %q, want %q", parsed.Format, payload.Format)
|
t.Errorf("Format = %q, want %q", parsed.Format, payload.Format)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.Quality != payload.Quality {
|
if parsed.Quality != payload.Quality {
|
||||||
t.Errorf("Quality = %d, want %d", parsed.Quality, payload.Quality)
|
t.Errorf("Quality = %d, want %d", parsed.Quality, payload.Quality)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.FitMode != payload.FitMode {
|
if parsed.FitMode != payload.FitMode {
|
||||||
t.Errorf("FitMode = %q, want %q", parsed.FitMode, payload.FitMode)
|
t.Errorf("FitMode = %q, want %q", parsed.FitMode, payload.FitMode)
|
||||||
}
|
}
|
||||||
|
|
||||||
if parsed.ExpiresAt != payload.ExpiresAt {
|
if parsed.ExpiresAt != payload.ExpiresAt {
|
||||||
t.Errorf("ExpiresAt = %d, want %d", parsed.ExpiresAt, payload.ExpiresAt)
|
t.Errorf("ExpiresAt = %d, want %d", parsed.ExpiresAt, payload.ExpiresAt)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerator_Parse_Expired(t *testing.T) {
|
func TestGenerator_Parse_Expired(t *testing.T) {
|
||||||
gen, _ := NewGenerator("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
payload := &Payload{
|
gen, _ := encurl.NewGenerator("test-signing-key-12345")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/images/photo.jpg",
|
payload := &encurl.Payload{
|
||||||
|
SourceHost: testSourceHost,
|
||||||
|
SourcePath: testSourcePath,
|
||||||
ExpiresAt: time.Now().Add(-time.Hour).Unix(), // Already expired
|
ExpiresAt: time.Now().Add(-time.Hour).Unix(), // Already expired
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -88,13 +109,15 @@ func TestGenerator_Parse_Expired(t *testing.T) {
|
|||||||
t.Error("Parse() should fail for expired token")
|
t.Error("Parse() should fail for expired token")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrExpired {
|
if !errors.Is(err, encurl.ErrExpired) {
|
||||||
t.Errorf("Parse() error = %v, want %v", err, ErrExpired)
|
t.Errorf("Parse() error = %v, want %v", err, encurl.ErrExpired)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerator_Parse_InvalidToken(t *testing.T) {
|
func TestGenerator_Parse_InvalidToken(t *testing.T) {
|
||||||
gen, _ := NewGenerator("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
gen, _ := encurl.NewGenerator("test-signing-key-12345")
|
||||||
|
|
||||||
_, err := gen.Parse("not-a-valid-token")
|
_, err := gen.Parse("not-a-valid-token")
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -103,11 +126,13 @@ func TestGenerator_Parse_InvalidToken(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerator_Parse_TamperedToken(t *testing.T) {
|
func TestGenerator_Parse_TamperedToken(t *testing.T) {
|
||||||
gen, _ := NewGenerator("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
payload := &Payload{
|
gen, _ := encurl.NewGenerator("test-signing-key-12345")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/images/photo.jpg",
|
payload := &encurl.Payload{
|
||||||
|
SourceHost: testSourceHost,
|
||||||
|
SourcePath: testSourcePath,
|
||||||
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -126,12 +151,14 @@ func TestGenerator_Parse_TamperedToken(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerator_Parse_WrongKey(t *testing.T) {
|
func TestGenerator_Parse_WrongKey(t *testing.T) {
|
||||||
gen1, _ := NewGenerator("signing-key-1")
|
t.Parallel()
|
||||||
gen2, _ := NewGenerator("signing-key-2")
|
|
||||||
|
|
||||||
payload := &Payload{
|
gen1, _ := encurl.NewGenerator("signing-key-1")
|
||||||
SourceHost: "cdn.example.com",
|
gen2, _ := encurl.NewGenerator("signing-key-2")
|
||||||
SourcePath: "/images/photo.jpg",
|
|
||||||
|
payload := &encurl.Payload{
|
||||||
|
SourceHost: testSourceHost,
|
||||||
|
SourcePath: testSourcePath,
|
||||||
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -144,10 +171,12 @@ func TestGenerator_Parse_WrongKey(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestPayload_ToImageRequest(t *testing.T) {
|
func TestPayload_ToImageRequest(t *testing.T) {
|
||||||
payload := &Payload{
|
t.Parallel()
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/images/photo.jpg",
|
payload := &encurl.Payload{
|
||||||
SourceQuery: "v=2",
|
SourceHost: testSourceHost,
|
||||||
|
SourcePath: testSourcePath,
|
||||||
|
SourceQuery: testSourceQuery,
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: imgcache.FormatWebP,
|
Format: imgcache.FormatWebP,
|
||||||
@@ -161,55 +190,68 @@ func TestPayload_ToImageRequest(t *testing.T) {
|
|||||||
if req.SourceHost != payload.SourceHost {
|
if req.SourceHost != payload.SourceHost {
|
||||||
t.Errorf("SourceHost = %q, want %q", req.SourceHost, payload.SourceHost)
|
t.Errorf("SourceHost = %q, want %q", req.SourceHost, payload.SourceHost)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.SourcePath != payload.SourcePath {
|
if req.SourcePath != payload.SourcePath {
|
||||||
t.Errorf("SourcePath = %q, want %q", req.SourcePath, payload.SourcePath)
|
t.Errorf("SourcePath = %q, want %q", req.SourcePath, payload.SourcePath)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.SourceQuery != payload.SourceQuery {
|
if req.SourceQuery != payload.SourceQuery {
|
||||||
t.Errorf("SourceQuery = %q, want %q", req.SourceQuery, payload.SourceQuery)
|
t.Errorf("SourceQuery = %q, want %q", req.SourceQuery, payload.SourceQuery)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Size.Width != payload.Width {
|
if req.Size.Width != payload.Width {
|
||||||
t.Errorf("Width = %d, want %d", req.Size.Width, payload.Width)
|
t.Errorf("Width = %d, want %d", req.Size.Width, payload.Width)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Size.Height != payload.Height {
|
if req.Size.Height != payload.Height {
|
||||||
t.Errorf("Height = %d, want %d", req.Size.Height, payload.Height)
|
t.Errorf("Height = %d, want %d", req.Size.Height, payload.Height)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Format != payload.Format {
|
if req.Format != payload.Format {
|
||||||
t.Errorf("Format = %q, want %q", req.Format, payload.Format)
|
t.Errorf("Format = %q, want %q", req.Format, payload.Format)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Quality != payload.Quality {
|
if req.Quality != payload.Quality {
|
||||||
t.Errorf("Quality = %d, want %d", req.Quality, payload.Quality)
|
t.Errorf("Quality = %d, want %d", req.Quality, payload.Quality)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.FitMode != payload.FitMode {
|
if req.FitMode != payload.FitMode {
|
||||||
t.Errorf("FitMode = %q, want %q", req.FitMode, payload.FitMode)
|
t.Errorf("FitMode = %q, want %q", req.FitMode, payload.FitMode)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPayload_ToImageRequest_Defaults(t *testing.T) {
|
func TestPayload_ToImageRequest_Defaults(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Payload with only required fields - should get defaults
|
// Payload with only required fields - should get defaults
|
||||||
payload := &Payload{
|
payload := &encurl.Payload{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testSourceHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testSourcePath,
|
||||||
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
||||||
}
|
}
|
||||||
|
|
||||||
req := payload.ToImageRequest()
|
req := payload.ToImageRequest()
|
||||||
|
|
||||||
if req.Format != DefaultFormat {
|
if req.Format != encurl.DefaultFormat {
|
||||||
t.Errorf("Format = %q, want default %q", req.Format, DefaultFormat)
|
t.Errorf("Format = %q, want default %q", req.Format, encurl.DefaultFormat)
|
||||||
}
|
}
|
||||||
if req.Quality != DefaultQuality {
|
|
||||||
t.Errorf("Quality = %d, want default %d", req.Quality, DefaultQuality)
|
if req.Quality != encurl.DefaultQuality {
|
||||||
|
t.Errorf("Quality = %d, want default %d", req.Quality, encurl.DefaultQuality)
|
||||||
}
|
}
|
||||||
if req.FitMode != DefaultFitMode {
|
|
||||||
t.Errorf("FitMode = %q, want default %q", req.FitMode, DefaultFitMode)
|
if req.FitMode != encurl.DefaultFitMode {
|
||||||
|
t.Errorf("FitMode = %q, want default %q", req.FitMode, encurl.DefaultFitMode)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFromImageRequest(t *testing.T) {
|
func TestFromImageRequest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &imgcache.ImageRequest{
|
req := &imgcache.ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testSourceHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testSourcePath,
|
||||||
SourceQuery: "v=2",
|
SourceQuery: testSourceQuery,
|
||||||
Size: imgcache.Size{Width: 800, Height: 600},
|
Size: imgcache.Size{Width: 800, Height: 600},
|
||||||
Format: imgcache.FormatWebP,
|
Format: imgcache.FormatWebP,
|
||||||
Quality: 90,
|
Quality: 90,
|
||||||
@@ -217,52 +259,62 @@ func TestFromImageRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
expiresAt := time.Now().Add(time.Hour)
|
expiresAt := time.Now().Add(time.Hour)
|
||||||
payload := FromImageRequest(req, expiresAt)
|
payload := encurl.FromImageRequest(req, expiresAt)
|
||||||
|
|
||||||
if payload.SourceHost != req.SourceHost {
|
if payload.SourceHost != req.SourceHost {
|
||||||
t.Errorf("SourceHost = %q, want %q", payload.SourceHost, req.SourceHost)
|
t.Errorf("SourceHost = %q, want %q", payload.SourceHost, req.SourceHost)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.SourcePath != req.SourcePath {
|
if payload.SourcePath != req.SourcePath {
|
||||||
t.Errorf("SourcePath = %q, want %q", payload.SourcePath, req.SourcePath)
|
t.Errorf("SourcePath = %q, want %q", payload.SourcePath, req.SourcePath)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.Width != req.Size.Width {
|
if payload.Width != req.Size.Width {
|
||||||
t.Errorf("Width = %d, want %d", payload.Width, req.Size.Width)
|
t.Errorf("Width = %d, want %d", payload.Width, req.Size.Width)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.ExpiresAt != expiresAt.Unix() {
|
if payload.ExpiresAt != expiresAt.Unix() {
|
||||||
t.Errorf("ExpiresAt = %d, want %d", payload.ExpiresAt, expiresAt.Unix())
|
t.Errorf("ExpiresAt = %d, want %d", payload.ExpiresAt, expiresAt.Unix())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFromImageRequest_OmitsDefaults(t *testing.T) {
|
func TestFromImageRequest_OmitsDefaults(t *testing.T) {
|
||||||
// Request with default values - payload should omit them for smaller encoding
|
t.Parallel()
|
||||||
|
|
||||||
|
// Request with default values - payload should omit them for
|
||||||
|
// smaller encoding
|
||||||
req := &imgcache.ImageRequest{
|
req := &imgcache.ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testSourceHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testSourcePath,
|
||||||
Format: DefaultFormat,
|
Format: encurl.DefaultFormat,
|
||||||
Quality: DefaultQuality,
|
Quality: encurl.DefaultQuality,
|
||||||
FitMode: DefaultFitMode,
|
FitMode: encurl.DefaultFitMode,
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := FromImageRequest(req, time.Now().Add(time.Hour))
|
payload := encurl.FromImageRequest(req, time.Now().Add(time.Hour))
|
||||||
|
|
||||||
// These should be zero/empty because they match defaults
|
// These should be zero/empty because they match defaults
|
||||||
if payload.Format != "" {
|
if payload.Format != "" {
|
||||||
t.Errorf("Format should be empty for default, got %q", payload.Format)
|
t.Errorf("Format should be empty for default, got %q", payload.Format)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.Quality != 0 {
|
if payload.Quality != 0 {
|
||||||
t.Errorf("Quality should be 0 for default, got %d", payload.Quality)
|
t.Errorf("Quality should be 0 for default, got %d", payload.Quality)
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.FitMode != "" {
|
if payload.FitMode != "" {
|
||||||
t.Errorf("FitMode should be empty for default, got %q", payload.FitMode)
|
t.Errorf("FitMode should be empty for default, got %q", payload.FitMode)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerator_TokenIsURLSafe(t *testing.T) {
|
func TestGenerator_TokenIsURLSafe(t *testing.T) {
|
||||||
gen, _ := NewGenerator("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
payload := &Payload{
|
gen, _ := encurl.NewGenerator("test-signing-key-12345")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/images/photo.jpg",
|
payload := &encurl.Payload{
|
||||||
|
SourceHost: testSourceHost,
|
||||||
|
SourcePath: testSourcePath,
|
||||||
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
ExpiresAt: time.Now().Add(time.Hour).Unix(),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -35,7 +35,8 @@ func (s *Handlers) HandleRoot() http.HandlerFunc {
|
|||||||
|
|
||||||
// handleLoginPost handles login form submission.
|
// handleLoginPost handles login form submission.
|
||||||
func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
||||||
if err := r.ParseForm(); err != nil {
|
err := r.ParseForm()
|
||||||
|
if err != nil {
|
||||||
s.renderLogin(w, "Invalid form data")
|
s.renderLogin(w, "Invalid form data")
|
||||||
|
|
||||||
return
|
return
|
||||||
@@ -52,7 +53,8 @@ func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Create session
|
// Create session
|
||||||
if err := s.sessMgr.CreateSession(w); err != nil {
|
err = s.sessMgr.CreateSession(w)
|
||||||
|
if err != nil {
|
||||||
s.log.Error("failed to create session", "error", err)
|
s.log.Error("failed to create session", "error", err)
|
||||||
s.renderLogin(w, "Failed to create session")
|
s.renderLogin(w, "Failed to create session")
|
||||||
|
|
||||||
@@ -83,20 +85,14 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := r.ParseForm(); err != nil {
|
err := r.ParseForm()
|
||||||
|
if err != nil {
|
||||||
s.renderGenerator(w, &generatorData{Error: "Invalid form data"})
|
s.renderGenerator(w, &generatorData{Error: "Invalid form data"})
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse form values
|
|
||||||
sourceURL := r.FormValue("url")
|
sourceURL := r.FormValue("url")
|
||||||
widthStr := r.FormValue("width")
|
|
||||||
heightStr := r.FormValue("height")
|
|
||||||
format := r.FormValue("format")
|
|
||||||
qualityStr := r.FormValue("quality")
|
|
||||||
fit := r.FormValue("fit")
|
|
||||||
ttlStr := r.FormValue("ttl")
|
|
||||||
|
|
||||||
// Validate source URL
|
// Validate source URL
|
||||||
parsed, err := url.Parse(sourceURL)
|
parsed, err := url.Parse(sourceURL)
|
||||||
@@ -106,38 +102,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parse dimensions
|
payload, expiresAt, ttl := buildGeneratePayload(parsed, r.Form)
|
||||||
width, _ := strconv.Atoi(widthStr)
|
|
||||||
height, _ := strconv.Atoi(heightStr)
|
|
||||||
quality, _ := strconv.Atoi(qualityStr)
|
|
||||||
ttl, _ := strconv.Atoi(ttlStr)
|
|
||||||
|
|
||||||
if quality <= 0 {
|
|
||||||
quality = 85
|
|
||||||
}
|
|
||||||
|
|
||||||
// Create payload
|
|
||||||
// ttl=0 means never expires
|
|
||||||
var expiresAt time.Time
|
|
||||||
var expiresAtUnix int64
|
|
||||||
|
|
||||||
if ttl > 0 {
|
|
||||||
expiresAt = time.Now().Add(time.Duration(ttl) * time.Second)
|
|
||||||
expiresAtUnix = expiresAt.Unix()
|
|
||||||
}
|
|
||||||
// else expiresAtUnix stays 0 (never expires)
|
|
||||||
|
|
||||||
payload := &encurl.Payload{
|
|
||||||
SourceHost: parsed.Host,
|
|
||||||
SourcePath: parsed.Path,
|
|
||||||
SourceQuery: parsed.RawQuery,
|
|
||||||
Width: width,
|
|
||||||
Height: height,
|
|
||||||
Format: imgcache.ImageFormat(format),
|
|
||||||
Quality: quality,
|
|
||||||
FitMode: imgcache.FitMode(fit),
|
|
||||||
ExpiresAt: expiresAtUnix,
|
|
||||||
}
|
|
||||||
|
|
||||||
// Generate encrypted token
|
// Generate encrypted token
|
||||||
token, err := s.encGen.Generate(payload)
|
token, err := s.encGen.Generate(payload)
|
||||||
@@ -148,20 +113,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build full URL (URL-encode the token for safety)
|
generatedURL := s.buildGeneratedURL(r, token, r.FormValue("format"))
|
||||||
scheme := "https"
|
|
||||||
if s.config.Debug {
|
|
||||||
scheme = "http"
|
|
||||||
}
|
|
||||||
|
|
||||||
// Determine file extension for the trailing filename
|
|
||||||
ext := format
|
|
||||||
if ext == "" || ext == "orig" {
|
|
||||||
ext = "jpg" // Default extension
|
|
||||||
}
|
|
||||||
|
|
||||||
host := r.Host
|
|
||||||
generatedURL := scheme + "://" + host + "/v1/e/" + url.PathEscape(token) + "/img." + ext
|
|
||||||
|
|
||||||
// Format expiry for display
|
// Format expiry for display
|
||||||
expiresAtStr := "Never"
|
expiresAtStr := "Never"
|
||||||
@@ -173,16 +125,55 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
GeneratedURL: generatedURL,
|
GeneratedURL: generatedURL,
|
||||||
ExpiresAt: expiresAtStr,
|
ExpiresAt: expiresAtStr,
|
||||||
FormURL: sourceURL,
|
FormURL: sourceURL,
|
||||||
FormWidth: widthStr,
|
FormWidth: r.FormValue("width"),
|
||||||
FormHeight: heightStr,
|
FormHeight: r.FormValue("height"),
|
||||||
FormFormat: format,
|
FormFormat: r.FormValue("format"),
|
||||||
FormQuality: qualityStr,
|
FormQuality: r.FormValue("quality"),
|
||||||
FormFit: fit,
|
FormFit: r.FormValue("fit"),
|
||||||
FormTTL: ttlStr,
|
FormTTL: r.FormValue("ttl"),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildGeneratePayload parses the numeric form fields and assembles the
|
||||||
|
// encrypted URL payload. ttl=0 means never expires (ExpiresAt stays 0).
|
||||||
|
func buildGeneratePayload(
|
||||||
|
parsed *url.URL, form url.Values,
|
||||||
|
) (*encurl.Payload, time.Time, int) {
|
||||||
|
width, _ := strconv.Atoi(form.Get("width"))
|
||||||
|
height, _ := strconv.Atoi(form.Get("height"))
|
||||||
|
quality, _ := strconv.Atoi(form.Get("quality"))
|
||||||
|
ttl, _ := strconv.Atoi(form.Get("ttl"))
|
||||||
|
|
||||||
|
if quality <= 0 {
|
||||||
|
quality = 85
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
expiresAt time.Time
|
||||||
|
expiresAtUnix int64
|
||||||
|
)
|
||||||
|
|
||||||
|
if ttl > 0 {
|
||||||
|
expiresAt = time.Now().Add(time.Duration(ttl) * time.Second)
|
||||||
|
expiresAtUnix = expiresAt.Unix()
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := &encurl.Payload{
|
||||||
|
SourceHost: parsed.Host,
|
||||||
|
SourcePath: parsed.Path,
|
||||||
|
SourceQuery: parsed.RawQuery,
|
||||||
|
Width: width,
|
||||||
|
Height: height,
|
||||||
|
Format: imgcache.ImageFormat(form.Get("format")),
|
||||||
|
Quality: quality,
|
||||||
|
FitMode: imgcache.FitMode(form.Get("fit")),
|
||||||
|
ExpiresAt: expiresAtUnix,
|
||||||
|
}
|
||||||
|
|
||||||
|
return payload, expiresAt, ttl
|
||||||
|
}
|
||||||
|
|
||||||
// generatorData holds template data for the generator page.
|
// generatorData holds template data for the generator page.
|
||||||
type generatorData struct {
|
type generatorData struct {
|
||||||
GeneratedURL string
|
GeneratedURL string
|
||||||
@@ -206,7 +197,8 @@ func (s *Handlers) renderLogin(w http.ResponseWriter, errorMsg string) {
|
|||||||
Error: errorMsg,
|
Error: errorMsg,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := templates.Render(w, "login.html", data); err != nil {
|
err := templates.Render(w, "login.html", data)
|
||||||
|
if err != nil {
|
||||||
s.log.Error("failed to render login template", "error", err)
|
s.log.Error("failed to render login template", "error", err)
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
@@ -219,13 +211,16 @@ func (s *Handlers) renderGenerator(w http.ResponseWriter, data *generatorData) {
|
|||||||
data = &generatorData{}
|
data = &generatorData{}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := templates.Render(w, "generator.html", data); err != nil {
|
err := templates.Render(w, "generator.html", data)
|
||||||
|
if err != nil {
|
||||||
s.log.Error("failed to render generator template", "error", err)
|
s.log.Error("failed to render generator template", "error", err)
|
||||||
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
http.Error(w, "Internal server error", http.StatusInternalServerError)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) renderGeneratorWithForm(w http.ResponseWriter, errorMsg string, form url.Values) {
|
func (s *Handlers) renderGeneratorWithForm(
|
||||||
|
w http.ResponseWriter, errorMsg string, form url.Values,
|
||||||
|
) {
|
||||||
s.renderGenerator(w, &generatorData{
|
s.renderGenerator(w, &generatorData{
|
||||||
Error: errorMsg,
|
Error: errorMsg,
|
||||||
FormURL: form.Get("url"),
|
FormURL: form.Get("url"),
|
||||||
@@ -237,3 +232,19 @@ func (s *Handlers) renderGeneratorWithForm(w http.ResponseWriter, errorMsg strin
|
|||||||
FormTTL: form.Get("ttl"),
|
FormTTL: form.Get("ttl"),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Handlers) buildGeneratedURL(r *http.Request, token, format string) string {
|
||||||
|
// Build full URL (URL-encode the token for safety)
|
||||||
|
scheme := "https"
|
||||||
|
if s.config.Debug {
|
||||||
|
scheme = "http"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Determine file extension for the trailing filename
|
||||||
|
ext := format
|
||||||
|
if ext == "" || ext == "orig" {
|
||||||
|
ext = "jpg" // Default extension
|
||||||
|
}
|
||||||
|
|
||||||
|
return scheme + "://" + r.Host + "/v1/e/" + url.PathEscape(token) + "/img." + ext
|
||||||
|
}
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ import (
|
|||||||
// Params defines dependencies for Handlers.
|
// Params defines dependencies for Handlers.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
Healthcheck *healthcheck.Healthcheck
|
Healthcheck *healthcheck.Healthcheck
|
||||||
Database *database.Database
|
Database *database.Database
|
||||||
@@ -75,6 +76,7 @@ func (s *Handlers) initImageService() error {
|
|||||||
// Create the fetcher config
|
// Create the fetcher config
|
||||||
fetcherCfg := httpfetcher.DefaultConfig()
|
fetcherCfg := httpfetcher.DefaultConfig()
|
||||||
fetcherCfg.AllowHTTP = s.config.AllowHTTP
|
fetcherCfg.AllowHTTP = s.config.AllowHTTP
|
||||||
|
|
||||||
if s.config.UpstreamConnectionsPerHost > 0 {
|
if s.config.UpstreamConnectionsPerHost > 0 {
|
||||||
fetcherCfg.MaxConnectionsPerHost = s.config.UpstreamConnectionsPerHost
|
fetcherCfg.MaxConnectionsPerHost = s.config.UpstreamConnectionsPerHost
|
||||||
}
|
}
|
||||||
@@ -100,6 +102,7 @@ func (s *Handlers) initImageService() error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
s.sessMgr = sessMgr
|
s.sessMgr = sessMgr
|
||||||
|
|
||||||
// Initialize encrypted URL generator
|
// Initialize encrypted URL generator
|
||||||
@@ -107,6 +110,7 @@ func (s *Handlers) initImageService() error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
s.encGen = encGen
|
s.encGen = encGen
|
||||||
|
|
||||||
s.log.Info("session manager and URL generator initialized")
|
s.log.Info("session manager and URL generator initialized")
|
||||||
@@ -114,9 +118,10 @@ func (s *Handlers) initImageService() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) respondJSON(w http.ResponseWriter, data interface{}, status int) {
|
func (s *Handlers) respondJSON(w http.ResponseWriter, data any, status int) {
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
w.WriteHeader(status)
|
w.WriteHeader(status)
|
||||||
|
|
||||||
if data != nil {
|
if data != nil {
|
||||||
err := json.NewEncoder(w).Encode(data)
|
err := json.NewEncoder(w).Encode(data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -126,7 +131,7 @@ func (s *Handlers) respondJSON(w http.ResponseWriter, data interface{}, status i
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) respondError(w http.ResponseWriter, message string, status int) {
|
func (s *Handlers) respondError(w http.ResponseWriter, message string, status int) {
|
||||||
s.respondJSON(w, map[string]interface{}{
|
s.respondJSON(w, map[string]any{
|
||||||
"error": message,
|
"error": message,
|
||||||
"status": status,
|
"status": status,
|
||||||
"timestamp": time.Now().UTC().Format(time.RFC3339),
|
"timestamp": time.Now().UTC().Format(time.RFC3339),
|
||||||
|
|||||||
@@ -83,7 +83,8 @@ func setupTestDB(t *testing.T) *sql.DB {
|
|||||||
t.Fatalf("failed to open test db: %v", err)
|
t.Fatalf("failed to open test db: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := database.ApplyMigrations(context.Background(), db, nil); err != nil {
|
err = database.ApplyMigrations(context.Background(), db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to apply migrations: %v", err)
|
t.Fatalf("failed to apply migrations: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -94,14 +95,16 @@ func generateTestJPEG(t *testing.T, width, height int, c color.Color) []byte {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||||
for y := 0; y < height; y++ {
|
for y := range height {
|
||||||
for x := 0; x < width; x++ {
|
for x := range width {
|
||||||
img.Set(x, y, c)
|
img.Set(x, y, c)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 85}); err != nil {
|
|
||||||
|
err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 85})
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test JPEG: %v", err)
|
t.Fatalf("failed to encode test JPEG: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,7 +120,9 @@ func newMockFetcher(fs fs.FS) *mockFetcher {
|
|||||||
return &mockFetcher{fs: fs}
|
return &mockFetcher{fs: fs}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *mockFetcher) Fetch(ctx context.Context, url string) (*httpfetcher.FetchResult, error) {
|
func (f *mockFetcher) Fetch(
|
||||||
|
_ context.Context, url string,
|
||||||
|
) (*httpfetcher.FetchResult, error) {
|
||||||
// Remove https:// prefix
|
// Remove https:// prefix
|
||||||
path := url[8:] // Remove "https://"
|
path := url[8:] // Remove "https://"
|
||||||
|
|
||||||
@@ -134,13 +139,16 @@ func (f *mockFetcher) Fetch(ctx context.Context, url string) (*httpfetcher.Fetch
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandleImage_HEAD_ReturnsHeadersOnly(t *testing.T) {
|
func TestHandleImage_HEAD_ReturnsHeadersOnly(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
fix := setupTestHandler(t)
|
fix := setupTestHandler(t)
|
||||||
|
|
||||||
// Create a chi router to properly handle wildcards
|
// Create a chi router to properly handle wildcards
|
||||||
r := chi.NewRouter()
|
r := chi.NewRouter()
|
||||||
r.Head("/v1/image/*", fix.handler.HandleImage())
|
r.Head("/v1/image/*", fix.handler.HandleImage())
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodHead, "/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodHead,
|
||||||
|
"/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
r.ServeHTTP(rec, req)
|
r.ServeHTTP(rec, req)
|
||||||
@@ -167,13 +175,16 @@ func TestHandleImage_HEAD_ReturnsHeadersOnly(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandleImage_ConditionalRequest_IfNoneMatch_Returns304(t *testing.T) {
|
func TestHandleImage_ConditionalRequest_IfNoneMatch_Returns304(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
fix := setupTestHandler(t)
|
fix := setupTestHandler(t)
|
||||||
|
|
||||||
r := chi.NewRouter()
|
r := chi.NewRouter()
|
||||||
r.Get("/v1/image/*", fix.handler.HandleImage())
|
r.Get("/v1/image/*", fix.handler.HandleImage())
|
||||||
|
|
||||||
// First request to get the ETag
|
// First request to get the ETag
|
||||||
req1 := httptest.NewRequest(http.MethodGet, "/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
req1 := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
||||||
rec1 := httptest.NewRecorder()
|
rec1 := httptest.NewRecorder()
|
||||||
|
|
||||||
r.ServeHTTP(rec1, req1)
|
r.ServeHTTP(rec1, req1)
|
||||||
@@ -188,15 +199,18 @@ func TestHandleImage_ConditionalRequest_IfNoneMatch_Returns304(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Second request with If-None-Match header
|
// Second request with If-None-Match header
|
||||||
req2 := httptest.NewRequest(http.MethodGet, "/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
req2 := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
||||||
req2.Header.Set("If-None-Match", etag)
|
req2.Header.Set("If-None-Match", etag)
|
||||||
|
|
||||||
rec2 := httptest.NewRecorder()
|
rec2 := httptest.NewRecorder()
|
||||||
|
|
||||||
r.ServeHTTP(rec2, req2)
|
r.ServeHTTP(rec2, req2)
|
||||||
|
|
||||||
// Should return 304 Not Modified
|
// Should return 304 Not Modified
|
||||||
if rec2.Code != http.StatusNotModified {
|
if rec2.Code != http.StatusNotModified {
|
||||||
t.Errorf("Conditional request status = %d, want %d", rec2.Code, http.StatusNotModified)
|
t.Errorf("Conditional request status = %d, want %d",
|
||||||
|
rec2.Code, http.StatusNotModified)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Body should be empty for 304 response
|
// Body should be empty for 304 response
|
||||||
@@ -206,21 +220,26 @@ func TestHandleImage_ConditionalRequest_IfNoneMatch_Returns304(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandleImage_ConditionalRequest_IfNoneMatch_DifferentETag(t *testing.T) {
|
func TestHandleImage_ConditionalRequest_IfNoneMatch_DifferentETag(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
fix := setupTestHandler(t)
|
fix := setupTestHandler(t)
|
||||||
|
|
||||||
r := chi.NewRouter()
|
r := chi.NewRouter()
|
||||||
r.Get("/v1/image/*", fix.handler.HandleImage())
|
r.Get("/v1/image/*", fix.handler.HandleImage())
|
||||||
|
|
||||||
// Request with non-matching ETag
|
// Request with non-matching ETag
|
||||||
req := httptest.NewRequest(http.MethodGet, "/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
||||||
req.Header.Set("If-None-Match", `"different-etag"`)
|
req.Header.Set("If-None-Match", `"different-etag"`)
|
||||||
|
|
||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
r.ServeHTTP(rec, req)
|
r.ServeHTTP(rec, req)
|
||||||
|
|
||||||
// Should return 200 OK with full response
|
// Should return 200 OK with full response
|
||||||
if rec.Code != http.StatusOK {
|
if rec.Code != http.StatusOK {
|
||||||
t.Errorf("Request with non-matching ETag status = %d, want %d", rec.Code, http.StatusOK)
|
t.Errorf("Request with non-matching ETag status = %d, want %d",
|
||||||
|
rec.Code, http.StatusOK)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Body should not be empty
|
// Body should not be empty
|
||||||
@@ -230,12 +249,15 @@ func TestHandleImage_ConditionalRequest_IfNoneMatch_DifferentETag(t *testing.T)
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandleImage_ETagHeader(t *testing.T) {
|
func TestHandleImage_ETagHeader(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
fix := setupTestHandler(t)
|
fix := setupTestHandler(t)
|
||||||
|
|
||||||
r := chi.NewRouter()
|
r := chi.NewRouter()
|
||||||
r.Get("/v1/image/*", fix.handler.HandleImage())
|
r.Get("/v1/image/*", fix.handler.HandleImage())
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||||
|
"/v1/image/"+fix.goodHost+"/images/photo.jpg/50x50.jpeg", nil)
|
||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
r.ServeHTTP(rec, req)
|
r.ServeHTTP(rec, req)
|
||||||
@@ -16,64 +16,14 @@ import (
|
|||||||
// /v1/image/<host>/<path>/<width>x<height>.<format>
|
// /v1/image/<host>/<path>/<width>x<height>.<format>
|
||||||
func (s *Handlers) HandleImage() http.HandlerFunc {
|
func (s *Handlers) HandleImage() http.HandlerFunc {
|
||||||
return func(w http.ResponseWriter, r *http.Request) {
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
ctx := r.Context()
|
req, ok := s.parseImageRequest(w, r)
|
||||||
|
if !ok {
|
||||||
// Get the wildcard path from chi
|
|
||||||
pathParam := chi.URLParam(r, "*")
|
|
||||||
|
|
||||||
// Parse the URL path
|
|
||||||
parsed, err := imgcache.ParseImagePath(pathParam)
|
|
||||||
if err != nil {
|
|
||||||
s.log.Warn("failed to parse image URL",
|
|
||||||
"path", pathParam,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
s.respondError(w, "invalid image URL: "+err.Error(), http.StatusBadRequest)
|
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Convert to ImageRequest
|
|
||||||
req := parsed.ToImageRequest()
|
|
||||||
|
|
||||||
// Parse signature params from query string
|
|
||||||
query := r.URL.Query()
|
|
||||||
req.Signature = query.Get("sig")
|
|
||||||
|
|
||||||
if expStr := query.Get("exp"); expStr != "" {
|
|
||||||
if exp, err := strconv.ParseInt(expStr, 10, 64); err == nil {
|
|
||||||
req.Expires = time.Unix(exp, 0)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Parse optional quality and fit params
|
|
||||||
if qStr := query.Get("q"); qStr != "" {
|
|
||||||
if q, err := strconv.Atoi(qStr); err == nil && q > 0 && q <= 100 {
|
|
||||||
req.Quality = q
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if fit := query.Get("fit"); fit != "" {
|
|
||||||
req.FitMode = imgcache.FitMode(fit)
|
|
||||||
if err := imgcache.ValidateFitMode(req.FitMode); err != nil {
|
|
||||||
s.respondError(w, "invalid fit mode: "+fit, http.StatusBadRequest)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Default quality if not set
|
|
||||||
if req.Quality == 0 {
|
|
||||||
req.Quality = 85
|
|
||||||
}
|
|
||||||
|
|
||||||
// Default fit mode if not set
|
|
||||||
if req.FitMode == "" {
|
|
||||||
req.FitMode = imgcache.FitCover
|
|
||||||
}
|
|
||||||
|
|
||||||
// Validate signature if required
|
// Validate signature if required
|
||||||
if err := s.imgSvc.ValidateRequest(req); err != nil {
|
err := s.imgSvc.ValidateRequest(req)
|
||||||
|
if err != nil {
|
||||||
s.log.Warn("signature validation failed",
|
s.log.Warn("signature validation failed",
|
||||||
"host", req.SourceHost,
|
"host", req.SourceHost,
|
||||||
"path", req.SourcePath,
|
"path", req.SourcePath,
|
||||||
@@ -89,83 +39,17 @@ func (s *Handlers) HandleImage() http.HandlerFunc {
|
|||||||
|
|
||||||
// Get the image (from cache or fetch/process)
|
// Get the image (from cache or fetch/process)
|
||||||
startTime := time.Now()
|
startTime := time.Now()
|
||||||
resp, err := s.imgSvc.Get(ctx, req)
|
|
||||||
|
resp, err := s.imgSvc.Get(r.Context(), req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Error("failed to get image",
|
s.respondImageError(w, req, err)
|
||||||
"host", req.SourceHost,
|
|
||||||
"path", req.SourcePath,
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
|
|
||||||
// Check for specific error types
|
|
||||||
if errors.Is(err, httpfetcher.ErrSSRFBlocked) {
|
|
||||||
s.respondError(w, "forbidden", http.StatusForbidden)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if errors.Is(err, httpfetcher.ErrUpstreamError) {
|
|
||||||
s.respondError(w, "upstream error", http.StatusBadGateway)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
s.respondError(w, "internal error", http.StatusInternalServerError)
|
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() { _ = resp.Content.Close() }()
|
defer func() { _ = resp.Content.Close() }()
|
||||||
|
|
||||||
// Set response headers
|
s.writeImageResponse(w, r, req, resp, cacheKey, startTime)
|
||||||
w.Header().Set("Content-Type", resp.ContentType)
|
|
||||||
if resp.ContentLength > 0 {
|
|
||||||
w.Header().Set("Content-Length", strconv.FormatInt(resp.ContentLength, 10))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Cache control headers
|
|
||||||
w.Header().Set("Cache-Control", "public, max-age=31536000, immutable")
|
|
||||||
w.Header().Set("X-Pixa-Cache", string(resp.CacheStatus))
|
|
||||||
|
|
||||||
if resp.ETag != "" {
|
|
||||||
w.Header().Set("ETag", resp.ETag)
|
|
||||||
|
|
||||||
// Check for conditional request (If-None-Match)
|
|
||||||
if ifNoneMatch := r.Header.Get("If-None-Match"); ifNoneMatch != "" {
|
|
||||||
if ifNoneMatch == resp.ETag {
|
|
||||||
w.WriteHeader(http.StatusNotModified)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle HEAD request - return headers only
|
|
||||||
if r.Method == http.MethodHead {
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stream the response
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
|
|
||||||
servedBytes, err := io.Copy(w, resp.Content)
|
|
||||||
if err != nil {
|
|
||||||
s.log.Error("failed to write response",
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Log cache status and timing after serving
|
|
||||||
duration := time.Since(startTime)
|
|
||||||
s.log.Info("image served",
|
|
||||||
"cache_key", cacheKey,
|
|
||||||
"cache_status", resp.CacheStatus,
|
|
||||||
"duration_ms", duration.Milliseconds(),
|
|
||||||
"format", req.Format,
|
|
||||||
"served_bytes", servedBytes,
|
|
||||||
"fetched_bytes", resp.FetchedBytes,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -180,3 +64,156 @@ func (s *Handlers) HandleRobotsTxt() http.HandlerFunc {
|
|||||||
_, _ = w.Write(robotsTxt)
|
_, _ = w.Write(robotsTxt)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseImageRequest parses the wildcard path and query parameters into
|
||||||
|
// an ImageRequest. On invalid input it writes an error response and
|
||||||
|
// returns false.
|
||||||
|
func (s *Handlers) parseImageRequest(
|
||||||
|
w http.ResponseWriter, r *http.Request,
|
||||||
|
) (*imgcache.ImageRequest, bool) {
|
||||||
|
// Get the wildcard path from chi
|
||||||
|
pathParam := chi.URLParam(r, "*")
|
||||||
|
|
||||||
|
// Parse the URL path
|
||||||
|
parsed, err := imgcache.ParseImagePath(pathParam)
|
||||||
|
if err != nil {
|
||||||
|
s.log.Warn("failed to parse image URL",
|
||||||
|
"path", pathParam,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
s.respondError(w, "invalid image URL: "+err.Error(), http.StatusBadRequest)
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Convert to ImageRequest
|
||||||
|
req := parsed.ToImageRequest()
|
||||||
|
|
||||||
|
// Parse signature params from query string
|
||||||
|
query := r.URL.Query()
|
||||||
|
req.Signature = query.Get("sig")
|
||||||
|
|
||||||
|
if expStr := query.Get("exp"); expStr != "" {
|
||||||
|
exp, parseErr := strconv.ParseInt(expStr, 10, 64)
|
||||||
|
if parseErr == nil {
|
||||||
|
req.Expires = time.Unix(exp, 0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse optional quality and fit params
|
||||||
|
if qStr := query.Get("q"); qStr != "" {
|
||||||
|
q, parseErr := strconv.Atoi(qStr)
|
||||||
|
if parseErr == nil && q > 0 && q <= 100 {
|
||||||
|
req.Quality = q
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if fit := query.Get("fit"); fit != "" {
|
||||||
|
req.FitMode = imgcache.FitMode(fit)
|
||||||
|
|
||||||
|
fitErr := imgcache.ValidateFitMode(req.FitMode)
|
||||||
|
if fitErr != nil {
|
||||||
|
s.respondError(w, "invalid fit mode: "+fit, http.StatusBadRequest)
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default quality if not set
|
||||||
|
if req.Quality == 0 {
|
||||||
|
req.Quality = 85
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default fit mode if not set
|
||||||
|
if req.FitMode == "" {
|
||||||
|
req.FitMode = imgcache.FitCover
|
||||||
|
}
|
||||||
|
|
||||||
|
return req, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// respondImageError maps image retrieval errors to HTTP responses.
|
||||||
|
func (s *Handlers) respondImageError(
|
||||||
|
w http.ResponseWriter, req *imgcache.ImageRequest, err error,
|
||||||
|
) {
|
||||||
|
s.log.Error("failed to get image",
|
||||||
|
"host", req.SourceHost,
|
||||||
|
"path", req.SourcePath,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
|
||||||
|
// Check for specific error types
|
||||||
|
if errors.Is(err, httpfetcher.ErrSSRFBlocked) {
|
||||||
|
s.respondError(w, "forbidden", http.StatusForbidden)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if errors.Is(err, httpfetcher.ErrUpstreamError) {
|
||||||
|
s.respondError(w, "upstream error", http.StatusBadGateway)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s.respondError(w, "internal error", http.StatusInternalServerError)
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeImageResponse writes headers and streams the image content,
|
||||||
|
// handling conditional and HEAD requests.
|
||||||
|
func (s *Handlers) writeImageResponse(
|
||||||
|
w http.ResponseWriter, r *http.Request,
|
||||||
|
req *imgcache.ImageRequest, resp *imgcache.ImageResponse,
|
||||||
|
cacheKey imgcache.VariantKey, startTime time.Time,
|
||||||
|
) {
|
||||||
|
// Set response headers
|
||||||
|
w.Header().Set("Content-Type", resp.ContentType)
|
||||||
|
|
||||||
|
if resp.ContentLength > 0 {
|
||||||
|
w.Header().Set("Content-Length", strconv.FormatInt(resp.ContentLength, 10))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache control headers
|
||||||
|
w.Header().Set("Cache-Control", "public, max-age=31536000, immutable")
|
||||||
|
w.Header().Set("X-Pixa-Cache", string(resp.CacheStatus))
|
||||||
|
|
||||||
|
if resp.ETag != "" {
|
||||||
|
w.Header().Set("ETag", resp.ETag)
|
||||||
|
|
||||||
|
// Check for conditional request (If-None-Match)
|
||||||
|
if ifNoneMatch := r.Header.Get("If-None-Match"); ifNoneMatch != "" {
|
||||||
|
if ifNoneMatch == resp.ETag {
|
||||||
|
w.WriteHeader(http.StatusNotModified)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle HEAD request - return headers only
|
||||||
|
if r.Method == http.MethodHead {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stream the response
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
|
||||||
|
servedBytes, err := io.Copy(w, resp.Content)
|
||||||
|
if err != nil {
|
||||||
|
s.log.Error("failed to write response",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Log cache status and timing after serving
|
||||||
|
duration := time.Since(startTime)
|
||||||
|
s.log.Info("image served",
|
||||||
|
"cache_key", cacheKey,
|
||||||
|
"cache_status", resp.CacheStatus,
|
||||||
|
"duration_ms", duration.Milliseconds(),
|
||||||
|
"format", req.Format,
|
||||||
|
"served_bytes", servedBytes,
|
||||||
|
"fetched_bytes", resp.FetchedBytes,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|||||||
@@ -15,8 +15,9 @@ import (
|
|||||||
"sneak.berlin/go/pixa/internal/imgcache"
|
"sneak.berlin/go/pixa/internal/imgcache"
|
||||||
)
|
)
|
||||||
|
|
||||||
// HandleImageEnc handles requests to /v1/e/{token}/* for encrypted image URLs.
|
// HandleImageEnc handles requests to /v1/e/{token}/* for encrypted
|
||||||
// The trailing path (e.g., /img.jpg) is ignored but helps browsers identify the content type.
|
// image URLs. The trailing path (e.g., /img.jpg) is ignored but helps
|
||||||
|
// browsers identify the content type.
|
||||||
func (s *Handlers) HandleImageEnc() http.HandlerFunc {
|
func (s *Handlers) HandleImageEnc() http.HandlerFunc {
|
||||||
return func(w http.ResponseWriter, r *http.Request) {
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
ctx := r.Context()
|
ctx := r.Context()
|
||||||
@@ -57,7 +58,8 @@ func (s *Handlers) HandleImageEnc() http.HandlerFunc {
|
|||||||
"format", req.Format,
|
"format", req.Format,
|
||||||
)
|
)
|
||||||
|
|
||||||
// Fetch and process the image (no signature validation needed - encrypted URL is trusted)
|
// Fetch and process the image (no signature validation
|
||||||
|
// needed - encrypted URL is trusted)
|
||||||
resp, err := s.imgSvc.Get(ctx, req)
|
resp, err := s.imgSvc.Get(ctx, req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.handleImageError(w, err)
|
s.handleImageError(w, err)
|
||||||
@@ -68,6 +70,7 @@ func (s *Handlers) HandleImageEnc() http.HandlerFunc {
|
|||||||
|
|
||||||
// Set response headers
|
// Set response headers
|
||||||
w.Header().Set("Content-Type", resp.ContentType)
|
w.Header().Set("Content-Type", resp.ContentType)
|
||||||
|
|
||||||
if resp.ContentLength > 0 {
|
if resp.ContentLength > 0 {
|
||||||
w.Header().Set("Content-Length", strconv.FormatInt(resp.ContentLength, 10))
|
w.Header().Set("Content-Length", strconv.FormatInt(resp.ContentLength, 10))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
// Params defines dependencies for Healthcheck.
|
// Params defines dependencies for Healthcheck.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Globals *globals.Globals
|
Globals *globals.Globals
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
@@ -53,6 +54,8 @@ func New(lc fx.Lifecycle, params Params) (*Healthcheck, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Response is the JSON response for health checks.
|
// Response is the JSON response for health checks.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // health endpoint response format uses snake_case
|
||||||
type Response struct {
|
type Response struct {
|
||||||
Status string `json:"status"`
|
Status string `json:"status"`
|
||||||
Now string `json:"now"`
|
Now string `json:"now"`
|
||||||
@@ -63,10 +66,6 @@ type Response struct {
|
|||||||
Maintenance bool `json:"maintenance_mode"`
|
Maintenance bool `json:"maintenance_mode"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Healthcheck) uptime() time.Duration {
|
|
||||||
return time.Since(s.StartupTime)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Healthcheck returns the current health status.
|
// Healthcheck returns the current health status.
|
||||||
func (s *Healthcheck) Healthcheck() *Response {
|
func (s *Healthcheck) Healthcheck() *Response {
|
||||||
resp := &Response{
|
resp := &Response{
|
||||||
@@ -81,3 +80,7 @@ func (s *Healthcheck) Healthcheck() *Response {
|
|||||||
|
|
||||||
return resp
|
return resp
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Healthcheck) uptime() time.Duration {
|
||||||
|
return time.Since(s.StartupTime)
|
||||||
|
}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptrace"
|
"net/http/httptrace"
|
||||||
neturl "net/url"
|
neturl "net/url"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -28,6 +29,23 @@ const (
|
|||||||
DefaultMaxConnectionsPerHost = 20
|
DefaultMaxConnectionsPerHost = 20
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// MIME content types.
|
||||||
|
const (
|
||||||
|
contentTypeJPEG = "image/jpeg"
|
||||||
|
contentTypePNG = "image/png"
|
||||||
|
contentTypeGIF = "image/gif"
|
||||||
|
contentTypeWebP = "image/webp"
|
||||||
|
contentTypeAVIF = "image/avif"
|
||||||
|
contentTypeSVG = "image/svg+xml"
|
||||||
|
contentTypeOctetStream = "application/octet-stream"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Loopback addresses blocked by SSRF protection.
|
||||||
|
const (
|
||||||
|
localhostIPv4 = "127.0.0.1"
|
||||||
|
localhostIPv6 = "::1"
|
||||||
|
)
|
||||||
|
|
||||||
// Fetcher errors.
|
// Fetcher errors.
|
||||||
var (
|
var (
|
||||||
ErrSSRFBlocked = errors.New("request blocked: private or internal IP")
|
ErrSSRFBlocked = errors.New("request blocked: private or internal IP")
|
||||||
@@ -39,6 +57,12 @@ var (
|
|||||||
ErrUpstreamTimeout = errors.New("upstream request timeout")
|
ErrUpstreamTimeout = errors.New("upstream request timeout")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Internal fetcher errors.
|
||||||
|
var (
|
||||||
|
errTooManyRedirects = errors.New("too many redirects")
|
||||||
|
errConnectFailed = errors.New("failed to connect")
|
||||||
|
)
|
||||||
|
|
||||||
// Fetcher retrieves content from upstream origins.
|
// Fetcher retrieves content from upstream origins.
|
||||||
type Fetcher interface {
|
type Fetcher interface {
|
||||||
// Fetch retrieves content from the given URL.
|
// Fetch retrieves content from the given URL.
|
||||||
@@ -92,12 +116,12 @@ func DefaultConfig() *Config {
|
|||||||
MaxResponseSize: DefaultMaxResponseSize,
|
MaxResponseSize: DefaultMaxResponseSize,
|
||||||
UserAgent: "pixa/1.0",
|
UserAgent: "pixa/1.0",
|
||||||
AllowedContentTypes: []string{
|
AllowedContentTypes: []string{
|
||||||
"image/jpeg",
|
contentTypeJPEG,
|
||||||
"image/png",
|
contentTypePNG,
|
||||||
"image/gif",
|
contentTypeGIF,
|
||||||
"image/webp",
|
contentTypeWebP,
|
||||||
"image/avif",
|
contentTypeAVIF,
|
||||||
"image/svg+xml",
|
contentTypeSVG,
|
||||||
},
|
},
|
||||||
AllowHTTP: false,
|
AllowHTTP: false,
|
||||||
MaxConnectionsPerHost: DefaultMaxConnectionsPerHost,
|
MaxConnectionsPerHost: DefaultMaxConnectionsPerHost,
|
||||||
@@ -132,10 +156,12 @@ func New(config *Config) *HTTPFetcher {
|
|||||||
// Don't follow redirects automatically - we need to validate each hop
|
// Don't follow redirects automatically - we need to validate each hop
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||||
if len(via) >= DefaultMaxRedirects {
|
if len(via) >= DefaultMaxRedirects {
|
||||||
return errors.New("too many redirects")
|
return errTooManyRedirects
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate the redirect target
|
// Validate the redirect target
|
||||||
if err := validateURL(req.URL.String(), config.AllowHTTP); err != nil {
|
err := validateURL(req.Context(), req.URL.String(), config.AllowHTTP)
|
||||||
|
if err != nil {
|
||||||
return fmt.Errorf("redirect blocked: %w", err)
|
return fmt.Errorf("redirect blocked: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,24 +176,11 @@ func New(config *Config) *HTTPFetcher {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// getHostSemaphore returns the semaphore for a host, creating it if necessary.
|
|
||||||
func (f *HTTPFetcher) getHostSemaphore(host string) chan struct{} {
|
|
||||||
f.hostSemMu.Lock()
|
|
||||||
defer f.hostSemMu.Unlock()
|
|
||||||
|
|
||||||
sem, ok := f.hostSems[host]
|
|
||||||
if !ok {
|
|
||||||
sem = make(chan struct{}, f.config.MaxConnectionsPerHost)
|
|
||||||
f.hostSems[host] = sem
|
|
||||||
}
|
|
||||||
|
|
||||||
return sem
|
|
||||||
}
|
|
||||||
|
|
||||||
// Fetch retrieves content from the given URL with SSRF protection.
|
// Fetch retrieves content from the given URL with SSRF protection.
|
||||||
func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, error) {
|
func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, error) {
|
||||||
// Validate URL before making request
|
// Validate URL before making request
|
||||||
if err := validateURL(url, f.config.AllowHTTP); err != nil {
|
err := validateURL(ctx, url, f.config.AllowHTTP)
|
||||||
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -201,7 +214,6 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro
|
|||||||
URL: parsedURL,
|
URL: parsedURL,
|
||||||
Header: make(http.Header),
|
Header: make(http.Header),
|
||||||
}
|
}
|
||||||
req = req.WithContext(ctx)
|
|
||||||
|
|
||||||
req.Header.Set("User-Agent", f.config.UserAgent)
|
req.Header.Set("User-Agent", f.config.UserAgent)
|
||||||
req.Header.Set("Accept", strings.Join(f.config.AllowedContentTypes, ", "))
|
req.Header.Set("Accept", strings.Join(f.config.AllowedContentTypes, ", "))
|
||||||
@@ -216,11 +228,10 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
req = req.WithContext(httptrace.WithClientTrace(req.Context(), trace))
|
req = req.WithContext(httptrace.WithClientTrace(ctx, trace))
|
||||||
|
|
||||||
startTime := time.Now()
|
startTime := time.Now()
|
||||||
|
|
||||||
//nolint:gosec // G704: URL validated by validateURL() above
|
|
||||||
resp, err := f.client.Do(req)
|
resp, err := f.client.Do(req)
|
||||||
|
|
||||||
fetchDuration := time.Since(startTime)
|
fetchDuration := time.Since(startTime)
|
||||||
@@ -233,6 +244,39 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro
|
|||||||
return nil, fmt.Errorf("upstream request failed: %w", err)
|
return nil, fmt.Errorf("upstream request failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
result, err := f.buildResult(resp, remoteAddr, fetchDuration, sem)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark success so defer doesn't release the semaphore
|
||||||
|
success = true
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// getHostSemaphore returns the semaphore for a host, creating it if necessary.
|
||||||
|
func (f *HTTPFetcher) getHostSemaphore(host string) chan struct{} {
|
||||||
|
f.hostSemMu.Lock()
|
||||||
|
defer f.hostSemMu.Unlock()
|
||||||
|
|
||||||
|
sem, ok := f.hostSems[host]
|
||||||
|
if !ok {
|
||||||
|
sem = make(chan struct{}, f.config.MaxConnectionsPerHost)
|
||||||
|
f.hostSems[host] = sem
|
||||||
|
}
|
||||||
|
|
||||||
|
return sem
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildResult validates the upstream response and assembles a FetchResult
|
||||||
|
// whose Content releases the host semaphore slot when closed.
|
||||||
|
func (f *HTTPFetcher) buildResult(
|
||||||
|
resp *http.Response,
|
||||||
|
remoteAddr string,
|
||||||
|
fetchDuration time.Duration,
|
||||||
|
sem chan struct{},
|
||||||
|
) (*FetchResult, error) {
|
||||||
// Extract HTTP version (strip "HTTP/" prefix)
|
// Extract HTTP version (strip "HTTP/" prefix)
|
||||||
httpVersion := strings.TrimPrefix(resp.Proto, "HTTP/")
|
httpVersion := strings.TrimPrefix(resp.Proto, "HTTP/")
|
||||||
|
|
||||||
@@ -265,9 +309,6 @@ func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro
|
|||||||
remaining: f.config.MaxResponseSize,
|
remaining: f.config.MaxResponseSize,
|
||||||
}
|
}
|
||||||
|
|
||||||
// Mark success so defer doesn't release the semaphore
|
|
||||||
success = true
|
|
||||||
|
|
||||||
return &FetchResult{
|
return &FetchResult{
|
||||||
Content: &semaphoreReleasingReadCloser{limitedBody, resp.Body, sem},
|
Content: &semaphoreReleasingReadCloser{limitedBody, resp.Body, sem},
|
||||||
ContentLength: resp.ContentLength,
|
ContentLength: resp.ContentLength,
|
||||||
@@ -297,7 +338,7 @@ func (f *HTTPFetcher) isAllowedContentType(contentType string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// validateURL checks if a URL is safe to fetch (not internal/private).
|
// validateURL checks if a URL is safe to fetch (not internal/private).
|
||||||
func validateURL(rawURL string, allowHTTP bool) error {
|
func validateURL(ctx context.Context, rawURL string, allowHTTP bool) error {
|
||||||
if !allowHTTP && !strings.HasPrefix(rawURL, "https://") {
|
if !allowHTTP && !strings.HasPrefix(rawURL, "https://") {
|
||||||
return ErrUnsupportedScheme
|
return ErrUnsupportedScheme
|
||||||
}
|
}
|
||||||
@@ -309,7 +350,8 @@ func validateURL(rawURL string, allowHTTP bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Remove port if present
|
// Remove port if present
|
||||||
if h, _, err := net.SplitHostPort(host); err == nil {
|
h, _, err := net.SplitHostPort(host)
|
||||||
|
if err == nil {
|
||||||
host = h
|
host = h
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -319,15 +361,16 @@ func validateURL(rawURL string, allowHTTP bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Resolve the host to check IP addresses
|
// Resolve the host to check IP addresses
|
||||||
ips, err := net.LookupIP(host)
|
addrs, err := net.DefaultResolver.LookupIPAddr(ctx, host)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("%w: %s", ErrInvalidHost, host)
|
return fmt.Errorf("%w: %s", ErrInvalidHost, host)
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, ip := range ips {
|
private := slices.ContainsFunc(addrs, func(addr net.IPAddr) bool {
|
||||||
if isPrivateIP(ip) {
|
return isPrivateIP(addr.IP)
|
||||||
return ErrSSRFBlocked
|
})
|
||||||
}
|
if private {
|
||||||
|
return ErrSSRFBlocked
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -340,9 +383,11 @@ func extractHost(rawURL string) string {
|
|||||||
if idx := strings.Index(url, "://"); idx != -1 {
|
if idx := strings.Index(url, "://"); idx != -1 {
|
||||||
url = url[idx+3:]
|
url = url[idx+3:]
|
||||||
}
|
}
|
||||||
|
|
||||||
if idx := strings.Index(url, "/"); idx != -1 {
|
if idx := strings.Index(url, "/"); idx != -1 {
|
||||||
url = url[:idx]
|
url = url[:idx]
|
||||||
}
|
}
|
||||||
|
|
||||||
if idx := strings.Index(url, "?"); idx != -1 {
|
if idx := strings.Index(url, "?"); idx != -1 {
|
||||||
url = url[:idx]
|
url = url[:idx]
|
||||||
}
|
}
|
||||||
@@ -355,8 +400,8 @@ func isLocalhost(host string) bool {
|
|||||||
host = strings.ToLower(host)
|
host = strings.ToLower(host)
|
||||||
|
|
||||||
return host == "localhost" ||
|
return host == "localhost" ||
|
||||||
host == "127.0.0.1" ||
|
host == localhostIPv4 ||
|
||||||
host == "::1" ||
|
host == localhostIPv6 ||
|
||||||
host == "[::1]" ||
|
host == "[::1]" ||
|
||||||
strings.HasSuffix(host, ".localhost") ||
|
strings.HasSuffix(host, ".localhost") ||
|
||||||
strings.HasSuffix(host, ".local")
|
strings.HasSuffix(host, ".local")
|
||||||
@@ -422,23 +467,23 @@ func ssrfSafeDialer(ctx context.Context, network, addr string) (net.Conn, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Check all resolved IPs
|
// Check all resolved IPs
|
||||||
for _, ip := range ips {
|
if slices.ContainsFunc(ips, isPrivateIP) {
|
||||||
if isPrivateIP(ip) {
|
return nil, ErrSSRFBlocked
|
||||||
return nil, ErrSSRFBlocked
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connect using the first valid IP
|
// Connect using the first valid IP
|
||||||
var dialer net.Dialer
|
var dialer net.Dialer
|
||||||
|
|
||||||
for _, ip := range ips {
|
for _, ip := range ips {
|
||||||
addr := net.JoinHostPort(ip.String(), port)
|
addr := net.JoinHostPort(ip.String(), port)
|
||||||
|
|
||||||
conn, err := dialer.DialContext(ctx, network, addr)
|
conn, err := dialer.DialContext(ctx, network, addr)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return conn, nil
|
return conn, nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, fmt.Errorf("failed to connect to %s", host)
|
return nil, fmt.Errorf("%w to %s", errConnectFailed, host)
|
||||||
}
|
}
|
||||||
|
|
||||||
// limitedReader wraps a reader and limits the number of bytes read.
|
// limitedReader wraps a reader and limits the number of bytes read.
|
||||||
@@ -465,6 +510,7 @@ func (r *limitedReader) Read(p []byte) (int, error) {
|
|||||||
// semaphoreReleasingReadCloser releases a semaphore slot when closed.
|
// semaphoreReleasingReadCloser releases a semaphore slot when closed.
|
||||||
type semaphoreReleasingReadCloser struct {
|
type semaphoreReleasingReadCloser struct {
|
||||||
*limitedReader
|
*limitedReader
|
||||||
|
|
||||||
closer io.Closer
|
closer io.Closer
|
||||||
sem chan struct{}
|
sem chan struct{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,7 +9,12 @@ import (
|
|||||||
"testing/fstest"
|
"testing/fstest"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// testHost is the hostname used by mock fetch tests.
|
||||||
|
const testHost = "example.com"
|
||||||
|
|
||||||
func TestDefaultConfig(t *testing.T) {
|
func TestDefaultConfig(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
cfg := DefaultConfig()
|
cfg := DefaultConfig()
|
||||||
|
|
||||||
if cfg.Timeout != DefaultFetchTimeout {
|
if cfg.Timeout != DefaultFetchTimeout {
|
||||||
@@ -35,6 +40,8 @@ func TestDefaultConfig(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestNewWithNilConfigUsesDefaults(t *testing.T) {
|
func TestNewWithNilConfigUsesDefaults(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
f := New(nil)
|
f := New(nil)
|
||||||
|
|
||||||
if f == nil {
|
if f == nil {
|
||||||
@@ -51,24 +58,28 @@ func TestNewWithNilConfigUsesDefaults(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIsAllowedContentType(t *testing.T) {
|
func TestIsAllowedContentType(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
f := New(DefaultConfig())
|
f := New(DefaultConfig())
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
contentType string
|
contentType string
|
||||||
want bool
|
want bool
|
||||||
}{
|
}{
|
||||||
{"image/jpeg", true},
|
{contentTypeJPEG, true},
|
||||||
{"image/png", true},
|
{contentTypePNG, true},
|
||||||
{"image/webp", true},
|
{contentTypeWebP, true},
|
||||||
{"image/jpeg; charset=utf-8", true},
|
{"image/jpeg; charset=utf-8", true},
|
||||||
{"IMAGE/JPEG", true},
|
{"IMAGE/JPEG", true},
|
||||||
{"text/html", false},
|
{"text/html", false},
|
||||||
{"application/octet-stream", false},
|
{contentTypeOctetStream, false},
|
||||||
{"", false},
|
{"", false},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.contentType, func(t *testing.T) {
|
t.Run(tc.contentType, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := f.isAllowedContentType(tc.contentType)
|
got := f.isAllowedContentType(tc.contentType)
|
||||||
if got != tc.want {
|
if got != tc.want {
|
||||||
t.Errorf("isAllowedContentType(%q) = %v, want %v", tc.contentType, got, tc.want)
|
t.Errorf("isAllowedContentType(%q) = %v, want %v", tc.contentType, got, tc.want)
|
||||||
@@ -78,20 +89,24 @@ func TestIsAllowedContentType(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestExtractHost(t *testing.T) {
|
func TestExtractHost(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
url string
|
url string
|
||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{"https://example.com/path", "example.com"},
|
{"https://example.com/path", testHost},
|
||||||
{"http://example.com:8080/path", "example.com:8080"},
|
{"http://example.com:8080/path", "example.com:8080"},
|
||||||
{"https://example.com", "example.com"},
|
{"https://example.com", testHost},
|
||||||
{"https://example.com?q=1", "example.com"},
|
{"https://example.com?q=1", testHost},
|
||||||
{"example.com/path", "example.com"},
|
{"example.com/path", testHost},
|
||||||
{"", ""},
|
{"", ""},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.url, func(t *testing.T) {
|
t.Run(tc.url, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := extractHost(tc.url)
|
got := extractHost(tc.url)
|
||||||
if got != tc.want {
|
if got != tc.want {
|
||||||
t.Errorf("extractHost(%q) = %q, want %q", tc.url, got, tc.want)
|
t.Errorf("extractHost(%q) = %q, want %q", tc.url, got, tc.want)
|
||||||
@@ -101,23 +116,27 @@ func TestExtractHost(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIsLocalhost(t *testing.T) {
|
func TestIsLocalhost(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
host string
|
host string
|
||||||
want bool
|
want bool
|
||||||
}{
|
}{
|
||||||
{"localhost", true},
|
{"localhost", true},
|
||||||
{"LOCALHOST", true},
|
{"LOCALHOST", true},
|
||||||
{"127.0.0.1", true},
|
{localhostIPv4, true},
|
||||||
{"::1", true},
|
{localhostIPv6, true},
|
||||||
{"[::1]", true},
|
{"[::1]", true},
|
||||||
{"foo.localhost", true},
|
{"foo.localhost", true},
|
||||||
{"foo.local", true},
|
{"foo.local", true},
|
||||||
{"example.com", false},
|
{testHost, false},
|
||||||
{"127.0.0.2", false}, // Handled by isPrivateIP, not isLocalhost string match
|
{"127.0.0.2", false}, // Handled by isPrivateIP, not isLocalhost string match
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.host, func(t *testing.T) {
|
t.Run(tc.host, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := isLocalhost(tc.host)
|
got := isLocalhost(tc.host)
|
||||||
if got != tc.want {
|
if got != tc.want {
|
||||||
t.Errorf("isLocalhost(%q) = %v, want %v", tc.host, got, tc.want)
|
t.Errorf("isLocalhost(%q) = %v, want %v", tc.host, got, tc.want)
|
||||||
@@ -127,18 +146,20 @@ func TestIsLocalhost(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIsPrivateIP(t *testing.T) {
|
func TestIsPrivateIP(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
ip string
|
ip string
|
||||||
want bool
|
want bool
|
||||||
}{
|
}{
|
||||||
{"127.0.0.1", true}, // loopback
|
{localhostIPv4, true}, // loopback
|
||||||
{"10.0.0.1", true}, // private
|
{"10.0.0.1", true}, // private
|
||||||
{"192.168.1.1", true}, // private
|
{"192.168.1.1", true}, // private
|
||||||
{"172.16.0.1", true}, // private
|
{"172.16.0.1", true}, // private
|
||||||
{"169.254.1.1", true}, // link-local
|
{"169.254.1.1", true}, // link-local
|
||||||
{"0.0.0.0", true}, // unspecified
|
{"0.0.0.0", true}, // unspecified
|
||||||
{"224.0.0.1", true}, // multicast
|
{"224.0.0.1", true}, // multicast
|
||||||
{"::1", true}, // IPv6 loopback
|
{localhostIPv6, true}, // IPv6 loopback
|
||||||
{"fe80::1", true}, // IPv6 link-local
|
{"fe80::1", true}, // IPv6 link-local
|
||||||
{"8.8.8.8", false}, // public
|
{"8.8.8.8", false}, // public
|
||||||
{"2001:4860:4860::8888", false}, // public IPv6
|
{"2001:4860:4860::8888", false}, // public IPv6
|
||||||
@@ -146,6 +167,8 @@ func TestIsPrivateIP(t *testing.T) {
|
|||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.ip, func(t *testing.T) {
|
t.Run(tc.ip, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
ip := net.ParseIP(tc.ip)
|
ip := net.ParseIP(tc.ip)
|
||||||
if ip == nil {
|
if ip == nil {
|
||||||
t.Fatalf("failed to parse IP %q", tc.ip)
|
t.Fatalf("failed to parse IP %q", tc.ip)
|
||||||
@@ -164,15 +187,19 @@ func TestIsPrivateIP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateURL_RejectsNonHTTPS(t *testing.T) {
|
func TestValidateURL_RejectsNonHTTPS(t *testing.T) {
|
||||||
err := validateURL("http://example.com/path", false)
|
t.Parallel()
|
||||||
|
|
||||||
|
err := validateURL(t.Context(), "http://example.com/path", false)
|
||||||
if !errors.Is(err, ErrUnsupportedScheme) {
|
if !errors.Is(err, ErrUnsupportedScheme) {
|
||||||
t.Errorf("validateURL http = %v, want ErrUnsupportedScheme", err)
|
t.Errorf("validateURL http = %v, want ErrUnsupportedScheme", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateURL_AllowsHTTPWhenConfigured(t *testing.T) {
|
func TestValidateURL_AllowsHTTPWhenConfigured(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Use a host that won't resolve (explicit .invalid TLD) so we don't hit DNS.
|
// Use a host that won't resolve (explicit .invalid TLD) so we don't hit DNS.
|
||||||
err := validateURL("http://nonexistent.invalid/path", true)
|
err := validateURL(t.Context(), "http://nonexistent.invalid/path", true)
|
||||||
// We expect a host resolution error, not ErrUnsupportedScheme.
|
// We expect a host resolution error, not ErrUnsupportedScheme.
|
||||||
if errors.Is(err, ErrUnsupportedScheme) {
|
if errors.Is(err, ErrUnsupportedScheme) {
|
||||||
t.Error("validateURL with AllowHTTP should not return ErrUnsupportedScheme")
|
t.Error("validateURL with AllowHTTP should not return ErrUnsupportedScheme")
|
||||||
@@ -180,20 +207,26 @@ func TestValidateURL_AllowsHTTPWhenConfigured(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateURL_RejectsLocalhost(t *testing.T) {
|
func TestValidateURL_RejectsLocalhost(t *testing.T) {
|
||||||
err := validateURL("https://localhost/path", false)
|
t.Parallel()
|
||||||
|
|
||||||
|
err := validateURL(t.Context(), "https://localhost/path", false)
|
||||||
if !errors.Is(err, ErrSSRFBlocked) {
|
if !errors.Is(err, ErrSSRFBlocked) {
|
||||||
t.Errorf("validateURL localhost = %v, want ErrSSRFBlocked", err)
|
t.Errorf("validateURL localhost = %v, want ErrSSRFBlocked", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateURL_EmptyHost(t *testing.T) {
|
func TestValidateURL_EmptyHost(t *testing.T) {
|
||||||
err := validateURL("https:///path", false)
|
t.Parallel()
|
||||||
|
|
||||||
|
err := validateURL(t.Context(), "https:///path", false)
|
||||||
if !errors.Is(err, ErrInvalidHost) {
|
if !errors.Is(err, ErrInvalidHost) {
|
||||||
t.Errorf("validateURL empty host = %v, want ErrInvalidHost", err)
|
t.Errorf("validateURL empty host = %v, want ErrInvalidHost", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMockFetcher_FetchesFile(t *testing.T) {
|
func TestMockFetcher_FetchesFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
mockFS := fstest.MapFS{
|
mockFS := fstest.MapFS{
|
||||||
"example.com/images/photo.jpg": &fstest.MapFile{Data: []byte("fake-jpeg-data")},
|
"example.com/images/photo.jpg": &fstest.MapFile{Data: []byte("fake-jpeg-data")},
|
||||||
}
|
}
|
||||||
@@ -206,7 +239,7 @@ func TestMockFetcher_FetchesFile(t *testing.T) {
|
|||||||
}
|
}
|
||||||
defer func() { _ = result.Content.Close() }()
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
if result.ContentType != "image/jpeg" {
|
if result.ContentType != contentTypeJPEG {
|
||||||
t.Errorf("ContentType = %q, want image/jpeg", result.ContentType)
|
t.Errorf("ContentType = %q, want image/jpeg", result.ContentType)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -225,6 +258,8 @@ func TestMockFetcher_FetchesFile(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestMockFetcher_MissingFileReturnsUpstreamError(t *testing.T) {
|
func TestMockFetcher_MissingFileReturnsUpstreamError(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
mockFS := fstest.MapFS{}
|
mockFS := fstest.MapFS{}
|
||||||
m := NewMock(mockFS)
|
m := NewMock(mockFS)
|
||||||
|
|
||||||
@@ -235,6 +270,8 @@ func TestMockFetcher_MissingFileReturnsUpstreamError(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestMockFetcher_RespectsContextCancellation(t *testing.T) {
|
func TestMockFetcher_RespectsContextCancellation(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
mockFS := fstest.MapFS{
|
mockFS := fstest.MapFS{
|
||||||
"example.com/photo.jpg": &fstest.MapFile{Data: []byte("data")},
|
"example.com/photo.jpg": &fstest.MapFile{Data: []byte("data")},
|
||||||
}
|
}
|
||||||
@@ -250,24 +287,28 @@ func TestMockFetcher_RespectsContextCancellation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDetectContentTypeFromPath(t *testing.T) {
|
func TestDetectContentTypeFromPath(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
path string
|
path string
|
||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{"foo/bar.jpg", "image/jpeg"},
|
{"foo/bar.jpg", contentTypeJPEG},
|
||||||
{"foo/bar.JPG", "image/jpeg"},
|
{"foo/bar.JPG", contentTypeJPEG},
|
||||||
{"foo/bar.jpeg", "image/jpeg"},
|
{"foo/bar.jpeg", contentTypeJPEG},
|
||||||
{"foo/bar.png", "image/png"},
|
{"foo/bar.png", contentTypePNG},
|
||||||
{"foo/bar.gif", "image/gif"},
|
{"foo/bar.gif", contentTypeGIF},
|
||||||
{"foo/bar.webp", "image/webp"},
|
{"foo/bar.webp", contentTypeWebP},
|
||||||
{"foo/bar.avif", "image/avif"},
|
{"foo/bar.avif", contentTypeAVIF},
|
||||||
{"foo/bar.svg", "image/svg+xml"},
|
{"foo/bar.svg", contentTypeSVG},
|
||||||
{"foo/bar.bin", "application/octet-stream"},
|
{"foo/bar.bin", contentTypeOctetStream},
|
||||||
{"foo/bar", "application/octet-stream"},
|
{"foo/bar", contentTypeOctetStream},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tc := range tests {
|
for _, tc := range tests {
|
||||||
t.Run(tc.path, func(t *testing.T) {
|
t.Run(tc.path, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := detectContentTypeFromPath(tc.path)
|
got := detectContentTypeFromPath(tc.path)
|
||||||
if got != tc.want {
|
if got != tc.want {
|
||||||
t.Errorf("detectContentTypeFromPath(%q) = %q, want %q", tc.path, got, tc.want)
|
t.Errorf("detectContentTypeFromPath(%q) = %q, want %q", tc.path, got, tc.want)
|
||||||
@@ -277,6 +318,8 @@ func TestDetectContentTypeFromPath(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestLimitedReader_EnforcesLimit(t *testing.T) {
|
func TestLimitedReader_EnforcesLimit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
src := make([]byte, 100)
|
src := make([]byte, 100)
|
||||||
r := &limitedReader{
|
r := &limitedReader{
|
||||||
reader: &byteReader{data: src},
|
reader: &byteReader{data: src},
|
||||||
@@ -298,10 +341,11 @@ func TestLimitedReader_EnforcesLimit(t *testing.T) {
|
|||||||
total := n
|
total := n
|
||||||
for total < 50 {
|
for total < 50 {
|
||||||
nn, err := r.Read(buf)
|
nn, err := r.Read(buf)
|
||||||
total += nn
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("during drain: %v", err)
|
t.Fatalf("during drain: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
total += nn
|
||||||
}
|
}
|
||||||
|
|
||||||
// Now the limit is exhausted — next read should error.
|
// Now the limit is exhausted — next read should error.
|
||||||
@@ -4,12 +4,14 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// errEmptyURLPath is returned when a mock URL has no usable path.
|
||||||
|
var errEmptyURLPath = errors.New("empty URL path")
|
||||||
|
|
||||||
// MockFetcher implements Fetcher using an embedded filesystem.
|
// MockFetcher implements Fetcher using an embedded filesystem.
|
||||||
// Files are organized as: hostname/path/to/file.ext
|
// Files are organized as: hostname/path/to/file.ext
|
||||||
// URLs like https://example.com/images/photo.jpg map to example.com/images/photo.jpg.
|
// URLs like https://example.com/images/photo.jpg map to example.com/images/photo.jpg.
|
||||||
@@ -59,7 +61,7 @@ func (m *MockFetcher) Fetch(ctx context.Context, url string) (*FetchResult, erro
|
|||||||
contentType := detectContentTypeFromPath(path)
|
contentType := detectContentTypeFromPath(path)
|
||||||
|
|
||||||
return &FetchResult{
|
return &FetchResult{
|
||||||
Content: f.(io.ReadCloser),
|
Content: f,
|
||||||
ContentLength: stat.Size(),
|
ContentLength: stat.Size(),
|
||||||
ContentType: contentType,
|
ContentType: contentType,
|
||||||
Headers: make(http.Header),
|
Headers: make(http.Header),
|
||||||
@@ -86,7 +88,7 @@ func urlToFSPath(rawURL string) (string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if url == "" {
|
if url == "" {
|
||||||
return "", errors.New("empty URL path")
|
return "", errEmptyURLPath
|
||||||
}
|
}
|
||||||
|
|
||||||
return url, nil
|
return url, nil
|
||||||
@@ -98,18 +100,18 @@ func detectContentTypeFromPath(path string) string {
|
|||||||
|
|
||||||
switch {
|
switch {
|
||||||
case strings.HasSuffix(path, ".jpg"), strings.HasSuffix(path, ".jpeg"):
|
case strings.HasSuffix(path, ".jpg"), strings.HasSuffix(path, ".jpeg"):
|
||||||
return "image/jpeg"
|
return contentTypeJPEG
|
||||||
case strings.HasSuffix(path, ".png"):
|
case strings.HasSuffix(path, ".png"):
|
||||||
return "image/png"
|
return contentTypePNG
|
||||||
case strings.HasSuffix(path, ".gif"):
|
case strings.HasSuffix(path, ".gif"):
|
||||||
return "image/gif"
|
return contentTypeGIF
|
||||||
case strings.HasSuffix(path, ".webp"):
|
case strings.HasSuffix(path, ".webp"):
|
||||||
return "image/webp"
|
return contentTypeWebP
|
||||||
case strings.HasSuffix(path, ".avif"):
|
case strings.HasSuffix(path, ".avif"):
|
||||||
return "image/avif"
|
return contentTypeAVIF
|
||||||
case strings.HasSuffix(path, ".svg"):
|
case strings.HasSuffix(path, ".svg"):
|
||||||
return "image/svg+xml"
|
return contentTypeSVG
|
||||||
default:
|
default:
|
||||||
return "application/octet-stream"
|
return contentTypeOctetStream
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -13,7 +13,9 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// vipsOnce ensures vips is initialized exactly once.
|
// vipsOnce ensures vips is initialized exactly once.
|
||||||
var vipsOnce sync.Once //nolint:gochecknoglobals // package-level sync.Once for one-time vips init
|
//
|
||||||
|
//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.
|
||||||
func initVips() {
|
func initVips() {
|
||||||
@@ -96,10 +98,12 @@ const DefaultMaxInputBytes = 50 << 20
|
|||||||
// ErrInputTooLarge is returned when input image dimensions exceed MaxInputDimension.
|
// ErrInputTooLarge is returned when input image dimensions exceed MaxInputDimension.
|
||||||
var ErrInputTooLarge = errors.New("input image dimensions exceed maximum")
|
var ErrInputTooLarge = errors.New("input image dimensions exceed maximum")
|
||||||
|
|
||||||
// ErrInputDataTooLarge is returned when the raw input data exceeds the configured byte limit.
|
// ErrInputDataTooLarge is returned when the raw input data exceeds the
|
||||||
|
// configured byte limit.
|
||||||
var ErrInputDataTooLarge = errors.New("input data exceeds maximum allowed size")
|
var ErrInputDataTooLarge = errors.New("input data exceeds maximum allowed size")
|
||||||
|
|
||||||
// ErrUnsupportedOutputFormat is returned when the requested output format is not supported.
|
// ErrUnsupportedOutputFormat is returned when the requested output format is
|
||||||
|
// not supported.
|
||||||
var ErrUnsupportedOutputFormat = errors.New("unsupported output format")
|
var ErrUnsupportedOutputFormat = errors.New("unsupported output format")
|
||||||
|
|
||||||
// ImageProcessor implements image transformation using libvips via govips.
|
// ImageProcessor implements image transformation using libvips via govips.
|
||||||
@@ -170,25 +174,12 @@ func (p *ImageProcessor) Process(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Determine target dimensions
|
// Determine target dimensions
|
||||||
targetWidth := req.Size.Width
|
targetWidth, targetHeight := targetDimensions(req.Size, origWidth, origHeight)
|
||||||
targetHeight := req.Size.Height
|
|
||||||
|
|
||||||
// Handle dimension calculation
|
|
||||||
if targetWidth == 0 && targetHeight == 0 {
|
|
||||||
// Both are 0: keep original size
|
|
||||||
targetWidth = origWidth
|
|
||||||
targetHeight = origHeight
|
|
||||||
} else if targetWidth == 0 {
|
|
||||||
// Only height specified: calculate width proportionally
|
|
||||||
targetWidth = origWidth * targetHeight / origHeight
|
|
||||||
} else if targetHeight == 0 {
|
|
||||||
// Only width specified: calculate height proportionally
|
|
||||||
targetHeight = origHeight * targetWidth / origWidth
|
|
||||||
}
|
|
||||||
|
|
||||||
// Resize if needed
|
// Resize if needed
|
||||||
if targetWidth != origWidth || targetHeight != origHeight {
|
if targetWidth != origWidth || targetHeight != origHeight {
|
||||||
if err := p.resize(img, targetWidth, targetHeight, req.FitMode); err != nil {
|
err := p.resize(img, targetWidth, targetHeight, req.FitMode)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to resize: %w", err)
|
return nil, fmt.Errorf("failed to resize: %w", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -217,14 +208,42 @@ func (p *ImageProcessor) Process(
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// targetDimensions calculates the output dimensions for a requested size,
|
||||||
|
// scaling proportionally when only one dimension is given and keeping the
|
||||||
|
// original dimensions when both are zero.
|
||||||
|
func targetDimensions(size Size, origWidth, origHeight int) (int, int) {
|
||||||
|
switch {
|
||||||
|
case size.Width == 0 && size.Height == 0:
|
||||||
|
// Both are 0: keep original size
|
||||||
|
return origWidth, origHeight
|
||||||
|
case size.Width == 0:
|
||||||
|
// Only height specified: calculate width proportionally
|
||||||
|
return origWidth * size.Height / origHeight, size.Height
|
||||||
|
case size.Height == 0:
|
||||||
|
// Only width specified: calculate height proportionally
|
||||||
|
return size.Width, origHeight * size.Width / origWidth
|
||||||
|
default:
|
||||||
|
return size.Width, size.Height
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MIME types for the supported image formats.
|
||||||
|
const (
|
||||||
|
mimeJPEG = "image/jpeg"
|
||||||
|
mimePNG = "image/png"
|
||||||
|
mimeGIF = "image/gif"
|
||||||
|
mimeWebP = "image/webp"
|
||||||
|
mimeAVIF = "image/avif"
|
||||||
|
)
|
||||||
|
|
||||||
// SupportedInputFormats returns MIME types this processor can read.
|
// SupportedInputFormats returns MIME types this processor can read.
|
||||||
func (p *ImageProcessor) SupportedInputFormats() []string {
|
func (p *ImageProcessor) SupportedInputFormats() []string {
|
||||||
return []string{
|
return []string{
|
||||||
"image/jpeg",
|
mimeJPEG,
|
||||||
"image/png",
|
mimePNG,
|
||||||
"image/gif",
|
mimeGIF,
|
||||||
"image/webp",
|
mimeWebP,
|
||||||
"image/avif",
|
mimeAVIF,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -243,15 +262,17 @@ func (p *ImageProcessor) SupportedOutputFormats() []Format {
|
|||||||
func FormatToMIME(format Format) string {
|
func FormatToMIME(format Format) string {
|
||||||
switch format {
|
switch format {
|
||||||
case FormatJPEG:
|
case FormatJPEG:
|
||||||
return "image/jpeg"
|
return mimeJPEG
|
||||||
case FormatPNG:
|
case FormatPNG:
|
||||||
return "image/png"
|
return mimePNG
|
||||||
case FormatWebP:
|
case FormatWebP:
|
||||||
return "image/webp"
|
return mimeWebP
|
||||||
case FormatGIF:
|
case FormatGIF:
|
||||||
return "image/gif"
|
return mimeGIF
|
||||||
case FormatAVIF:
|
case FormatAVIF:
|
||||||
return "image/avif"
|
return mimeAVIF
|
||||||
|
case FormatOriginal:
|
||||||
|
return "application/octet-stream"
|
||||||
default:
|
default:
|
||||||
return "application/octet-stream"
|
return "application/octet-stream"
|
||||||
}
|
}
|
||||||
@@ -270,14 +291,20 @@ func (p *ImageProcessor) detectFormat(img *vips.ImageRef) string {
|
|||||||
case vips.ImageTypeWEBP:
|
case vips.ImageTypeWEBP:
|
||||||
return "webp"
|
return "webp"
|
||||||
case vips.ImageTypeAVIF, vips.ImageTypeHEIF:
|
case vips.ImageTypeAVIF, vips.ImageTypeHEIF:
|
||||||
return "avif"
|
return string(FormatAVIF)
|
||||||
|
case vips.ImageTypeUnknown, vips.ImageTypeMagick, vips.ImageTypePDF,
|
||||||
|
vips.ImageTypeSVG, vips.ImageTypeTIFF, vips.ImageTypeBMP,
|
||||||
|
vips.ImageTypeJP2K, vips.ImageTypeJXL:
|
||||||
|
return "unknown"
|
||||||
default:
|
default:
|
||||||
return "unknown"
|
return "unknown"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// resize resizes the image according to the fit mode.
|
// resize resizes the image according to the fit mode.
|
||||||
func (p *ImageProcessor) resize(img *vips.ImageRef, width, height int, fit FitMode) error {
|
func (p *ImageProcessor) resize(
|
||||||
|
img *vips.ImageRef, width, height int, fit FitMode,
|
||||||
|
) error {
|
||||||
switch fit {
|
switch fit {
|
||||||
case FitCover, "":
|
case FitCover, "":
|
||||||
// Resize and crop to fill exact dimensions (default)
|
// Resize and crop to fill exact dimensions (default)
|
||||||
@@ -303,6 +330,7 @@ func (p *ImageProcessor) resize(img *vips.ImageRef, width, height int, fit FitMo
|
|||||||
if img.Width() <= width && img.Height() <= height {
|
if img.Width() <= width && img.Height() <= height {
|
||||||
return nil // Already fits
|
return nil // Already fits
|
||||||
}
|
}
|
||||||
|
|
||||||
imgW, imgH := img.Width(), img.Height()
|
imgW, imgH := img.Width(), img.Height()
|
||||||
scaleW := float64(width) / float64(imgW)
|
scaleW := float64(width) / float64(imgW)
|
||||||
scaleH := float64(height) / float64(imgH)
|
scaleH := float64(height) / float64(imgH)
|
||||||
@@ -331,7 +359,9 @@ func (p *ImageProcessor) resize(img *vips.ImageRef, width, height int, fit FitMo
|
|||||||
const defaultQuality = 85
|
const defaultQuality = 85
|
||||||
|
|
||||||
// encode encodes an image to the specified format.
|
// encode encodes an image to the specified format.
|
||||||
func (p *ImageProcessor) encode(img *vips.ImageRef, format Format, quality int) ([]byte, error) {
|
func (p *ImageProcessor) encode(
|
||||||
|
img *vips.ImageRef, format Format, quality int,
|
||||||
|
) ([]byte, error) {
|
||||||
if quality <= 0 {
|
if quality <= 0 {
|
||||||
quality = defaultQuality
|
quality = defaultQuality
|
||||||
}
|
}
|
||||||
@@ -367,8 +397,11 @@ func (p *ImageProcessor) encode(img *vips.ImageRef, format Format, quality int)
|
|||||||
Quality: quality,
|
Quality: quality,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
case FormatOriginal:
|
||||||
|
return nil, fmt.Errorf("%w: %s", ErrUnsupportedOutputFormat, format)
|
||||||
|
|
||||||
default:
|
default:
|
||||||
return nil, fmt.Errorf("unsupported output format: %s", format)
|
return nil, fmt.Errorf("%w: %s", ErrUnsupportedOutputFormat, format)
|
||||||
}
|
}
|
||||||
|
|
||||||
output, _, err := img.Export(¶ms)
|
output, _, err := img.Export(¶ms)
|
||||||
@@ -390,7 +423,7 @@ func (p *ImageProcessor) formatFromString(format string) Format {
|
|||||||
return FormatGIF
|
return FormatGIF
|
||||||
case "webp":
|
case "webp":
|
||||||
return FormatWebP
|
return FormatWebP
|
||||||
case "avif":
|
case string(FormatAVIF):
|
||||||
return FormatAVIF
|
return FormatAVIF
|
||||||
default:
|
default:
|
||||||
return FormatJPEG
|
return FormatJPEG
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package imageprocessor
|
|||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"image"
|
"image"
|
||||||
"image/color"
|
"image/color"
|
||||||
"image/jpeg"
|
"image/jpeg"
|
||||||
@@ -16,7 +17,9 @@ import (
|
|||||||
|
|
||||||
func TestMain(m *testing.M) {
|
func TestMain(m *testing.M) {
|
||||||
initVips()
|
initVips()
|
||||||
|
|
||||||
code := m.Run()
|
code := m.Run()
|
||||||
|
|
||||||
vips.Shutdown()
|
vips.Shutdown()
|
||||||
os.Exit(code)
|
os.Exit(code)
|
||||||
}
|
}
|
||||||
@@ -27,11 +30,11 @@ func createTestJPEG(t *testing.T, width, height int) []byte {
|
|||||||
|
|
||||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||||
// Fill with a gradient
|
// Fill with a gradient
|
||||||
for y := 0; y < height; y++ {
|
for y := range height {
|
||||||
for x := 0; x < width; x++ {
|
for x := range width {
|
||||||
img.Set(x, y, color.RGBA{
|
img.Set(x, y, color.RGBA{
|
||||||
R: uint8(x * 255 / width),
|
R: uint8((x * 255 / width) & 0xff),
|
||||||
G: uint8(y * 255 / height),
|
G: uint8((y * 255 / height) & 0xff),
|
||||||
B: 128,
|
B: 128,
|
||||||
A: 255,
|
A: 255,
|
||||||
})
|
})
|
||||||
@@ -39,7 +42,9 @@ func createTestJPEG(t *testing.T, width, height int) []byte {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90}); err != nil {
|
|
||||||
|
err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90})
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test JPEG: %v", err)
|
t.Fatalf("failed to encode test JPEG: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -51,11 +56,11 @@ func createTestPNG(t *testing.T, width, height int) []byte {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||||
for y := 0; y < height; y++ {
|
for y := range height {
|
||||||
for x := 0; x < width; x++ {
|
for x := range width {
|
||||||
img.Set(x, y, color.RGBA{
|
img.Set(x, y, color.RGBA{
|
||||||
R: uint8(x * 255 / width),
|
R: uint8((x * 255 / width) & 0xff),
|
||||||
G: uint8(y * 255 / height),
|
G: uint8((y * 255 / height) & 0xff),
|
||||||
B: 128,
|
B: 128,
|
||||||
A: 255,
|
A: 255,
|
||||||
})
|
})
|
||||||
@@ -63,37 +68,54 @@ func createTestPNG(t *testing.T, width, height int) []byte {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := png.Encode(&buf, img); err != nil {
|
|
||||||
|
err := png.Encode(&buf, img)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test PNG: %v", err)
|
t.Fatalf("failed to encode test PNG: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return buf.Bytes()
|
return buf.Bytes()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isAVIF reports whether data starts with an AVIF ftyp box.
|
||||||
|
func isAVIF(data []byte) bool {
|
||||||
|
if len(data) < 12 || string(data[4:8]) != "ftyp" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
brand := string(data[8:12])
|
||||||
|
|
||||||
|
return brand == string(FormatAVIF) || brand == "avis"
|
||||||
|
}
|
||||||
|
|
||||||
// detectMIME is a minimal magic-byte detector for test assertions.
|
// detectMIME is a minimal magic-byte detector for test assertions.
|
||||||
func detectMIME(data []byte) string {
|
func detectMIME(data []byte) string {
|
||||||
if len(data) >= 3 && data[0] == 0xFF && data[1] == 0xD8 && data[2] == 0xFF {
|
if len(data) >= 3 && data[0] == 0xFF && data[1] == 0xD8 && data[2] == 0xFF {
|
||||||
return "image/jpeg"
|
return mimeJPEG
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(data) >= 8 && string(data[:8]) == "\x89PNG\r\n\x1a\n" {
|
if len(data) >= 8 && string(data[:8]) == "\x89PNG\r\n\x1a\n" {
|
||||||
return "image/png"
|
return mimePNG
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(data) >= 4 && string(data[:4]) == "GIF8" {
|
if len(data) >= 4 && string(data[:4]) == "GIF8" {
|
||||||
return "image/gif"
|
return mimeGIF
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(data) >= 12 && string(data[:4]) == "RIFF" && string(data[8:12]) == "WEBP" {
|
if len(data) >= 12 && string(data[:4]) == "RIFF" && string(data[8:12]) == "WEBP" {
|
||||||
return "image/webp"
|
return mimeWebP
|
||||||
}
|
}
|
||||||
if len(data) >= 12 && string(data[4:8]) == "ftyp" {
|
|
||||||
brand := string(data[8:12])
|
if isAVIF(data) {
|
||||||
if brand == "avif" || brand == "avis" {
|
return mimeAVIF
|
||||||
return "image/avif"
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_ResizeJPEG(t *testing.T) {
|
func TestImageProcessor_ResizeJPEG(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -110,7 +132,8 @@ func TestImageProcessor_ResizeJPEG(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v", err)
|
t.Fatalf("Process() error = %v", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
if result.Width != 400 {
|
if result.Width != 400 {
|
||||||
t.Errorf("Process() width = %d, want 400", result.Width)
|
t.Errorf("Process() width = %d, want 400", result.Width)
|
||||||
@@ -131,12 +154,14 @@ func TestImageProcessor_ResizeJPEG(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mime := detectMIME(data)
|
mime := detectMIME(data)
|
||||||
if mime != "image/jpeg" {
|
if mime != mimeJPEG {
|
||||||
t.Errorf("Output format = %v, want image/jpeg", mime)
|
t.Errorf("Output format = %v, want image/jpeg", mime)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_ConvertToPNG(t *testing.T) {
|
func TestImageProcessor_ConvertToPNG(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -152,7 +177,8 @@ func TestImageProcessor_ConvertToPNG(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v", err)
|
t.Fatalf("Process() error = %v", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
data, err := io.ReadAll(result.Content)
|
data, err := io.ReadAll(result.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -160,19 +186,25 @@ func TestImageProcessor_ConvertToPNG(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mime := detectMIME(data)
|
mime := detectMIME(data)
|
||||||
if mime != "image/png" {
|
if mime != mimePNG {
|
||||||
t.Errorf("Output format = %v, want image/png", mime)
|
t.Errorf("Output format = %v, want image/png", mime)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_OriginalSize(t *testing.T) {
|
// processAndCheckSize processes a test JPEG of the given input dimensions
|
||||||
|
// with the requested size and asserts the resulting dimensions.
|
||||||
|
func processAndCheckSize(
|
||||||
|
t *testing.T, inputW, inputH int, size Size, wantW, wantH int,
|
||||||
|
) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
input := createTestJPEG(t, 640, 480)
|
input := createTestJPEG(t, inputW, inputH)
|
||||||
|
|
||||||
req := &Request{
|
req := &Request{
|
||||||
Size: Size{Width: 0, Height: 0}, // Original size
|
Size: size,
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
FitMode: FitCover,
|
FitMode: FitCover,
|
||||||
@@ -182,18 +214,28 @@ func TestImageProcessor_OriginalSize(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v", err)
|
t.Fatalf("Process() error = %v", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
|
||||||
if result.Width != 640 {
|
defer func() { _ = result.Content.Close() }()
|
||||||
t.Errorf("Process() width = %d, want 640", result.Width)
|
|
||||||
|
if result.Width != wantW {
|
||||||
|
t.Errorf("Process() width = %d, want %d", result.Width, wantW)
|
||||||
}
|
}
|
||||||
|
|
||||||
if result.Height != 480 {
|
if result.Height != wantH {
|
||||||
t.Errorf("Process() height = %d, want 480", result.Height)
|
t.Errorf("Process() height = %d, want %d", result.Height, wantH)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestImageProcessor_OriginalSize(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
// Width and height 0: keep original size
|
||||||
|
processAndCheckSize(t, 640, 480, Size{Width: 0, Height: 0}, 640, 480)
|
||||||
|
}
|
||||||
|
|
||||||
func TestImageProcessor_FitContain(t *testing.T) {
|
func TestImageProcessor_FitContain(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -212,7 +254,8 @@ func TestImageProcessor_FitContain(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v", err)
|
t.Fatalf("Process() error = %v", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
// With contain, the image should fit within the box
|
// With contain, the image should fit within the box
|
||||||
if result.Width > 400 || result.Height > 400 {
|
if result.Width > 400 || result.Height > 400 {
|
||||||
@@ -221,66 +264,24 @@ func TestImageProcessor_FitContain(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_ProportionalScale_WidthOnly(t *testing.T) {
|
func TestImageProcessor_ProportionalScale_WidthOnly(t *testing.T) {
|
||||||
proc := New(Params{})
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// 800x600 image, request width=400 height=0
|
// 800x600 image, request width=400 height=0
|
||||||
// Should scale proportionally to 400x300
|
// Should scale proportionally to 400x300
|
||||||
input := createTestJPEG(t, 800, 600)
|
processAndCheckSize(t, 800, 600, Size{Width: 400, Height: 0}, 400, 300)
|
||||||
|
|
||||||
req := &Request{
|
|
||||||
Size: Size{Width: 400, Height: 0},
|
|
||||||
Format: FormatJPEG,
|
|
||||||
Quality: 85,
|
|
||||||
FitMode: FitCover,
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := proc.Process(ctx, bytes.NewReader(input), req)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Process() error = %v", err)
|
|
||||||
}
|
|
||||||
defer result.Content.Close()
|
|
||||||
|
|
||||||
if result.Width != 400 {
|
|
||||||
t.Errorf("Process() width = %d, want 400", result.Width)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.Height != 300 {
|
|
||||||
t.Errorf("Process() height = %d, want 300", result.Height)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_ProportionalScale_HeightOnly(t *testing.T) {
|
func TestImageProcessor_ProportionalScale_HeightOnly(t *testing.T) {
|
||||||
proc := New(Params{})
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// 800x600 image, request width=0 height=300
|
// 800x600 image, request width=0 height=300
|
||||||
// Should scale proportionally to 400x300
|
// Should scale proportionally to 400x300
|
||||||
input := createTestJPEG(t, 800, 600)
|
processAndCheckSize(t, 800, 600, Size{Width: 0, Height: 300}, 400, 300)
|
||||||
|
|
||||||
req := &Request{
|
|
||||||
Size: Size{Width: 0, Height: 300},
|
|
||||||
Format: FormatJPEG,
|
|
||||||
Quality: 85,
|
|
||||||
FitMode: FitCover,
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := proc.Process(ctx, bytes.NewReader(input), req)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Process() error = %v", err)
|
|
||||||
}
|
|
||||||
defer result.Content.Close()
|
|
||||||
|
|
||||||
if result.Width != 400 {
|
|
||||||
t.Errorf("Process() width = %d, want 400", result.Width)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result.Height != 300 {
|
|
||||||
t.Errorf("Process() height = %d, want 300", result.Height)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_ProcessPNG(t *testing.T) {
|
func TestImageProcessor_ProcessPNG(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -296,7 +297,8 @@ func TestImageProcessor_ProcessPNG(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v", err)
|
t.Fatalf("Process() error = %v", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
if result.Width != 200 {
|
if result.Width != 200 {
|
||||||
t.Errorf("Process() width = %d, want 200", result.Width)
|
t.Errorf("Process() width = %d, want 200", result.Width)
|
||||||
@@ -308,6 +310,8 @@ func TestImageProcessor_ProcessPNG(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_SupportedFormats(t *testing.T) {
|
func TestImageProcessor_SupportedFormats(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
|
|
||||||
inputFormats := proc.SupportedInputFormats()
|
inputFormats := proc.SupportedInputFormats()
|
||||||
@@ -322,55 +326,49 @@ func TestImageProcessor_SupportedFormats(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_RejectsOversizedInput(t *testing.T) {
|
func TestImageProcessor_RejectsOversizedInput(t *testing.T) {
|
||||||
proc := New(Params{})
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// Create an image that exceeds MaxInputDimension (e.g., 10000x100)
|
// Images exceeding MaxInputDimension in either dimension must be
|
||||||
// This should be rejected before processing to prevent DoS
|
// rejected before processing to prevent DoS.
|
||||||
input := createTestJPEG(t, 10000, 100)
|
tests := []struct {
|
||||||
|
name string
|
||||||
req := &Request{
|
width int
|
||||||
Size: Size{Width: 100, Height: 100},
|
height int
|
||||||
Format: FormatJPEG,
|
}{
|
||||||
Quality: 85,
|
{name: "oversized width", width: 10000, height: 100},
|
||||||
FitMode: FitCover,
|
{name: "oversized height", width: 100, height: 10000},
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := proc.Process(ctx, bytes.NewReader(input), req)
|
for _, tt := range tests {
|
||||||
if err == nil {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
t.Error("Process() should reject oversized input images")
|
t.Parallel()
|
||||||
}
|
|
||||||
|
|
||||||
if err != ErrInputTooLarge {
|
proc := New(Params{})
|
||||||
t.Errorf("Process() error = %v, want ErrInputTooLarge", err)
|
ctx := context.Background()
|
||||||
}
|
input := createTestJPEG(t, tt.width, tt.height)
|
||||||
}
|
|
||||||
|
|
||||||
func TestImageProcessor_RejectsOversizedInputHeight(t *testing.T) {
|
req := &Request{
|
||||||
proc := New(Params{})
|
Size: Size{Width: 100, Height: 100},
|
||||||
ctx := context.Background()
|
Format: FormatJPEG,
|
||||||
|
Quality: 85,
|
||||||
|
FitMode: FitCover,
|
||||||
|
}
|
||||||
|
|
||||||
// Create an image with oversized height
|
_, err := proc.Process(ctx, bytes.NewReader(input), req)
|
||||||
input := createTestJPEG(t, 100, 10000)
|
if err == nil {
|
||||||
|
t.Error("Process() should reject oversized input images")
|
||||||
|
}
|
||||||
|
|
||||||
req := &Request{
|
if !errors.Is(err, ErrInputTooLarge) {
|
||||||
Size: Size{Width: 100, Height: 100},
|
t.Errorf("Process() error = %v, want ErrInputTooLarge", err)
|
||||||
Format: FormatJPEG,
|
}
|
||||||
Quality: 85,
|
})
|
||||||
FitMode: FitCover,
|
|
||||||
}
|
|
||||||
|
|
||||||
_, err := proc.Process(ctx, bytes.NewReader(input), req)
|
|
||||||
if err == nil {
|
|
||||||
t.Error("Process() should reject oversized input images")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != ErrInputTooLarge {
|
|
||||||
t.Errorf("Process() error = %v, want ErrInputTooLarge", err)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_AcceptsMaxDimensionInput(t *testing.T) {
|
func TestImageProcessor_AcceptsMaxDimensionInput(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -386,12 +384,20 @@ func TestImageProcessor_AcceptsMaxDimensionInput(t *testing.T) {
|
|||||||
|
|
||||||
result, err := proc.Process(ctx, bytes.NewReader(input), req)
|
result, err := proc.Process(ctx, bytes.NewReader(input), req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() should accept images at MaxInputDimension, got error: %v", err)
|
t.Fatalf(
|
||||||
|
"Process() should accept images at MaxInputDimension, got error: %v",
|
||||||
|
err,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_EncodeWebP(t *testing.T) {
|
// encodeAndCheck processes a 200x150 test JPEG into a 100x75 output of the
|
||||||
|
// given format and asserts the output MIME type and dimensions.
|
||||||
|
func encodeAndCheck(t *testing.T, format Format, quality int, wantMIME string) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -399,8 +405,8 @@ func TestImageProcessor_EncodeWebP(t *testing.T) {
|
|||||||
|
|
||||||
req := &Request{
|
req := &Request{
|
||||||
Size: Size{Width: 100, Height: 75},
|
Size: Size{Width: 100, Height: 75},
|
||||||
Format: FormatWebP,
|
Format: format,
|
||||||
Quality: 80,
|
Quality: quality,
|
||||||
FitMode: FitCover,
|
FitMode: FitCover,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -408,29 +414,39 @@ func TestImageProcessor_EncodeWebP(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v, want nil", err)
|
t.Fatalf("Process() error = %v, want nil", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
|
||||||
// Verify output is valid WebP
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
|
// Verify output format
|
||||||
data, err := io.ReadAll(result.Content)
|
data, err := io.ReadAll(result.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to read result: %v", err)
|
t.Fatalf("failed to read result: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
mime := detectMIME(data)
|
mime := detectMIME(data)
|
||||||
if mime != "image/webp" {
|
if mime != wantMIME {
|
||||||
t.Errorf("Output format = %v, want image/webp", mime)
|
t.Errorf("Output format = %v, want %v", mime, wantMIME)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify dimensions
|
// Verify dimensions
|
||||||
if result.Width != 100 {
|
if result.Width != 100 {
|
||||||
t.Errorf("Width = %d, want 100", result.Width)
|
t.Errorf("Width = %d, want 100", result.Width)
|
||||||
}
|
}
|
||||||
|
|
||||||
if result.Height != 75 {
|
if result.Height != 75 {
|
||||||
t.Errorf("Height = %d, want 75", result.Height)
|
t.Errorf("Height = %d, want 75", result.Height)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestImageProcessor_EncodeWebP(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
encodeAndCheck(t, FormatWebP, 80, mimeWebP)
|
||||||
|
}
|
||||||
|
|
||||||
func TestImageProcessor_DecodeAVIF(t *testing.T) {
|
func TestImageProcessor_DecodeAVIF(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -452,7 +468,8 @@ func TestImageProcessor_DecodeAVIF(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v, want nil (AVIF decoding should work)", err)
|
t.Fatalf("Process() error = %v, want nil (AVIF decoding should work)", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
|
|
||||||
// Verify output is valid JPEG
|
// Verify output is valid JPEG
|
||||||
data, err := io.ReadAll(result.Content)
|
data, err := io.ReadAll(result.Content)
|
||||||
@@ -461,14 +478,17 @@ func TestImageProcessor_DecodeAVIF(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mime := detectMIME(data)
|
mime := detectMIME(data)
|
||||||
if mime != "image/jpeg" {
|
if mime != mimeJPEG {
|
||||||
t.Errorf("Output format = %v, want image/jpeg", mime)
|
t.Errorf("Output format = %v, want image/jpeg", mime)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_RejectsOversizedInputData(t *testing.T) {
|
func TestImageProcessor_RejectsOversizedInputData(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Create a processor with a very small byte limit
|
// Create a processor with a very small byte limit
|
||||||
const limit = 1024
|
const limit = 1024
|
||||||
|
|
||||||
proc := New(Params{MaxInputBytes: limit})
|
proc := New(Params{MaxInputBytes: limit})
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -490,12 +510,14 @@ func TestImageProcessor_RejectsOversizedInputData(t *testing.T) {
|
|||||||
t.Fatal("Process() should reject input exceeding maxInputBytes")
|
t.Fatal("Process() should reject input exceeding maxInputBytes")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrInputDataTooLarge {
|
if !errors.Is(err, ErrInputDataTooLarge) {
|
||||||
t.Errorf("Process() error = %v, want ErrInputDataTooLarge", err)
|
t.Errorf("Process() error = %v, want ErrInputDataTooLarge", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_AcceptsInputWithinLimit(t *testing.T) {
|
func TestImageProcessor_AcceptsInputWithinLimit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Create a small image and set limit well above its size
|
// Create a small image and set limit well above its size
|
||||||
input := createTestJPEG(t, 10, 10)
|
input := createTestJPEG(t, 10, 10)
|
||||||
limit := int64(len(input)) * 10 // 10× headroom
|
limit := int64(len(input)) * 10 // 10× headroom
|
||||||
@@ -514,10 +536,13 @@ func TestImageProcessor_AcceptsInputWithinLimit(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Process() error = %v, want nil", err)
|
t.Fatalf("Process() error = %v, want nil", err)
|
||||||
}
|
}
|
||||||
defer result.Content.Close()
|
|
||||||
|
defer func() { _ = result.Content.Close() }()
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_DefaultMaxInputBytes(t *testing.T) {
|
func TestImageProcessor_DefaultMaxInputBytes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Passing 0 should use the default
|
// Passing 0 should use the default
|
||||||
proc := New(Params{})
|
proc := New(Params{})
|
||||||
if proc.maxInputBytes != DefaultMaxInputBytes {
|
if proc.maxInputBytes != DefaultMaxInputBytes {
|
||||||
@@ -532,40 +557,7 @@ func TestImageProcessor_DefaultMaxInputBytes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageProcessor_EncodeAVIF(t *testing.T) {
|
func TestImageProcessor_EncodeAVIF(t *testing.T) {
|
||||||
proc := New(Params{})
|
t.Parallel()
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
input := createTestJPEG(t, 200, 150)
|
encodeAndCheck(t, FormatAVIF, 85, mimeAVIF)
|
||||||
|
|
||||||
req := &Request{
|
|
||||||
Size: Size{Width: 100, Height: 75},
|
|
||||||
Format: FormatAVIF,
|
|
||||||
Quality: 85,
|
|
||||||
FitMode: FitCover,
|
|
||||||
}
|
|
||||||
|
|
||||||
result, err := proc.Process(ctx, bytes.NewReader(input), req)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Process() error = %v, want nil (AVIF encoding should work)", err)
|
|
||||||
}
|
|
||||||
defer result.Content.Close()
|
|
||||||
|
|
||||||
// Verify output is valid AVIF
|
|
||||||
data, err := io.ReadAll(result.Content)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to read result: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
mime := detectMIME(data)
|
|
||||||
if mime != "image/avif" {
|
|
||||||
t.Errorf("Output format = %v, want image/avif", mime)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify dimensions
|
|
||||||
if result.Width != 100 {
|
|
||||||
t.Errorf("Width = %d, want 100", result.Width)
|
|
||||||
}
|
|
||||||
if result.Height != 75 {
|
|
||||||
t.Errorf("Height = %d, want 75", result.Height)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
@@ -43,23 +43,30 @@ type Cache struct {
|
|||||||
srcMetadata *MetadataStorage // source metadata by host/path
|
srcMetadata *MetadataStorage // source metadata by host/path
|
||||||
config CacheConfig
|
config CacheConfig
|
||||||
|
|
||||||
// In-memory cache of variant metadata (content type, size) to avoid reading .meta files
|
// In-memory cache of variant metadata (content type, size) to avoid
|
||||||
|
// reading .meta files
|
||||||
metaCache map[VariantKey]variantMeta
|
metaCache map[VariantKey]variantMeta
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewCache creates a new cache instance.
|
// NewCache creates a new cache instance.
|
||||||
func NewCache(db *sql.DB, config CacheConfig) (*Cache, error) {
|
func NewCache(db *sql.DB, config CacheConfig) (*Cache, error) {
|
||||||
srcContent, err := NewContentStorage(filepath.Join(config.StateDir, "cache", "sources"))
|
srcContent, err := NewContentStorage(
|
||||||
|
filepath.Join(config.StateDir, "cache", "sources"),
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create source content storage: %w", err)
|
return nil, fmt.Errorf("failed to create source content storage: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
variants, err := NewVariantStorage(filepath.Join(config.StateDir, "cache", "variants"))
|
variants, err := NewVariantStorage(
|
||||||
|
filepath.Join(config.StateDir, "cache", "variants"),
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create variant storage: %w", err)
|
return nil, fmt.Errorf("failed to create variant storage: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
srcMetadata, err := NewMetadataStorage(filepath.Join(config.StateDir, "cache", "metadata"))
|
srcMetadata, err := NewMetadataStorage(
|
||||||
|
filepath.Join(config.StateDir, "cache", "metadata"),
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create source metadata storage: %w", err)
|
return nil, fmt.Errorf("failed to create source metadata storage: %w", err)
|
||||||
}
|
}
|
||||||
@@ -123,7 +130,11 @@ func (c *Cache) StoreSource(
|
|||||||
|
|
||||||
// Store in database
|
// Store in database
|
||||||
pathHash := HashPath(req.SourcePath + "?" + req.SourceQuery)
|
pathHash := HashPath(req.SourcePath + "?" + req.SourceQuery)
|
||||||
headersJSON, _ := json.Marshal(result.Headers)
|
|
||||||
|
headersJSON, err := json.Marshal(result.Headers)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to marshal response headers: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
_, err = c.db.ExecContext(ctx, `
|
_, err = c.db.ExecContext(ctx, `
|
||||||
INSERT INTO source_content (content_hash, content_type, size_bytes)
|
INSERT INTO source_content (content_hash, content_type, size_bytes)
|
||||||
@@ -166,16 +177,16 @@ func (c *Cache) StoreSource(
|
|||||||
RemoteAddr: result.RemoteAddr,
|
RemoteAddr: result.RemoteAddr,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := c.srcMetadata.Store(req.SourceHost, pathHash, meta); err != nil {
|
// A failure here is non-fatal; the metadata is in the database.
|
||||||
// Non-fatal, we have it in the database
|
_ = c.srcMetadata.Store(req.SourceHost, pathHash, meta)
|
||||||
_ = err
|
|
||||||
}
|
|
||||||
|
|
||||||
return contentHash, nil
|
return contentHash, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// StoreVariant stores a processed variant by its cache key.
|
// StoreVariant stores a processed variant by its cache key.
|
||||||
func (c *Cache) StoreVariant(cacheKey VariantKey, content io.Reader, contentType string) error {
|
func (c *Cache) StoreVariant(
|
||||||
|
cacheKey VariantKey, content io.Reader, contentType string,
|
||||||
|
) error {
|
||||||
_, err := c.variants.Store(cacheKey, content, contentType)
|
_, err := c.variants.Store(cacheKey, content, contentType)
|
||||||
|
|
||||||
return err
|
return err
|
||||||
@@ -183,7 +194,9 @@ func (c *Cache) StoreVariant(cacheKey VariantKey, content io.Reader, contentType
|
|||||||
|
|
||||||
// LookupSource checks if we have cached source content for a request.
|
// LookupSource checks if we have cached source content for a request.
|
||||||
// Returns the content hash and content type if found, or empty values if not.
|
// Returns the content hash and content type if found, or empty values if not.
|
||||||
func (c *Cache) LookupSource(ctx context.Context, req *ImageRequest) (ContentHash, string, error) {
|
func (c *Cache) LookupSource(
|
||||||
|
ctx context.Context, req *ImageRequest,
|
||||||
|
) (ContentHash, string, error) {
|
||||||
var hashStr, contentType string
|
var hashStr, contentType string
|
||||||
|
|
||||||
err := c.db.QueryRowContext(ctx, `
|
err := c.db.QueryRowContext(ctx, `
|
||||||
@@ -210,11 +223,15 @@ func (c *Cache) LookupSource(ctx context.Context, req *ImageRequest) (ContentHas
|
|||||||
}
|
}
|
||||||
|
|
||||||
// StoreNegative stores a negative cache entry for a failed fetch.
|
// StoreNegative stores a negative cache entry for a failed fetch.
|
||||||
func (c *Cache) StoreNegative(ctx context.Context, req *ImageRequest, statusCode int, errMsg string) error {
|
func (c *Cache) StoreNegative(
|
||||||
|
ctx context.Context, req *ImageRequest, statusCode int, errMsg string,
|
||||||
|
) error {
|
||||||
expiresAt := time.Now().UTC().Add(c.config.NegativeTTL)
|
expiresAt := time.Now().UTC().Add(c.config.NegativeTTL)
|
||||||
|
|
||||||
_, err := c.db.ExecContext(ctx, `
|
_, err := c.db.ExecContext(ctx, `
|
||||||
INSERT INTO negative_cache (source_host, source_path, source_query, status_code, error_message, expires_at)
|
INSERT INTO negative_cache
|
||||||
|
(source_host, source_path, source_query, status_code,
|
||||||
|
error_message, expires_at)
|
||||||
VALUES (?, ?, ?, ?, ?, ?)
|
VALUES (?, ?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(source_host, source_path, source_query) DO UPDATE SET
|
ON CONFLICT(source_host, source_path, source_query) DO UPDATE SET
|
||||||
status_code = excluded.status_code,
|
status_code = excluded.status_code,
|
||||||
@@ -229,46 +246,16 @@ func (c *Cache) StoreNegative(ctx context.Context, req *ImageRequest, statusCode
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// checkNegativeCache checks if a request is in the negative cache.
|
|
||||||
func (c *Cache) checkNegativeCache(ctx context.Context, req *ImageRequest) (bool, error) {
|
|
||||||
var expiresAt time.Time
|
|
||||||
|
|
||||||
err := c.db.QueryRowContext(ctx, `
|
|
||||||
SELECT expires_at FROM negative_cache
|
|
||||||
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
|
||||||
`, req.SourceHost, req.SourcePath, req.SourceQuery).Scan(&expiresAt)
|
|
||||||
|
|
||||||
if errors.Is(err, sql.ErrNoRows) {
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return false, fmt.Errorf("failed to check negative cache: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check if expired
|
|
||||||
if time.Now().After(expiresAt) {
|
|
||||||
// Clean up expired entry
|
|
||||||
_, _ = c.db.ExecContext(ctx, `
|
|
||||||
DELETE FROM negative_cache
|
|
||||||
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
|
||||||
`, req.SourceHost, req.SourcePath, req.SourceQuery)
|
|
||||||
|
|
||||||
return false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return true, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetSourceMetadataID returns the source metadata ID for a request.
|
// GetSourceMetadataID returns the source metadata ID for a request.
|
||||||
func (c *Cache) GetSourceMetadataID(ctx context.Context, req *ImageRequest) (int64, error) {
|
func (c *Cache) GetSourceMetadataID(
|
||||||
|
ctx context.Context, req *ImageRequest,
|
||||||
|
) (int64, error) {
|
||||||
var id int64
|
var id int64
|
||||||
|
|
||||||
err := c.db.QueryRowContext(ctx, `
|
err := c.db.QueryRowContext(ctx, `
|
||||||
SELECT id FROM source_metadata
|
SELECT id FROM source_metadata
|
||||||
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
||||||
`, req.SourceHost, req.SourcePath, req.SourceQuery).Scan(&id)
|
`, req.SourceHost, req.SourcePath, req.SourceQuery).Scan(&id)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to get source metadata ID: %w", err)
|
return 0, fmt.Errorf("failed to get source metadata ID: %w", err)
|
||||||
}
|
}
|
||||||
@@ -309,8 +296,12 @@ func (c *Cache) Stats(ctx context.Context) (*CacheStats, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Get actual item count and total size from content tables
|
// Get actual item count and total size from content tables
|
||||||
_ = c.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM request_cache`).Scan(&stats.TotalItems)
|
_ = c.db.QueryRowContext(ctx,
|
||||||
_ = c.db.QueryRowContext(ctx, `SELECT COALESCE(SUM(size_bytes), 0) FROM output_content`).Scan(&stats.TotalSizeBytes)
|
`SELECT COUNT(*) FROM request_cache`,
|
||||||
|
).Scan(&stats.TotalItems)
|
||||||
|
_ = c.db.QueryRowContext(ctx,
|
||||||
|
`SELECT COALESCE(SUM(size_bytes), 0) FROM output_content`,
|
||||||
|
).Scan(&stats.TotalSizeBytes)
|
||||||
|
|
||||||
// Compute hit rate as a ratio
|
// Compute hit rate as a ratio
|
||||||
if stats.HitCount+stats.MissCount > 0 {
|
if stats.HitCount+stats.MissCount > 0 {
|
||||||
@@ -324,11 +315,17 @@ func (c *Cache) Stats(ctx context.Context) (*CacheStats, error) {
|
|||||||
func (c *Cache) IncrementStats(ctx context.Context, hit bool, fetchBytes int64) {
|
func (c *Cache) IncrementStats(ctx context.Context, hit bool, fetchBytes int64) {
|
||||||
if hit {
|
if hit {
|
||||||
_, _ = c.db.ExecContext(ctx, `
|
_, _ = c.db.ExecContext(ctx, `
|
||||||
UPDATE cache_stats SET hit_count = hit_count + 1, last_updated_at = CURRENT_TIMESTAMP WHERE id = 1
|
UPDATE cache_stats
|
||||||
|
SET hit_count = hit_count + 1,
|
||||||
|
last_updated_at = CURRENT_TIMESTAMP
|
||||||
|
WHERE id = 1
|
||||||
`)
|
`)
|
||||||
} else {
|
} else {
|
||||||
_, _ = c.db.ExecContext(ctx, `
|
_, _ = c.db.ExecContext(ctx, `
|
||||||
UPDATE cache_stats SET miss_count = miss_count + 1, last_updated_at = CURRENT_TIMESTAMP WHERE id = 1
|
UPDATE cache_stats
|
||||||
|
SET miss_count = miss_count + 1,
|
||||||
|
last_updated_at = CURRENT_TIMESTAMP
|
||||||
|
WHERE id = 1
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -342,3 +339,36 @@ func (c *Cache) IncrementStats(ctx context.Context, hit bool, fetchBytes int64)
|
|||||||
`, fetchBytes)
|
`, fetchBytes)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// checkNegativeCache checks if a request is in the negative cache.
|
||||||
|
func (c *Cache) checkNegativeCache(
|
||||||
|
ctx context.Context, req *ImageRequest,
|
||||||
|
) (bool, error) {
|
||||||
|
var expiresAt time.Time
|
||||||
|
|
||||||
|
err := c.db.QueryRowContext(ctx, `
|
||||||
|
SELECT expires_at FROM negative_cache
|
||||||
|
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
||||||
|
`, req.SourceHost, req.SourcePath, req.SourceQuery).Scan(&expiresAt)
|
||||||
|
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("failed to check negative cache: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if expired
|
||||||
|
if time.Now().After(expiresAt) {
|
||||||
|
// Clean up expired entry
|
||||||
|
_, _ = c.db.ExecContext(ctx, `
|
||||||
|
DELETE FROM negative_cache
|
||||||
|
WHERE source_host = ? AND source_path = ? AND source_query = ?
|
||||||
|
`, req.SourceHost, req.SourcePath, req.SourceQuery)
|
||||||
|
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -86,14 +86,15 @@ func setupTestDB(t *testing.T) *sql.DB {
|
|||||||
INSERT INTO cache_stats (id) VALUES (1);
|
INSERT INTO cache_stats (id) VALUES (1);
|
||||||
`
|
`
|
||||||
|
|
||||||
if _, err := db.Exec(schema); err != nil {
|
_, err = db.ExecContext(t.Context(), schema)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to create schema: %v", err)
|
t.Fatalf("failed to create schema: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return db
|
return db
|
||||||
}
|
}
|
||||||
|
|
||||||
func setupTestCache(t *testing.T) (*Cache, string) {
|
func setupTestCache(t *testing.T) *Cache {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
@@ -108,16 +109,18 @@ func setupTestCache(t *testing.T) (*Cache, string) {
|
|||||||
t.Fatalf("failed to create cache: %v", err)
|
t.Fatalf("failed to create cache: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return cache, tmpDir
|
return cache
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_LookupMiss(t *testing.T) {
|
func TestCache_LookupMiss(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -139,12 +142,14 @@ func TestCache_LookupMiss(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_StoreAndLookup(t *testing.T) {
|
func TestCache_StoreAndLookup(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -154,11 +159,12 @@ func TestCache_StoreAndLookup(t *testing.T) {
|
|||||||
// Store source content
|
// Store source content
|
||||||
sourceContent := []byte("fake jpeg data")
|
sourceContent := []byte("fake jpeg data")
|
||||||
fetchResult := &httpfetcher.FetchResult{
|
fetchResult := &httpfetcher.FetchResult{
|
||||||
ContentType: "image/jpeg",
|
ContentType: testContentTypeJPEG,
|
||||||
Headers: map[string][]string{"Content-Type": {"image/jpeg"}},
|
Headers: map[string][]string{"Content-Type": {testContentTypeJPEG}},
|
||||||
}
|
}
|
||||||
|
|
||||||
contentHash, err := cache.StoreSource(ctx, req, bytes.NewReader(sourceContent), fetchResult)
|
contentHash, err := cache.StoreSource(
|
||||||
|
ctx, req, bytes.NewReader(sourceContent), fetchResult)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("StoreSource() error = %v", err)
|
t.Fatalf("StoreSource() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -170,6 +176,7 @@ func TestCache_StoreAndLookup(t *testing.T) {
|
|||||||
// Store variant
|
// Store variant
|
||||||
cacheKey := CacheKey(req)
|
cacheKey := CacheKey(req)
|
||||||
outputContent := []byte("fake webp data")
|
outputContent := []byte("fake webp data")
|
||||||
|
|
||||||
err = cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
err = cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("StoreVariant() error = %v", err)
|
t.Fatalf("StoreVariant() error = %v", err)
|
||||||
@@ -195,11 +202,13 @@ func TestCache_StoreAndLookup(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_NegativeCache(t *testing.T) {
|
func TestCache_NegativeCache(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/notfound.jpg",
|
SourcePath: "/photos/notfound.jpg",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -223,6 +232,8 @@ func TestCache_NegativeCache(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_NegativeCacheExpiry(t *testing.T) {
|
func TestCache_NegativeCacheExpiry(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
db := setupTestDB(t)
|
db := setupTestDB(t)
|
||||||
|
|
||||||
@@ -239,7 +250,7 @@ func TestCache_NegativeCacheExpiry(t *testing.T) {
|
|||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/expired.jpg",
|
SourcePath: "/photos/expired.jpg",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -266,11 +277,13 @@ func TestCache_NegativeCacheExpiry(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_VariantLookup(t *testing.T) {
|
func TestCache_VariantLookup(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/variant.jpg",
|
SourcePath: "/photos/variant.jpg",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -281,6 +294,7 @@ func TestCache_VariantLookup(t *testing.T) {
|
|||||||
// Store variant
|
// Store variant
|
||||||
cacheKey := CacheKey(req)
|
cacheKey := CacheKey(req)
|
||||||
outputContent := []byte("output data")
|
outputContent := []byte("output data")
|
||||||
|
|
||||||
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("StoreVariant() error = %v", err)
|
t.Fatalf("StoreVariant() error = %v", err)
|
||||||
@@ -312,11 +326,13 @@ func TestCache_VariantLookup(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_GetVariant_ReturnsContentType(t *testing.T) {
|
func TestCache_GetVariant_ReturnsContentType(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/variantct.jpg",
|
SourcePath: "/photos/variantct.jpg",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -327,6 +343,7 @@ func TestCache_GetVariant_ReturnsContentType(t *testing.T) {
|
|||||||
// Store variant
|
// Store variant
|
||||||
cacheKey := CacheKey(req)
|
cacheKey := CacheKey(req)
|
||||||
outputContent := []byte("output webp data")
|
outputContent := []byte("output webp data")
|
||||||
|
|
||||||
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("StoreVariant() error = %v", err)
|
t.Fatalf("StoreVariant() error = %v", err)
|
||||||
@@ -347,7 +364,8 @@ func TestCache_GetVariant_ReturnsContentType(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetVariant() error = %v", err)
|
t.Fatalf("GetVariant() error = %v", err)
|
||||||
}
|
}
|
||||||
defer reader.Close()
|
|
||||||
|
defer func() { _ = reader.Close() }()
|
||||||
|
|
||||||
if contentType != "image/webp" {
|
if contentType != "image/webp" {
|
||||||
t.Errorf("GetVariant() ContentType = %q, want %q", contentType, "image/webp")
|
t.Errorf("GetVariant() ContentType = %q, want %q", contentType, "image/webp")
|
||||||
@@ -359,11 +377,13 @@ func TestCache_GetVariant_ReturnsContentType(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_GetVariant(t *testing.T) {
|
func TestCache_GetVariant(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/output.jpg",
|
SourcePath: "/photos/output.jpg",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -374,6 +394,7 @@ func TestCache_GetVariant(t *testing.T) {
|
|||||||
// Store variant
|
// Store variant
|
||||||
cacheKey := CacheKey(req)
|
cacheKey := CacheKey(req)
|
||||||
outputContent := []byte("the actual output content")
|
outputContent := []byte("the actual output content")
|
||||||
|
|
||||||
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
err := cache.StoreVariant(cacheKey, bytes.NewReader(outputContent), "image/webp")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("StoreVariant() error = %v", err)
|
t.Fatalf("StoreVariant() error = %v", err)
|
||||||
@@ -390,7 +411,8 @@ func TestCache_GetVariant(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetVariant() error = %v", err)
|
t.Fatalf("GetVariant() error = %v", err)
|
||||||
}
|
}
|
||||||
defer reader.Close()
|
|
||||||
|
defer func() { _ = reader.Close() }()
|
||||||
|
|
||||||
buf := make([]byte, 100)
|
buf := make([]byte, 100)
|
||||||
n, _ := reader.Read(buf)
|
n, _ := reader.Read(buf)
|
||||||
@@ -401,7 +423,9 @@ func TestCache_GetVariant(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_Stats(t *testing.T) {
|
func TestCache_Stats(t *testing.T) {
|
||||||
cache, _ := setupTestCache(t)
|
t.Parallel()
|
||||||
|
|
||||||
|
cache := setupTestCache(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
// Increment some stats
|
// Increment some stats
|
||||||
@@ -424,6 +448,8 @@ func TestCache_Stats(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_CleanExpired(t *testing.T) {
|
func TestCache_CleanExpired(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
db := setupTestDB(t)
|
db := setupTestDB(t)
|
||||||
|
|
||||||
@@ -436,7 +462,8 @@ func TestCache_CleanExpired(t *testing.T) {
|
|||||||
|
|
||||||
// Insert expired negative cache entry directly
|
// Insert expired negative cache entry directly
|
||||||
_, err := db.ExecContext(ctx, `
|
_, err := db.ExecContext(ctx, `
|
||||||
INSERT INTO negative_cache (source_host, source_path, source_query, status_code, expires_at)
|
INSERT INTO negative_cache
|
||||||
|
(source_host, source_path, source_query, status_code, expires_at)
|
||||||
VALUES ('example.com', '/old.jpg', '', 404, datetime('now', '-1 hour'))
|
VALUES ('example.com', '/old.jpg', '', 404, datetime('now', '-1 hour'))
|
||||||
`)
|
`)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -445,7 +472,12 @@ func TestCache_CleanExpired(t *testing.T) {
|
|||||||
|
|
||||||
// Verify it exists
|
// Verify it exists
|
||||||
var count int
|
var count int
|
||||||
db.QueryRowContext(ctx, `SELECT COUNT(*) FROM negative_cache`).Scan(&count)
|
|
||||||
|
err = db.QueryRowContext(ctx, `SELECT COUNT(*) FROM negative_cache`).Scan(&count)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to count negative cache entries: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
if count != 1 {
|
if count != 1 {
|
||||||
t.Fatalf("expected 1 negative cache entry, got %d", count)
|
t.Fatalf("expected 1 negative cache entry, got %d", count)
|
||||||
}
|
}
|
||||||
@@ -457,13 +489,19 @@ func TestCache_CleanExpired(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Verify it's gone
|
// Verify it's gone
|
||||||
db.QueryRowContext(ctx, `SELECT COUNT(*) FROM negative_cache`).Scan(&count)
|
err = db.QueryRowContext(ctx, `SELECT COUNT(*) FROM negative_cache`).Scan(&count)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to count negative cache entries: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
if count != 0 {
|
if count != 0 {
|
||||||
t.Errorf("expected 0 negative cache entries after clean, got %d", count)
|
t.Errorf("expected 0 negative cache entries after clean, got %d", count)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCache_StorageDirectoriesCreated(t *testing.T) {
|
func TestCache_StorageDirectoriesCreated(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
db := setupTestDB(t)
|
db := setupTestDB(t)
|
||||||
|
|
||||||
@@ -483,7 +521,9 @@ func TestCache_StorageDirectoriesCreated(t *testing.T) {
|
|||||||
|
|
||||||
for _, dir := range dirs {
|
for _, dir := range dirs {
|
||||||
path := tmpDir + "/" + dir
|
path := tmpDir + "/" + dir
|
||||||
if _, err := os.Stat(path); os.IsNotExist(err) {
|
|
||||||
|
_, err := os.Stat(path)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
t.Errorf("directory %s was not created", dir)
|
t.Errorf("directory %s was not created", dir)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestSizePercentSafeWithZeroFetchBytes(t *testing.T) {
|
func TestSizePercentSafeWithZeroFetchBytes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Simulate the calculation from processAndStore
|
// Simulate the calculation from processAndStore
|
||||||
fetchBytes := int64(0)
|
fetchBytes := int64(0)
|
||||||
outputSize := int64(100)
|
outputSize := int64(100)
|
||||||
@@ -29,6 +31,8 @@ func TestSizePercentSafeWithZeroFetchBytes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestSizePercentNormalCase(t *testing.T) {
|
func TestSizePercentNormalCase(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
fetchBytes := int64(1000)
|
fetchBytes := int64(1000)
|
||||||
outputSize := int64(500)
|
outputSize := int64(500)
|
||||||
|
|
||||||
@@ -90,6 +90,7 @@ func (r *ImageRequest) SourceURL() string {
|
|||||||
if r.AllowHTTP {
|
if r.AllowHTTP {
|
||||||
scheme = "http"
|
scheme = "http"
|
||||||
}
|
}
|
||||||
|
|
||||||
url := scheme + "://" + r.SourceHost + r.SourcePath
|
url := scheme + "://" + r.SourceHost + r.SourcePath
|
||||||
if r.SourceQuery != "" {
|
if r.SourceQuery != "" {
|
||||||
url += "?" + r.SourceQuery
|
url += "?" + r.SourceQuery
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestNegativeCache_StoreAndCheck(t *testing.T) {
|
func TestNegativeCache_StoreAndCheck(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
db := setupTestDB(t)
|
db := setupTestDB(t)
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
|
|
||||||
@@ -22,7 +24,7 @@ func TestNegativeCache_StoreAndCheck(t *testing.T) {
|
|||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "example.com",
|
SourceHost: testHostExample,
|
||||||
SourcePath: "/missing.jpg",
|
SourcePath: "/missing.jpg",
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -31,6 +33,7 @@ func TestNegativeCache_StoreAndCheck(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if hit {
|
if hit {
|
||||||
t.Error("expected no negative cache hit initially")
|
t.Error("expected no negative cache hit initially")
|
||||||
}
|
}
|
||||||
@@ -46,12 +49,15 @@ func TestNegativeCache_StoreAndCheck(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !hit {
|
if !hit {
|
||||||
t.Error("expected negative cache hit after storing")
|
t.Error("expected negative cache hit after storing")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestNegativeCache_Expired(t *testing.T) {
|
func TestNegativeCache_Expired(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
db := setupTestDB(t)
|
db := setupTestDB(t)
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
|
|
||||||
@@ -66,7 +72,7 @@ func TestNegativeCache_Expired(t *testing.T) {
|
|||||||
|
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "example.com",
|
SourceHost: testHostExample,
|
||||||
SourcePath: "/expired.jpg",
|
SourcePath: "/expired.jpg",
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,12 +90,15 @@ func TestNegativeCache_Expired(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if hit {
|
if hit {
|
||||||
t.Error("expected expired negative cache entry to be a miss")
|
t.Error("expected expired negative cache entry to be a miss")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_ReturnsErrorForNegativeCachedURL(t *testing.T) {
|
func TestService_Get_ReturnsErrorForNegativeCachedURL(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// This test verifies that Service.Get() checks the negative cache
|
// This test verifies that Service.Get() checks the negative cache
|
||||||
// We can't easily test the full pipeline without vips, but we can
|
// We can't easily test the full pipeline without vips, but we can
|
||||||
// verify the error type
|
// verify the error type
|
||||||
@@ -18,7 +18,8 @@ import (
|
|||||||
"sneak.berlin/go/pixa/internal/signature"
|
"sneak.berlin/go/pixa/internal/signature"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Service implements the ImageCache interface, orchestrating cache, fetcher, and processor.
|
// Service implements the ImageCache interface, orchestrating cache,
|
||||||
|
// fetcher, and processor.
|
||||||
type Service struct {
|
type Service struct {
|
||||||
cache *Cache
|
cache *Cache
|
||||||
fetcher httpfetcher.Fetcher
|
fetcher httpfetcher.Fetcher
|
||||||
@@ -46,14 +47,21 @@ type ServiceConfig struct {
|
|||||||
Logger *slog.Logger
|
Logger *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Static errors for service construction and unimplemented operations.
|
||||||
|
var (
|
||||||
|
errCacheRequired = errors.New("cache is required")
|
||||||
|
errSigningKeyRequired = errors.New("signing key is required")
|
||||||
|
errPurgeNotImplemented = errors.New("purge not implemented")
|
||||||
|
)
|
||||||
|
|
||||||
// NewService creates a new image service.
|
// NewService creates a new image service.
|
||||||
func NewService(cfg *ServiceConfig) (*Service, error) {
|
func NewService(cfg *ServiceConfig) (*Service, error) {
|
||||||
if cfg.Cache == nil {
|
if cfg.Cache == nil {
|
||||||
return nil, errors.New("cache is required")
|
return nil, errCacheRequired
|
||||||
}
|
}
|
||||||
|
|
||||||
if cfg.SigningKey == "" {
|
if cfg.SigningKey == "" {
|
||||||
return nil, errors.New("signing key is required")
|
return nil, errSigningKeyRequired
|
||||||
}
|
}
|
||||||
|
|
||||||
// Resolve fetcher config for defaults
|
// Resolve fetcher config for defaults
|
||||||
@@ -83,11 +91,14 @@ func NewService(cfg *ServiceConfig) (*Service, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
maxResponseSize := fetcherCfg.MaxResponseSize
|
maxResponseSize := fetcherCfg.MaxResponseSize
|
||||||
|
processor := imageprocessor.New(
|
||||||
|
imageprocessor.Params{MaxInputBytes: maxResponseSize},
|
||||||
|
)
|
||||||
|
|
||||||
return &Service{
|
return &Service{
|
||||||
cache: cfg.Cache,
|
cache: cfg.Cache,
|
||||||
fetcher: fetcher,
|
fetcher: fetcher,
|
||||||
processor: imageprocessor.New(imageprocessor.Params{MaxInputBytes: maxResponseSize}),
|
processor: processor,
|
||||||
signer: signer,
|
signer: signer,
|
||||||
allowlist: allowlist.New(cfg.Allowlist),
|
allowlist: allowlist.New(cfg.Allowlist),
|
||||||
log: log,
|
log: log,
|
||||||
@@ -109,6 +120,7 @@ func (s *Service) Get(ctx context.Context, req *ImageRequest) (*ImageResponse, e
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Warn("negative cache check failed", "error", err)
|
s.log.Warn("negative cache check failed", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if negHit {
|
if negHit {
|
||||||
s.log.Debug("negative cache hit",
|
s.log.Debug("negative cache hit",
|
||||||
"host", req.SourceHost,
|
"host", req.SourceHost,
|
||||||
@@ -145,6 +157,7 @@ func (s *Service) Get(ctx context.Context, req *ImageRequest) (*ImageResponse, e
|
|||||||
|
|
||||||
// Cache miss - check if we have source content cached
|
// Cache miss - check if we have source content cached
|
||||||
cacheKey := CacheKey(req)
|
cacheKey := CacheKey(req)
|
||||||
|
|
||||||
s.cache.IncrementStats(ctx, false, 0)
|
s.cache.IncrementStats(ctx, false, 0)
|
||||||
|
|
||||||
response, err := s.processFromSourceOrFetch(ctx, req, cacheKey)
|
response, err := s.processFromSourceOrFetch(ctx, req, cacheKey)
|
||||||
@@ -157,6 +170,57 @@ func (s *Service) Get(ctx context.Context, req *ImageRequest) (*ImageResponse, e
|
|||||||
return response, nil
|
return response, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Warm pre-fetches and caches an image without returning it.
|
||||||
|
func (s *Service) Warm(ctx context.Context, req *ImageRequest) error {
|
||||||
|
_, err := s.Get(ctx, req)
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Purge removes a cached image. Purging is not implemented yet.
|
||||||
|
func (s *Service) Purge(_ context.Context, _ *ImageRequest) error {
|
||||||
|
return errPurgeNotImplemented
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stats returns cache statistics.
|
||||||
|
func (s *Service) Stats(ctx context.Context) (*CacheStats, error) {
|
||||||
|
return s.cache.Stats(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateRequest validates the request signature if required.
|
||||||
|
func (s *Service) ValidateRequest(req *ImageRequest) error {
|
||||||
|
// Check if host is allowed (no signature required)
|
||||||
|
sourceURL := req.SourceURL()
|
||||||
|
|
||||||
|
parsedURL, err := url.Parse(sourceURL)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid source URL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.allowlist.IsAllowed(parsedURL) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Signature required for non-allowed hosts
|
||||||
|
return s.signer.Verify(signatureRequest(req))
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateSignedURL generates a signed URL for the given request.
|
||||||
|
func (s *Service) GenerateSignedURL(
|
||||||
|
baseURL string,
|
||||||
|
req *ImageRequest,
|
||||||
|
ttl time.Duration,
|
||||||
|
) (string, error) {
|
||||||
|
sigReq := signatureRequest(req)
|
||||||
|
path, sig, exp := s.signer.GenerateSignedURL(sigReq, ttl)
|
||||||
|
|
||||||
|
// Propagate the generated signature and expiration back onto the request.
|
||||||
|
req.Expires = sigReq.Expires
|
||||||
|
req.Signature = sigReq.Signature
|
||||||
|
|
||||||
|
return fmt.Sprintf("%s%s?sig=%s&exp=%d", baseURL, path, sig, exp), nil
|
||||||
|
}
|
||||||
|
|
||||||
// loadCachedSource attempts to load source content from cache, returning nil
|
// loadCachedSource attempts to load source content from cache, returning nil
|
||||||
// if the cached data is unavailable or exceeds maxResponseSize.
|
// if the cached data is unavailable or exceeds maxResponseSize.
|
||||||
func (s *Service) loadCachedSource(contentHash ContentHash) []byte {
|
func (s *Service) loadCachedSource(contentHash ContentHash) []byte {
|
||||||
@@ -191,7 +255,8 @@ func (s *Service) loadCachedSource(contentHash ContentHash) []byte {
|
|||||||
return data
|
return data
|
||||||
}
|
}
|
||||||
|
|
||||||
// processFromSourceOrFetch processes an image, using cached source content if available.
|
// processFromSourceOrFetch processes an image, using cached source content
|
||||||
|
// if available.
|
||||||
func (s *Service) processFromSourceOrFetch(
|
func (s *Service) processFromSourceOrFetch(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
req *ImageRequest,
|
req *ImageRequest,
|
||||||
@@ -203,8 +268,10 @@ func (s *Service) processFromSourceOrFetch(
|
|||||||
s.log.Warn("source lookup failed", "error", err)
|
s.log.Warn("source lookup failed", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var sourceData []byte
|
var (
|
||||||
var fetchBytes int64
|
sourceData []byte
|
||||||
|
fetchBytes int64
|
||||||
|
)
|
||||||
|
|
||||||
if contentHash != "" {
|
if contentHash != "" {
|
||||||
s.log.Debug("using cached source", "hash", contentHash)
|
s.log.Debug("using cached source", "hash", contentHash)
|
||||||
@@ -258,6 +325,7 @@ func (s *Service) fetchAndProcess(
|
|||||||
|
|
||||||
// Calculate download bitrate
|
// Calculate download bitrate
|
||||||
fetchBytes := int64(len(sourceData))
|
fetchBytes := int64(len(sourceData))
|
||||||
|
|
||||||
var downloadRate string
|
var downloadRate string
|
||||||
|
|
||||||
if fetchResult.FetchDurationMs > 0 {
|
if fetchResult.FetchDurationMs > 0 {
|
||||||
@@ -280,7 +348,8 @@ func (s *Service) fetchAndProcess(
|
|||||||
)
|
)
|
||||||
|
|
||||||
// Validate magic bytes match content type
|
// Validate magic bytes match content type
|
||||||
if err := magic.ValidateMagicBytes(sourceData, fetchResult.ContentType); err != nil {
|
err = magic.ValidateMagicBytes(sourceData, fetchResult.ContentType)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("content validation failed: %w", err)
|
return nil, fmt.Errorf("content validation failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -332,7 +401,8 @@ func (s *Service) processAndStore(
|
|||||||
|
|
||||||
var sizePercent float64
|
var sizePercent float64
|
||||||
if fetchBytes > 0 {
|
if fetchBytes > 0 {
|
||||||
sizePercent = float64(outputSize) / float64(fetchBytes) * 100.0 //nolint:mnd // percentage calculation
|
//nolint:mnd // percentage calculation
|
||||||
|
sizePercent = float64(outputSize) / float64(fetchBytes) * 100.0
|
||||||
}
|
}
|
||||||
|
|
||||||
s.log.Info("image converted",
|
s.log.Info("image converted",
|
||||||
@@ -342,8 +412,10 @@ func (s *Service) processAndStore(
|
|||||||
"dst_format", req.Format,
|
"dst_format", req.Format,
|
||||||
"src_bytes", fetchBytes,
|
"src_bytes", fetchBytes,
|
||||||
"dst_bytes", outputSize,
|
"dst_bytes", outputSize,
|
||||||
"src_dimensions", fmt.Sprintf("%dx%d", processResult.InputWidth, processResult.InputHeight),
|
"src_dimensions", fmt.Sprintf("%dx%d",
|
||||||
"dst_dimensions", fmt.Sprintf("%dx%d", processResult.Width, processResult.Height),
|
processResult.InputWidth, processResult.InputHeight),
|
||||||
|
"dst_dimensions", fmt.Sprintf("%dx%d",
|
||||||
|
processResult.Width, processResult.Height),
|
||||||
"size_ratio", fmt.Sprintf("%.1f%%", sizePercent),
|
"size_ratio", fmt.Sprintf("%.1f%%", sizePercent),
|
||||||
"convert_ms", processDuration.Milliseconds(),
|
"convert_ms", processDuration.Milliseconds(),
|
||||||
"quality", req.Quality,
|
"quality", req.Quality,
|
||||||
@@ -351,7 +423,10 @@ func (s *Service) processAndStore(
|
|||||||
)
|
)
|
||||||
|
|
||||||
// Store variant to cache
|
// Store variant to cache
|
||||||
if err := s.cache.StoreVariant(cacheKey, bytes.NewReader(processedData), processResult.ContentType); err != nil {
|
err = s.cache.StoreVariant(
|
||||||
|
cacheKey, bytes.NewReader(processedData), processResult.ContentType,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
s.log.Warn("failed to store variant", "error", err)
|
s.log.Warn("failed to store variant", "error", err)
|
||||||
// Continue even if caching fails
|
// Continue even if caching fails
|
||||||
}
|
}
|
||||||
@@ -365,58 +440,6 @@ func (s *Service) processAndStore(
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Warm pre-fetches and caches an image without returning it.
|
|
||||||
func (s *Service) Warm(ctx context.Context, req *ImageRequest) error {
|
|
||||||
_, err := s.Get(ctx, req)
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Purge removes a cached image.
|
|
||||||
func (s *Service) Purge(_ context.Context, _ *ImageRequest) error {
|
|
||||||
// TODO: Implement purge
|
|
||||||
return errors.New("purge not implemented")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Stats returns cache statistics.
|
|
||||||
func (s *Service) Stats(ctx context.Context) (*CacheStats, error) {
|
|
||||||
return s.cache.Stats(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ValidateRequest validates the request signature if required.
|
|
||||||
func (s *Service) ValidateRequest(req *ImageRequest) error {
|
|
||||||
// Check if host is allowed (no signature required)
|
|
||||||
sourceURL := req.SourceURL()
|
|
||||||
|
|
||||||
parsedURL, err := url.Parse(sourceURL)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("invalid source URL: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if s.allowlist.IsAllowed(parsedURL) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Signature required for non-allowed hosts
|
|
||||||
return s.signer.Verify(signatureRequest(req))
|
|
||||||
}
|
|
||||||
|
|
||||||
// GenerateSignedURL generates a signed URL for the given request.
|
|
||||||
func (s *Service) GenerateSignedURL(
|
|
||||||
baseURL string,
|
|
||||||
req *ImageRequest,
|
|
||||||
ttl time.Duration,
|
|
||||||
) (string, error) {
|
|
||||||
sigReq := signatureRequest(req)
|
|
||||||
path, sig, exp := s.signer.GenerateSignedURL(sigReq, ttl)
|
|
||||||
|
|
||||||
// Propagate the generated signature and expiration back onto the request.
|
|
||||||
req.Expires = sigReq.Expires
|
|
||||||
req.Signature = sigReq.Signature
|
|
||||||
|
|
||||||
return fmt.Sprintf("%s%s?sig=%s&exp=%d", baseURL, path, sig, exp), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// signatureRequest projects an ImageRequest onto the standalone
|
// signatureRequest projects an ImageRequest onto the standalone
|
||||||
// signature.Request type used by the signature package. This keeps the
|
// signature.Request type used by the signature package. This keeps the
|
||||||
// import edge one-way: imgcache depends on signature, never the reverse.
|
// import edge one-way: imgcache depends on signature, never the reverse.
|
||||||
|
|||||||
@@ -10,13 +10,22 @@ import (
|
|||||||
"sneak.berlin/go/pixa/internal/signature"
|
"sneak.berlin/go/pixa/internal/signature"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Test data literals used repeatedly in this file (goconst).
|
||||||
|
const (
|
||||||
|
testPathPhoto = "/images/photo.jpg"
|
||||||
|
testPathUpload = "/uploads/image.jpg"
|
||||||
|
testSigningKey = "test-signing-key-12345"
|
||||||
|
)
|
||||||
|
|
||||||
func TestService_Get_AllowlistedHost(t *testing.T) {
|
func TestService_Get_AllowlistedHost(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -27,7 +36,8 @@ func TestService_Get_AllowlistedHost(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() error = %v", err)
|
t.Fatalf("Get() error = %v", err)
|
||||||
}
|
}
|
||||||
defer resp.Content.Close()
|
|
||||||
|
defer func() { _ = resp.Content.Close() }()
|
||||||
|
|
||||||
// Verify we got content
|
// Verify we got content
|
||||||
data, err := io.ReadAll(resp.Content)
|
data, err := io.ReadAll(resp.Content)
|
||||||
@@ -39,17 +49,19 @@ func TestService_Get_AllowlistedHost(t *testing.T) {
|
|||||||
t.Error("expected non-empty response")
|
t.Error("expected non-empty response")
|
||||||
}
|
}
|
||||||
|
|
||||||
if resp.ContentType != "image/jpeg" {
|
if resp.ContentType != testContentTypeJPEG {
|
||||||
t.Errorf("ContentType = %q, want %q", resp.ContentType, "image/jpeg")
|
t.Errorf("ContentType = %q, want %q", resp.ContentType, testContentTypeJPEG)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_NonAllowlistedHost_NoSignature(t *testing.T) {
|
func TestService_Get_NonAllowlistedHost_NoSignature(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t, WithSigningKey("test-key"))
|
svc, fixtures := SetupTestService(t, WithSigningKey("test-key"))
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.OtherHost,
|
SourceHost: fixtures.OtherHost,
|
||||||
SourcePath: "/uploads/image.jpg",
|
SourcePath: testPathUpload,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -64,13 +76,15 @@ func TestService_Get_NonAllowlistedHost_NoSignature(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_NonAllowlistedHost_ValidSignature(t *testing.T) {
|
func TestService_Get_NonAllowlistedHost_ValidSignature(t *testing.T) {
|
||||||
signingKey := "test-signing-key-12345"
|
t.Parallel()
|
||||||
|
|
||||||
|
signingKey := testSigningKey
|
||||||
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.OtherHost,
|
SourceHost: fixtures.OtherHost,
|
||||||
SourcePath: "/uploads/image.jpg",
|
SourcePath: testPathUpload,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -93,7 +107,8 @@ func TestService_Get_NonAllowlistedHost_ValidSignature(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() error = %v", err)
|
t.Fatalf("Get() error = %v", err)
|
||||||
}
|
}
|
||||||
defer resp.Content.Close()
|
|
||||||
|
defer func() { _ = resp.Content.Close() }()
|
||||||
|
|
||||||
data, err := io.ReadAll(resp.Content)
|
data, err := io.ReadAll(resp.Content)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -106,12 +121,14 @@ func TestService_Get_NonAllowlistedHost_ValidSignature(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_NonAllowlistedHost_ExpiredSignature(t *testing.T) {
|
func TestService_Get_NonAllowlistedHost_ExpiredSignature(t *testing.T) {
|
||||||
signingKey := "test-signing-key-12345"
|
t.Parallel()
|
||||||
|
|
||||||
|
signingKey := testSigningKey
|
||||||
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.OtherHost,
|
SourceHost: fixtures.OtherHost,
|
||||||
SourcePath: "/uploads/image.jpg",
|
SourcePath: testPathUpload,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -131,12 +148,14 @@ func TestService_Get_NonAllowlistedHost_ExpiredSignature(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_NonAllowlistedHost_InvalidSignature(t *testing.T) {
|
func TestService_Get_NonAllowlistedHost_InvalidSignature(t *testing.T) {
|
||||||
signingKey := "test-signing-key-12345"
|
t.Parallel()
|
||||||
|
|
||||||
|
signingKey := testSigningKey
|
||||||
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
svc, fixtures := SetupTestService(t, WithSigningKey(signingKey))
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.OtherHost,
|
SourceHost: fixtures.OtherHost,
|
||||||
SourcePath: "/uploads/image.jpg",
|
SourcePath: testPathUpload,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -159,6 +178,8 @@ func TestService_Get_NonAllowlistedHost_InvalidSignature(t *testing.T) {
|
|||||||
// signature for one host must not verify for a different host, even
|
// signature for one host must not verify for a different host, even
|
||||||
// if they share a domain suffix.
|
// if they share a domain suffix.
|
||||||
func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
signingKey := "test-signing-key-must-be-32-chars"
|
signingKey := "test-signing-key-must-be-32-chars"
|
||||||
svc, _ := SetupTestService(t,
|
svc, _ := SetupTestService(t,
|
||||||
WithSigningKey(signingKey),
|
WithSigningKey(signingKey),
|
||||||
@@ -169,8 +190,8 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
|
|
||||||
// Sign a request for "cdn.example.com"
|
// Sign a request for "cdn.example.com"
|
||||||
signedReq := &ImageRequest{
|
signedReq := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -181,6 +202,8 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
|
|
||||||
// The original request should pass validation
|
// The original request should pass validation
|
||||||
t.Run("exact host passes", func(t *testing.T) {
|
t.Run("exact host passes", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
err := svc.ValidateRequest(signedReq)
|
err := svc.ValidateRequest(signedReq)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Errorf("ValidateRequest() exact host failed: %v", err)
|
t.Errorf("ValidateRequest() exact host failed: %v", err)
|
||||||
@@ -192,7 +215,7 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
name string
|
name string
|
||||||
host string
|
host string
|
||||||
}{
|
}{
|
||||||
{"parent domain", "example.com"},
|
{"parent domain", testHostExample},
|
||||||
{"sibling subdomain", "images.example.com"},
|
{"sibling subdomain", "images.example.com"},
|
||||||
{"deeper subdomain", "a.cdn.example.com"},
|
{"deeper subdomain", "a.cdn.example.com"},
|
||||||
{"evil suffix domain", "cdn.example.com.evil.com"},
|
{"evil suffix domain", "cdn.example.com.evil.com"},
|
||||||
@@ -201,6 +224,8 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name+" rejected", func(t *testing.T) {
|
t.Run(tt.name+" rejected", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: tt.host,
|
SourceHost: tt.host,
|
||||||
SourcePath: signedReq.SourcePath,
|
SourcePath: signedReq.SourcePath,
|
||||||
@@ -215,7 +240,8 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
|
|
||||||
err := svc.ValidateRequest(req)
|
err := svc.ValidateRequest(req)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Errorf("ValidateRequest() should reject signature for host %q (signed for %q)",
|
t.Errorf(
|
||||||
|
"ValidateRequest() should reject signature for host %q (signed for %q)",
|
||||||
tt.host, signedReq.SourceHost)
|
tt.host, signedReq.SourceHost)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -223,6 +249,8 @@ func TestService_ValidateRequest_SignatureExactHostMatch(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_InvalidFile(t *testing.T) {
|
func TestService_Get_InvalidFile(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -243,6 +271,8 @@ func TestService_Get_InvalidFile(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_NotFound(t *testing.T) {
|
func TestService_Get_NotFound(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -262,6 +292,8 @@ func TestService_Get_NotFound(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_FormatConversion(t *testing.T) {
|
func TestService_Get_FormatConversion(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -273,7 +305,7 @@ func TestService_Get_FormatConversion(t *testing.T) {
|
|||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "JPEG to PNG",
|
name: "JPEG to PNG",
|
||||||
sourcePath: "/images/photo.jpg",
|
sourcePath: testPathPhoto,
|
||||||
outFormat: FormatPNG,
|
outFormat: FormatPNG,
|
||||||
wantMIME: "image/png",
|
wantMIME: "image/png",
|
||||||
},
|
},
|
||||||
@@ -281,7 +313,7 @@ func TestService_Get_FormatConversion(t *testing.T) {
|
|||||||
name: "PNG to JPEG",
|
name: "PNG to JPEG",
|
||||||
sourcePath: "/images/logo.png",
|
sourcePath: "/images/logo.png",
|
||||||
outFormat: FormatJPEG,
|
outFormat: FormatJPEG,
|
||||||
wantMIME: "image/jpeg",
|
wantMIME: testContentTypeJPEG,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "GIF to PNG",
|
name: "GIF to PNG",
|
||||||
@@ -293,6 +325,8 @@ func TestService_Get_FormatConversion(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: tt.sourcePath,
|
SourcePath: tt.sourcePath,
|
||||||
@@ -306,7 +340,8 @@ func TestService_Get_FormatConversion(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() error = %v", err)
|
t.Fatalf("Get() error = %v", err)
|
||||||
}
|
}
|
||||||
defer resp.Content.Close()
|
|
||||||
|
defer func() { _ = resp.Content.Close() }()
|
||||||
|
|
||||||
if resp.ContentType != tt.wantMIME {
|
if resp.ContentType != tt.wantMIME {
|
||||||
t.Errorf("ContentType = %q, want %q", resp.ContentType, tt.wantMIME)
|
t.Errorf("ContentType = %q, want %q", resp.ContentType, tt.wantMIME)
|
||||||
@@ -341,12 +376,14 @@ func TestService_Get_FormatConversion(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_Caching(t *testing.T) {
|
func TestService_Get_Caching(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -367,7 +404,8 @@ func TestService_Get_Caching(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to read first response: %v", err)
|
t.Fatalf("failed to read first response: %v", err)
|
||||||
}
|
}
|
||||||
resp1.Content.Close()
|
|
||||||
|
_ = resp1.Content.Close()
|
||||||
|
|
||||||
// Second request - should be a cache hit
|
// Second request - should be a cache hit
|
||||||
resp2, err := svc.Get(ctx, req)
|
resp2, err := svc.Get(ctx, req)
|
||||||
@@ -383,7 +421,8 @@ func TestService_Get_Caching(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to read second response: %v", err)
|
t.Fatalf("failed to read second response: %v", err)
|
||||||
}
|
}
|
||||||
resp2.Content.Close()
|
|
||||||
|
_ = resp2.Content.Close()
|
||||||
|
|
||||||
// Content should be identical
|
// Content should be identical
|
||||||
if len(data1) != len(data2) {
|
if len(data1) != len(data2) {
|
||||||
@@ -392,6 +431,8 @@ func TestService_Get_Caching(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_DifferentSizes(t *testing.T) {
|
func TestService_Get_DifferentSizes(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -402,12 +443,12 @@ func TestService_Get_DifferentSizes(t *testing.T) {
|
|||||||
{Width: 75, Height: 75},
|
{Width: 75, Height: 75},
|
||||||
}
|
}
|
||||||
|
|
||||||
var responses [][]byte
|
responses := make([][]byte, 0, len(sizes))
|
||||||
|
|
||||||
for _, size := range sizes {
|
for _, size := range sizes {
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: size,
|
Size: size,
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -423,27 +464,31 @@ func TestService_Get_DifferentSizes(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("failed to read response: %v", err)
|
t.Fatalf("failed to read response: %v", err)
|
||||||
}
|
}
|
||||||
resp.Content.Close()
|
|
||||||
|
_ = resp.Content.Close()
|
||||||
|
|
||||||
responses = append(responses, data)
|
responses = append(responses, data)
|
||||||
}
|
}
|
||||||
|
|
||||||
// All responses should be different sizes (different cache entries)
|
// All responses should be different sizes (different cache entries)
|
||||||
for i := 0; i < len(responses)-1; i++ {
|
for i := range len(responses) - 1 {
|
||||||
if len(responses[i]) == len(responses[i+1]) {
|
if len(responses[i]) == len(responses[i+1]) {
|
||||||
// Not necessarily an error, but worth noting
|
// Not necessarily an error, but worth noting
|
||||||
t.Logf("responses %d and %d have same size: %d bytes", i, i+1, len(responses[i]))
|
t.Logf("responses %d and %d have same size: %d bytes",
|
||||||
|
i, i+1, len(responses[i]))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestService_ValidateRequest_NoSigningKey(t *testing.T) {
|
func TestService_ValidateRequest_NoSigningKey(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Service with no signing key - all non-allowlisted requests should fail
|
// Service with no signing key - all non-allowlisted requests should fail
|
||||||
svc, fixtures := SetupTestService(t, WithNoAllowlist())
|
svc, fixtures := SetupTestService(t, WithNoAllowlist())
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.OtherHost,
|
SourceHost: fixtures.OtherHost,
|
||||||
SourcePath: "/uploads/image.jpg",
|
SourcePath: testPathUpload,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -452,11 +497,15 @@ func TestService_ValidateRequest_NoSigningKey(t *testing.T) {
|
|||||||
|
|
||||||
err := svc.ValidateRequest(req)
|
err := svc.ValidateRequest(req)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("ValidateRequest() expected error when no signing key and host not allowlisted")
|
t.Error(
|
||||||
|
"ValidateRequest() expected error when no signing key and host not allowlisted",
|
||||||
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_ContextCancellation(t *testing.T) {
|
func TestService_Get_ContextCancellation(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
@@ -464,7 +513,7 @@ func TestService_Get_ContextCancellation(t *testing.T) {
|
|||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -478,12 +527,14 @@ func TestService_Get_ContextCancellation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_ReturnsETag(t *testing.T) {
|
func TestService_Get_ReturnsETag(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -494,7 +545,8 @@ func TestService_Get_ReturnsETag(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() error = %v", err)
|
t.Fatalf("Get() error = %v", err)
|
||||||
}
|
}
|
||||||
defer resp.Content.Close()
|
|
||||||
|
defer func() { _ = resp.Content.Close() }()
|
||||||
|
|
||||||
// ETag should be set
|
// ETag should be set
|
||||||
if resp.ETag == "" {
|
if resp.ETag == "" {
|
||||||
@@ -508,12 +560,14 @@ func TestService_Get_ReturnsETag(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_ETagConsistency(t *testing.T) {
|
func TestService_Get_ETagConsistency(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -525,16 +579,20 @@ func TestService_Get_ETagConsistency(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() first request error = %v", err)
|
t.Fatalf("Get() first request error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
etag1 := resp1.ETag
|
etag1 := resp1.ETag
|
||||||
resp1.Content.Close()
|
|
||||||
|
_ = resp1.Content.Close()
|
||||||
|
|
||||||
// Second request (from cache)
|
// Second request (from cache)
|
||||||
resp2, err := svc.Get(ctx, req)
|
resp2, err := svc.Get(ctx, req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() second request error = %v", err)
|
t.Fatalf("Get() second request error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
etag2 := resp2.ETag
|
etag2 := resp2.ETag
|
||||||
resp2.Content.Close()
|
|
||||||
|
_ = resp2.Content.Close()
|
||||||
|
|
||||||
// ETags should be identical for the same content
|
// ETags should be identical for the same content
|
||||||
if etag1 != etag2 {
|
if etag1 != etag2 {
|
||||||
@@ -543,13 +601,15 @@ func TestService_Get_ETagConsistency(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestService_Get_DifferentETagsForDifferentContent(t *testing.T) {
|
func TestService_Get_DifferentETagsForDifferentContent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
svc, fixtures := SetupTestService(t)
|
svc, fixtures := SetupTestService(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
// Request same image at different sizes - should get different ETags
|
// Request same image at different sizes - should get different ETags
|
||||||
req1 := &ImageRequest{
|
req1 := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 25, Height: 25},
|
Size: Size{Width: 25, Height: 25},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -558,7 +618,7 @@ func TestService_Get_DifferentETagsForDifferentContent(t *testing.T) {
|
|||||||
|
|
||||||
req2 := &ImageRequest{
|
req2 := &ImageRequest{
|
||||||
SourceHost: fixtures.GoodHost,
|
SourceHost: fixtures.GoodHost,
|
||||||
SourcePath: "/images/photo.jpg",
|
SourcePath: testPathPhoto,
|
||||||
Size: Size{Width: 50, Height: 50},
|
Size: Size{Width: 50, Height: 50},
|
||||||
Format: FormatJPEG,
|
Format: FormatJPEG,
|
||||||
Quality: 85,
|
Quality: 85,
|
||||||
@@ -569,15 +629,19 @@ func TestService_Get_DifferentETagsForDifferentContent(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() first request error = %v", err)
|
t.Fatalf("Get() first request error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
etag1 := resp1.ETag
|
etag1 := resp1.ETag
|
||||||
resp1.Content.Close()
|
|
||||||
|
_ = resp1.Content.Close()
|
||||||
|
|
||||||
resp2, err := svc.Get(ctx, req2)
|
resp2, err := svc.Get(ctx, req2)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Get() second request error = %v", err)
|
t.Fatalf("Get() second request error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
etag2 := resp2.ETag
|
etag2 := resp2.ETag
|
||||||
resp2.Content.Close()
|
|
||||||
|
_ = resp2.Content.Close()
|
||||||
|
|
||||||
// ETags should be different for different content
|
// ETags should be different for different content
|
||||||
if etag1 == etag2 {
|
if etag1 == etag2 {
|
||||||
@@ -3,13 +3,16 @@ package imgcache
|
|||||||
import "testing"
|
import "testing"
|
||||||
|
|
||||||
func TestImageRequest_SourceURL_DefaultHTTPS(t *testing.T) {
|
func TestImageRequest_SourceURL_DefaultHTTPS(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "v=2",
|
SourceQuery: "v=2",
|
||||||
}
|
}
|
||||||
|
|
||||||
got := req.SourceURL()
|
got := req.SourceURL()
|
||||||
|
|
||||||
want := "https://cdn.example.com/photos/cat.jpg?v=2"
|
want := "https://cdn.example.com/photos/cat.jpg?v=2"
|
||||||
if got != want {
|
if got != want {
|
||||||
t.Errorf("SourceURL() = %q, want %q", got, want)
|
t.Errorf("SourceURL() = %q, want %q", got, want)
|
||||||
@@ -17,13 +20,16 @@ func TestImageRequest_SourceURL_DefaultHTTPS(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageRequest_SourceURL_AllowHTTP(t *testing.T) {
|
func TestImageRequest_SourceURL_AllowHTTP(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "localhost:8080",
|
SourceHost: "localhost:8080",
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
AllowHTTP: true,
|
AllowHTTP: true,
|
||||||
}
|
}
|
||||||
|
|
||||||
got := req.SourceURL()
|
got := req.SourceURL()
|
||||||
|
|
||||||
want := "http://localhost:8080/photos/cat.jpg"
|
want := "http://localhost:8080/photos/cat.jpg"
|
||||||
if got != want {
|
if got != want {
|
||||||
t.Errorf("SourceURL() = %q, want %q", got, want)
|
t.Errorf("SourceURL() = %q, want %q", got, want)
|
||||||
@@ -31,8 +37,10 @@ func TestImageRequest_SourceURL_AllowHTTP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageRequest_SourceURL_AllowHTTPFalse(t *testing.T) {
|
func TestImageRequest_SourceURL_AllowHTTPFalse(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := &ImageRequest{
|
req := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/img.jpg",
|
SourcePath: "/img.jpg",
|
||||||
AllowHTTP: false,
|
AllowHTTP: false,
|
||||||
}
|
}
|
||||||
@@ -12,18 +12,25 @@ import (
|
|||||||
|
|
||||||
func setupStatsTestDB(t *testing.T) *sql.DB {
|
func setupStatsTestDB(t *testing.T) *sql.DB {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
db, err := sql.Open("sqlite", ":memory:")
|
db, err := sql.Open("sqlite", ":memory:")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := database.ApplyMigrations(context.Background(), db, nil); err != nil {
|
|
||||||
|
err = database.ApplyMigrations(context.Background(), db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { db.Close() })
|
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
return db
|
return db
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStats_HitRateIsRatio(t *testing.T) {
|
func TestStats_HitRateIsRatio(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
db := setupStatsTestDB(t)
|
db := setupStatsTestDB(t)
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
|
|
||||||
@@ -40,7 +47,9 @@ func TestStats_HitRateIsRatio(t *testing.T) {
|
|||||||
|
|
||||||
// Set some hit/miss counts and a transform_count
|
// Set some hit/miss counts and a transform_count
|
||||||
_, err = db.ExecContext(ctx, `
|
_, err = db.ExecContext(ctx, `
|
||||||
UPDATE cache_stats SET hit_count = 75, miss_count = 25, transform_count = 9999 WHERE id = 1
|
UPDATE cache_stats
|
||||||
|
SET hit_count = 75, miss_count = 25, transform_count = 9999
|
||||||
|
WHERE id = 1
|
||||||
`)
|
`)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -54,6 +63,7 @@ func TestStats_HitRateIsRatio(t *testing.T) {
|
|||||||
if stats.HitCount != 75 {
|
if stats.HitCount != 75 {
|
||||||
t.Errorf("HitCount = %d, want 75", stats.HitCount)
|
t.Errorf("HitCount = %d, want 75", stats.HitCount)
|
||||||
}
|
}
|
||||||
|
|
||||||
if stats.MissCount != 25 {
|
if stats.MissCount != 25 {
|
||||||
t.Errorf("MissCount = %d, want 25", stats.MissCount)
|
t.Errorf("MissCount = %d, want 25", stats.MissCount)
|
||||||
}
|
}
|
||||||
@@ -61,11 +71,14 @@ func TestStats_HitRateIsRatio(t *testing.T) {
|
|||||||
// HitRate should be 0.75, NOT 9999 (transform_count)
|
// HitRate should be 0.75, NOT 9999 (transform_count)
|
||||||
expectedRate := 0.75
|
expectedRate := 0.75
|
||||||
if math.Abs(stats.HitRate-expectedRate) > 0.001 {
|
if math.Abs(stats.HitRate-expectedRate) > 0.001 {
|
||||||
t.Errorf("HitRate = %f, want %f (was it scanning transform_count?)", stats.HitRate, expectedRate)
|
t.Errorf("HitRate = %f, want %f (was it scanning transform_count?)",
|
||||||
|
stats.HitRate, expectedRate)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestStats_ZeroCounts(t *testing.T) {
|
func TestStats_ZeroCounts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
db := setupStatsTestDB(t)
|
db := setupStatsTestDB(t)
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
|
|
||||||
@@ -44,7 +44,8 @@ type ContentStorage struct {
|
|||||||
|
|
||||||
// NewContentStorage creates a new content storage at the given base directory.
|
// NewContentStorage creates a new content storage at the given base directory.
|
||||||
func NewContentStorage(baseDir string) (*ContentStorage, error) {
|
func NewContentStorage(baseDir string) (*ContentStorage, error) {
|
||||||
if err := os.MkdirAll(baseDir, StorageDirPerm); err != nil {
|
err := os.MkdirAll(baseDir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create storage directory: %w", err)
|
return nil, fmt.Errorf("failed to create storage directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -53,7 +54,7 @@ func NewContentStorage(baseDir string) (*ContentStorage, error) {
|
|||||||
|
|
||||||
// Store writes content to storage and returns its SHA256 hash.
|
// Store writes content to storage and returns its SHA256 hash.
|
||||||
// The content is read fully into memory to compute the hash before writing.
|
// The content is read fully into memory to compute the hash before writing.
|
||||||
func (s *ContentStorage) Store(r io.Reader) (hash ContentHash, size int64, err error) {
|
func (s *ContentStorage) Store(r io.Reader) (ContentHash, int64, error) {
|
||||||
// Read all content to compute hash
|
// Read all content to compute hash
|
||||||
data, err := io.ReadAll(r)
|
data, err := io.ReadAll(r)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -62,20 +63,23 @@ func (s *ContentStorage) Store(r io.Reader) (hash ContentHash, size int64, err e
|
|||||||
|
|
||||||
// Compute hash
|
// Compute hash
|
||||||
h := sha256.Sum256(data)
|
h := sha256.Sum256(data)
|
||||||
hash = ContentHash(hex.EncodeToString(h[:]))
|
hash := ContentHash(hex.EncodeToString(h[:]))
|
||||||
size = int64(len(data))
|
size := int64(len(data))
|
||||||
|
|
||||||
// Build path: <basedir>/<ab>/<cd>/<hash>
|
// Build path: <basedir>/<ab>/<cd>/<hash>
|
||||||
path := s.hashToPath(hash)
|
path := s.hashToPath(hash)
|
||||||
|
|
||||||
// Check if already exists
|
// Check if already exists
|
||||||
if _, err := os.Stat(path); err == nil {
|
_, err = os.Stat(path)
|
||||||
|
if err == nil {
|
||||||
return hash, size, nil
|
return hash, size, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create directory structure
|
// Create directory structure
|
||||||
dir := filepath.Dir(path)
|
dir := filepath.Dir(path)
|
||||||
if err := os.MkdirAll(dir, StorageDirPerm); err != nil {
|
|
||||||
|
err = os.MkdirAll(dir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return "", 0, fmt.Errorf("failed to create directory: %w", err)
|
return "", 0, fmt.Errorf("failed to create directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,27 +88,29 @@ func (s *ContentStorage) Store(r io.Reader) (hash ContentHash, size int64, err e
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", 0, fmt.Errorf("failed to create temp file: %w", err)
|
return "", 0, fmt.Errorf("failed to create temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
tmpPath := tmpFile.Name()
|
tmpPath := tmpFile.Name()
|
||||||
|
|
||||||
defer func() {
|
_, err = tmpFile.Write(data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = os.Remove(tmpPath)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if _, err := tmpFile.Write(data); err != nil {
|
|
||||||
_ = tmpFile.Close()
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return "", 0, fmt.Errorf("failed to write content: %w", err)
|
return "", 0, fmt.Errorf("failed to write content: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := tmpFile.Close(); err != nil {
|
err = tmpFile.Close()
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return "", 0, fmt.Errorf("failed to close temp file: %w", err)
|
return "", 0, fmt.Errorf("failed to close temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Atomic rename
|
// Atomic rename
|
||||||
//nolint:gosec // G703: paths from internal SHA256 hashes
|
err = os.Rename(filepath.Clean(tmpPath), filepath.Clean(path))
|
||||||
if err := os.Rename(filepath.Clean(tmpPath), filepath.Clean(path)); err != nil {
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return "", 0, fmt.Errorf("failed to rename temp file: %w", err)
|
return "", 0, fmt.Errorf("failed to rename temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,7 +194,8 @@ type MetadataStorage struct {
|
|||||||
|
|
||||||
// NewMetadataStorage creates a new metadata storage at the given base directory.
|
// NewMetadataStorage creates a new metadata storage at the given base directory.
|
||||||
func NewMetadataStorage(baseDir string) (*MetadataStorage, error) {
|
func NewMetadataStorage(baseDir string) (*MetadataStorage, error) {
|
||||||
if err := os.MkdirAll(baseDir, StorageDirPerm); err != nil {
|
err := os.MkdirAll(baseDir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create metadata directory: %w", err)
|
return nil, fmt.Errorf("failed to create metadata directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,6 +203,8 @@ func NewMetadataStorage(baseDir string) (*MetadataStorage, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// SourceMetadata represents cached metadata about a source URL.
|
// SourceMetadata represents cached metadata about a source URL.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // stored metadata format uses snake_case
|
||||||
type SourceMetadata struct {
|
type SourceMetadata struct {
|
||||||
Host string `json:"host"`
|
Host string `json:"host"`
|
||||||
Path string `json:"path"`
|
Path string `json:"path"`
|
||||||
@@ -214,12 +223,16 @@ type SourceMetadata struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Store writes metadata to storage.
|
// Store writes metadata to storage.
|
||||||
func (s *MetadataStorage) Store(host string, pathHash PathHash, meta *SourceMetadata) error {
|
func (s *MetadataStorage) Store(
|
||||||
|
host string, pathHash PathHash, meta *SourceMetadata,
|
||||||
|
) error {
|
||||||
path := s.metaPath(host, pathHash)
|
path := s.metaPath(host, pathHash)
|
||||||
|
|
||||||
// Create directory structure
|
// Create directory structure
|
||||||
dir := filepath.Dir(path)
|
dir := filepath.Dir(path)
|
||||||
if err := os.MkdirAll(dir, StorageDirPerm); err != nil {
|
|
||||||
|
err := os.MkdirAll(dir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create directory: %w", err)
|
return fmt.Errorf("failed to create directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -234,27 +247,29 @@ func (s *MetadataStorage) Store(host string, pathHash PathHash, meta *SourceMeta
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create temp file: %w", err)
|
return fmt.Errorf("failed to create temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
tmpPath := tmpFile.Name()
|
tmpPath := tmpFile.Name()
|
||||||
|
|
||||||
defer func() {
|
_, err = tmpFile.Write(data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = os.Remove(tmpPath)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if _, err := tmpFile.Write(data); err != nil {
|
|
||||||
_ = tmpFile.Close()
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return fmt.Errorf("failed to write metadata: %w", err)
|
return fmt.Errorf("failed to write metadata: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := tmpFile.Close(); err != nil {
|
err = tmpFile.Close()
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return fmt.Errorf("failed to close temp file: %w", err)
|
return fmt.Errorf("failed to close temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Atomic rename
|
// Atomic rename
|
||||||
//nolint:gosec // G703: paths from internal SHA256 hashes
|
err = os.Rename(filepath.Clean(tmpPath), filepath.Clean(path))
|
||||||
if err := os.Rename(filepath.Clean(tmpPath), filepath.Clean(path)); err != nil {
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return fmt.Errorf("failed to rename temp file: %w", err)
|
return fmt.Errorf("failed to rename temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -262,7 +277,9 @@ func (s *MetadataStorage) Store(host string, pathHash PathHash, meta *SourceMeta
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Load reads metadata from storage.
|
// Load reads metadata from storage.
|
||||||
func (s *MetadataStorage) Load(host string, pathHash PathHash) (*SourceMetadata, error) {
|
func (s *MetadataStorage) Load(
|
||||||
|
host string, pathHash PathHash,
|
||||||
|
) (*SourceMetadata, error) {
|
||||||
path := s.metaPath(host, pathHash)
|
path := s.metaPath(host, pathHash)
|
||||||
|
|
||||||
data, err := os.ReadFile(path) //nolint:gosec // path derived from host+hash
|
data, err := os.ReadFile(path) //nolint:gosec // path derived from host+hash
|
||||||
@@ -275,7 +292,9 @@ func (s *MetadataStorage) Load(host string, pathHash PathHash) (*SourceMetadata,
|
|||||||
}
|
}
|
||||||
|
|
||||||
var meta SourceMetadata
|
var meta SourceMetadata
|
||||||
if err := json.Unmarshal(data, &meta); err != nil {
|
|
||||||
|
err = json.Unmarshal(data, &meta)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to unmarshal metadata: %w", err)
|
return nil, fmt.Errorf("failed to unmarshal metadata: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -341,6 +360,8 @@ type VariantStorage struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// VariantMeta contains metadata about a cached variant.
|
// VariantMeta contains metadata about a cached variant.
|
||||||
|
//
|
||||||
|
//nolint:tagliatelle // stored metadata format uses snake_case
|
||||||
type VariantMeta struct {
|
type VariantMeta struct {
|
||||||
ContentType string `json:"content_type"`
|
ContentType string `json:"content_type"`
|
||||||
Size int64 `json:"size"`
|
Size int64 `json:"size"`
|
||||||
@@ -349,7 +370,8 @@ type VariantMeta struct {
|
|||||||
|
|
||||||
// NewVariantStorage creates a new variant storage at the given base directory.
|
// NewVariantStorage creates a new variant storage at the given base directory.
|
||||||
func NewVariantStorage(baseDir string) (*VariantStorage, error) {
|
func NewVariantStorage(baseDir string) (*VariantStorage, error) {
|
||||||
if err := os.MkdirAll(baseDir, StorageDirPerm); err != nil {
|
err := os.MkdirAll(baseDir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create variant storage directory: %w", err)
|
return nil, fmt.Errorf("failed to create variant storage directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -357,19 +379,23 @@ func NewVariantStorage(baseDir string) (*VariantStorage, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Store writes content and metadata to storage at the given key.
|
// Store writes content and metadata to storage at the given key.
|
||||||
func (s *VariantStorage) Store(key VariantKey, r io.Reader, contentType string) (size int64, err error) {
|
func (s *VariantStorage) Store(
|
||||||
|
key VariantKey, r io.Reader, contentType string,
|
||||||
|
) (int64, error) {
|
||||||
data, err := io.ReadAll(r)
|
data, err := io.ReadAll(r)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to read content: %w", err)
|
return 0, fmt.Errorf("failed to read content: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
size = int64(len(data))
|
size := int64(len(data))
|
||||||
path := s.keyToPath(key)
|
path := s.keyToPath(key)
|
||||||
metaPath := path + ".meta"
|
metaPath := path + ".meta"
|
||||||
|
|
||||||
// Create directory structure
|
// Create directory structure
|
||||||
dir := filepath.Dir(path)
|
dir := filepath.Dir(path)
|
||||||
if err := os.MkdirAll(dir, StorageDirPerm); err != nil {
|
|
||||||
|
err = os.MkdirAll(dir, StorageDirPerm)
|
||||||
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to create directory: %w", err)
|
return 0, fmt.Errorf("failed to create directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -378,27 +404,29 @@ func (s *VariantStorage) Store(key VariantKey, r io.Reader, contentType string)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to create temp file: %w", err)
|
return 0, fmt.Errorf("failed to create temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
tmpPath := tmpFile.Name()
|
tmpPath := tmpFile.Name()
|
||||||
|
|
||||||
defer func() {
|
_, err = tmpFile.Write(data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = os.Remove(tmpPath)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if _, err := tmpFile.Write(data); err != nil {
|
|
||||||
_ = tmpFile.Close()
|
_ = tmpFile.Close()
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return 0, fmt.Errorf("failed to write content: %w", err)
|
return 0, fmt.Errorf("failed to write content: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := tmpFile.Close(); err != nil {
|
err = tmpFile.Close()
|
||||||
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return 0, fmt.Errorf("failed to close temp file: %w", err)
|
return 0, fmt.Errorf("failed to close temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Atomic rename content
|
// Atomic rename content
|
||||||
//nolint:gosec // G703: paths from internal SHA256 hashes
|
err = os.Rename(filepath.Clean(tmpPath), filepath.Clean(path))
|
||||||
if err := os.Rename(filepath.Clean(tmpPath), filepath.Clean(path)); err != nil {
|
if err != nil {
|
||||||
|
_ = os.Remove(tmpPath)
|
||||||
|
|
||||||
return 0, fmt.Errorf("failed to rename temp file: %w", err)
|
return 0, fmt.Errorf("failed to rename temp file: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -414,10 +442,8 @@ func (s *VariantStorage) Store(key VariantKey, r io.Reader, contentType string)
|
|||||||
return 0, fmt.Errorf("failed to marshal metadata: %w", err)
|
return 0, fmt.Errorf("failed to marshal metadata: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := os.WriteFile(metaPath, metaData, StorageFilePerm); err != nil {
|
// Metadata write failure is non-fatal; content is already stored.
|
||||||
// Non-fatal, content is stored
|
_ = os.WriteFile(metaPath, metaData, StorageFilePerm)
|
||||||
_ = err
|
|
||||||
}
|
|
||||||
|
|
||||||
return size, nil
|
return size, nil
|
||||||
}
|
}
|
||||||
@@ -438,8 +464,11 @@ func (s *VariantStorage) Load(key VariantKey) (io.ReadCloser, error) {
|
|||||||
return f, nil
|
return f, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// LoadWithMeta returns a reader, size, and content type for the content at the given key.
|
// LoadWithMeta returns a reader, size, and content type for the content at
|
||||||
func (s *VariantStorage) LoadWithMeta(key VariantKey) (io.ReadCloser, int64, string, error) {
|
// the given key.
|
||||||
|
func (s *VariantStorage) LoadWithMeta(
|
||||||
|
key VariantKey,
|
||||||
|
) (io.ReadCloser, int64, string, error) {
|
||||||
path := s.keyToPath(key)
|
path := s.keyToPath(key)
|
||||||
metaPath := path + ".meta"
|
metaPath := path + ".meta"
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package imgcache
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -9,13 +10,17 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestContentStorage_StoreAndLoad(t *testing.T) {
|
func TestContentStorage_StoreAndLoad(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
content := []byte("hello world")
|
content := []byte("hello world")
|
||||||
|
|
||||||
hash, size, err := storage.Store(bytes.NewReader(content))
|
hash, size, err := storage.Store(bytes.NewReader(content))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Store() error = %v", err)
|
t.Fatalf("Store() error = %v", err)
|
||||||
@@ -31,8 +36,11 @@ func TestContentStorage_StoreAndLoad(t *testing.T) {
|
|||||||
|
|
||||||
// Verify file exists at expected path
|
// Verify file exists at expected path
|
||||||
hashStr := string(hash)
|
hashStr := string(hash)
|
||||||
|
|
||||||
expectedPath := filepath.Join(tmpDir, hashStr[0:2], hashStr[2:4], hashStr)
|
expectedPath := filepath.Join(tmpDir, hashStr[0:2], hashStr[2:4], hashStr)
|
||||||
if _, err := os.Stat(expectedPath); err != nil {
|
|
||||||
|
_, err = os.Stat(expectedPath)
|
||||||
|
if err != nil {
|
||||||
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -41,7 +49,8 @@ func TestContentStorage_StoreAndLoad(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Load() error = %v", err)
|
t.Fatalf("Load() error = %v", err)
|
||||||
}
|
}
|
||||||
defer r.Close()
|
|
||||||
|
defer func() { _ = r.Close() }()
|
||||||
|
|
||||||
loaded, err := io.ReadAll(r)
|
loaded, err := io.ReadAll(r)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -54,7 +63,10 @@ func TestContentStorage_StoreAndLoad(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestContentStorage_StoreIdempotent(t *testing.T) {
|
func TestContentStorage_StoreIdempotent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
@@ -78,26 +90,33 @@ func TestContentStorage_StoreIdempotent(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestContentStorage_LoadNotFound(t *testing.T) {
|
func TestContentStorage_LoadNotFound(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = storage.Load(ContentHash("nonexistent"))
|
_, err = storage.Load(ContentHash("nonexistent"))
|
||||||
if err != ErrNotFound {
|
if !errors.Is(err, ErrNotFound) {
|
||||||
t.Errorf("Load() error = %v, want ErrNotFound", err)
|
t.Errorf("Load() error = %v, want ErrNotFound", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestContentStorage_Delete(t *testing.T) {
|
func TestContentStorage_Delete(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
content := []byte("to be deleted")
|
content := []byte("to be deleted")
|
||||||
|
|
||||||
hash, _, err := storage.Store(bytes.NewReader(content))
|
hash, _, err := storage.Store(bytes.NewReader(content))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Store() error = %v", err)
|
t.Fatalf("Store() error = %v", err)
|
||||||
@@ -107,7 +126,8 @@ func TestContentStorage_Delete(t *testing.T) {
|
|||||||
t.Error("Exists() = false, want true")
|
t.Error("Exists() = false, want true")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := storage.Delete(hash); err != nil {
|
err = storage.Delete(hash)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("Delete() error = %v", err)
|
t.Fatalf("Delete() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,20 +137,27 @@ func TestContentStorage_Delete(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestContentStorage_DeleteNonexistent(t *testing.T) {
|
func TestContentStorage_DeleteNonexistent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Should not error
|
// Should not error
|
||||||
if err := storage.Delete(ContentHash("nonexistent")); err != nil {
|
err = storage.Delete(ContentHash("nonexistent"))
|
||||||
|
if err != nil {
|
||||||
t.Errorf("Delete() error = %v, want nil", err)
|
t.Errorf("Delete() error = %v, want nil", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestContentStorage_HashToPath(t *testing.T) {
|
func TestContentStorage_HashToPath(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewContentStorage(tmpDir)
|
storage, err := NewContentStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewContentStorage() error = %v", err)
|
t.Fatalf("NewContentStorage() error = %v", err)
|
||||||
@@ -138,50 +165,59 @@ func TestContentStorage_HashToPath(t *testing.T) {
|
|||||||
|
|
||||||
// Test by storing and verifying the resulting path structure
|
// Test by storing and verifying the resulting path structure
|
||||||
content := []byte("test content for path verification")
|
content := []byte("test content for path verification")
|
||||||
|
|
||||||
hash, _, err := storage.Store(bytes.NewReader(content))
|
hash, _, err := storage.Store(bytes.NewReader(content))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Store() error = %v", err)
|
t.Fatalf("Store() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
hashStr := string(hash)
|
hashStr := string(hash)
|
||||||
|
|
||||||
expectedPath := filepath.Join(tmpDir, hashStr[0:2], hashStr[2:4], hashStr)
|
expectedPath := filepath.Join(tmpDir, hashStr[0:2], hashStr[2:4], hashStr)
|
||||||
if _, err := os.Stat(expectedPath); err != nil {
|
|
||||||
|
_, err = os.Stat(expectedPath)
|
||||||
|
if err != nil {
|
||||||
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMetadataStorage_StoreAndLoad(t *testing.T) {
|
func TestMetadataStorage_StoreAndLoad(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewMetadataStorage(tmpDir)
|
storage, err := NewMetadataStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewMetadataStorage() error = %v", err)
|
t.Fatalf("NewMetadataStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
meta := &SourceMetadata{
|
meta := &SourceMetadata{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN,
|
||||||
Path: "/photos/cat.jpg",
|
Path: testPathCat,
|
||||||
ContentHash: "abc123",
|
ContentHash: "abc123",
|
||||||
StatusCode: 200,
|
StatusCode: 200,
|
||||||
ContentType: "image/jpeg",
|
ContentType: testContentTypeJPEG,
|
||||||
FetchedAt: 1704067200,
|
FetchedAt: 1704067200,
|
||||||
ETag: `"etag123"`,
|
ETag: `"etag123"`,
|
||||||
}
|
}
|
||||||
|
|
||||||
pathHash := HashPath("/photos/cat.jpg")
|
pathHash := HashPath(testPathCat)
|
||||||
|
|
||||||
err = storage.Store("cdn.example.com", pathHash, meta)
|
err = storage.Store(testHostCDN, pathHash, meta)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Store() error = %v", err)
|
t.Fatalf("Store() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify file exists at expected path
|
// Verify file exists at expected path
|
||||||
expectedPath := filepath.Join(tmpDir, "cdn.example.com", string(pathHash)+".json")
|
expectedPath := filepath.Join(tmpDir, testHostCDN, string(pathHash)+".json")
|
||||||
if _, err := os.Stat(expectedPath); err != nil {
|
|
||||||
|
_, err = os.Stat(expectedPath)
|
||||||
|
if err != nil {
|
||||||
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
t.Errorf("File not at expected path %s: %v", expectedPath, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load and verify
|
// Load and verify
|
||||||
loaded, err := storage.Load("cdn.example.com", pathHash)
|
loaded, err := storage.Load(testHostCDN, pathHash)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Load() error = %v", err)
|
t.Fatalf("Load() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -208,55 +244,64 @@ func TestMetadataStorage_StoreAndLoad(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestMetadataStorage_LoadNotFound(t *testing.T) {
|
func TestMetadataStorage_LoadNotFound(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewMetadataStorage(tmpDir)
|
storage, err := NewMetadataStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewMetadataStorage() error = %v", err)
|
t.Fatalf("NewMetadataStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = storage.Load("example.com", PathHash("nonexistent"))
|
_, err = storage.Load(testHostExample, PathHash("nonexistent"))
|
||||||
if err != ErrNotFound {
|
if !errors.Is(err, ErrNotFound) {
|
||||||
t.Errorf("Load() error = %v, want ErrNotFound", err)
|
t.Errorf("Load() error = %v, want ErrNotFound", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMetadataStorage_Delete(t *testing.T) {
|
func TestMetadataStorage_Delete(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tmpDir := t.TempDir()
|
tmpDir := t.TempDir()
|
||||||
|
|
||||||
storage, err := NewMetadataStorage(tmpDir)
|
storage, err := NewMetadataStorage(tmpDir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewMetadataStorage() error = %v", err)
|
t.Fatalf("NewMetadataStorage() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
meta := &SourceMetadata{
|
meta := &SourceMetadata{
|
||||||
Host: "example.com",
|
Host: testHostExample,
|
||||||
Path: "/test.jpg",
|
Path: "/test.jpg",
|
||||||
StatusCode: 200,
|
StatusCode: 200,
|
||||||
}
|
}
|
||||||
|
|
||||||
pathHash := HashPath("/test.jpg")
|
pathHash := HashPath("/test.jpg")
|
||||||
|
|
||||||
err = storage.Store("example.com", pathHash, meta)
|
err = storage.Store(testHostExample, pathHash, meta)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Store() error = %v", err)
|
t.Fatalf("Store() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !storage.Exists("example.com", pathHash) {
|
if !storage.Exists(testHostExample, pathHash) {
|
||||||
t.Error("Exists() = false, want true")
|
t.Error("Exists() = false, want true")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := storage.Delete("example.com", pathHash); err != nil {
|
err = storage.Delete(testHostExample, pathHash)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("Delete() error = %v", err)
|
t.Fatalf("Delete() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if storage.Exists("example.com", pathHash) {
|
if storage.Exists(testHostExample, pathHash) {
|
||||||
t.Error("Exists() = true after delete, want false")
|
t.Error("Exists() = true after delete, want false")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHashPath(t *testing.T) {
|
func TestHashPath(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Same input should produce same hash
|
// Same input should produce same hash
|
||||||
hash1 := HashPath("/photos/cat.jpg")
|
hash1 := HashPath(testPathCat)
|
||||||
hash2 := HashPath("/photos/cat.jpg")
|
hash2 := HashPath(testPathCat)
|
||||||
|
|
||||||
if hash1 != hash2 {
|
if hash1 != hash2 {
|
||||||
t.Errorf("HashPath() not deterministic: %s vs %s", hash1, hash2)
|
t.Errorf("HashPath() not deterministic: %s vs %s", hash1, hash2)
|
||||||
@@ -276,9 +321,11 @@ func TestHashPath(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCacheKey(t *testing.T) {
|
func TestCacheKey(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req1 := &ImageRequest{
|
req1 := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -287,8 +334,8 @@ func TestCacheKey(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
req2 := &ImageRequest{
|
req2 := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -311,8 +358,8 @@ func TestCacheKey(t *testing.T) {
|
|||||||
|
|
||||||
// Different size should produce different key
|
// Different size should produce different key
|
||||||
req3 := &ImageRequest{
|
req3 := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Size: Size{Width: 400, Height: 300}, // Different size
|
Size: Size{Width: 400, Height: 300}, // Different size
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -327,8 +374,8 @@ func TestCacheKey(t *testing.T) {
|
|||||||
|
|
||||||
// Different format should produce different key
|
// Different format should produce different key
|
||||||
req4 := &ImageRequest{
|
req4 := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatPNG, // Different format
|
Format: FormatPNG, // Different format
|
||||||
@@ -343,8 +390,8 @@ func TestCacheKey(t *testing.T) {
|
|||||||
|
|
||||||
// Different quality should produce different key
|
// Different quality should produce different key
|
||||||
req5 := &ImageRequest{
|
req5 := &ImageRequest{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHostCDN,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPathCat,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -18,6 +18,14 @@ import (
|
|||||||
"sneak.berlin/go/pixa/internal/httpfetcher"
|
"sneak.berlin/go/pixa/internal/httpfetcher"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Shared test data literals, extracted as constants for goconst.
|
||||||
|
const (
|
||||||
|
testHostCDN = "cdn.example.com"
|
||||||
|
testHostExample = "example.com"
|
||||||
|
testPathCat = "/photos/cat.jpg"
|
||||||
|
testContentTypeJPEG = "image/jpeg"
|
||||||
|
)
|
||||||
|
|
||||||
// TestFixtures contains paths to test files in the mock filesystem.
|
// TestFixtures contains paths to test files in the mock filesystem.
|
||||||
type TestFixtures struct {
|
type TestFixtures struct {
|
||||||
// Valid image files
|
// Valid image files
|
||||||
@@ -89,14 +97,16 @@ func generateTestJPEG(t *testing.T, width, height int, c color.Color) []byte {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||||
for y := 0; y < height; y++ {
|
for y := range height {
|
||||||
for x := 0; x < width; x++ {
|
for x := range width {
|
||||||
img.Set(x, y, c)
|
img.Set(x, y, c)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 85}); err != nil {
|
|
||||||
|
err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 85})
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test JPEG: %v", err)
|
t.Fatalf("failed to encode test JPEG: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -108,14 +118,16 @@ func generateTestPNG(t *testing.T, width, height int, c color.Color) []byte {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
img := image.NewRGBA(image.Rect(0, 0, width, height))
|
||||||
for y := 0; y < height; y++ {
|
for y := range height {
|
||||||
for x := 0; x < width; x++ {
|
for x := range width {
|
||||||
img.Set(x, y, c)
|
img.Set(x, y, c)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := png.Encode(&buf, img); err != nil {
|
|
||||||
|
err := png.Encode(&buf, img)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test PNG: %v", err)
|
t.Fatalf("failed to encode test PNG: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -126,15 +138,20 @@ func generateTestPNG(t *testing.T, width, height int, c color.Color) []byte {
|
|||||||
func generateTestGIF(t *testing.T, width, height int, c color.Color) []byte {
|
func generateTestGIF(t *testing.T, width, height int, c color.Color) []byte {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
img := image.NewPaletted(image.Rect(0, 0, width, height), []color.Color{c, color.White})
|
img := image.NewPaletted(
|
||||||
for y := 0; y < height; y++ {
|
image.Rect(0, 0, width, height),
|
||||||
for x := 0; x < width; x++ {
|
[]color.Color{c, color.White},
|
||||||
|
)
|
||||||
|
for y := range height {
|
||||||
|
for x := range width {
|
||||||
img.SetColorIndex(x, y, 0)
|
img.SetColorIndex(x, y, 0)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
if err := gif.Encode(&buf, img, nil); err != nil {
|
|
||||||
|
err := gif.Encode(&buf, img, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to encode test GIF: %v", err)
|
t.Fatalf("failed to encode test GIF: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -142,7 +159,9 @@ func generateTestGIF(t *testing.T, width, height int, c color.Color) []byte {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// SetupTestService creates a Service with mock fetcher for testing.
|
// SetupTestService creates a Service with mock fetcher for testing.
|
||||||
func SetupTestService(t *testing.T, opts ...TestServiceOption) (*Service, *TestFixtures) {
|
func SetupTestService(
|
||||||
|
t *testing.T, opts ...TestServiceOption,
|
||||||
|
) (*Service, *TestFixtures) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
mockFS, fixtures := NewTestFS(t)
|
mockFS, fixtures := NewTestFS(t)
|
||||||
@@ -195,7 +214,8 @@ func setupServiceTestDB(t *testing.T) *sql.DB {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Use the real production schema via migrations
|
// Use the real production schema via migrations
|
||||||
if err := database.ApplyMigrations(context.Background(), db, nil); err != nil {
|
err = database.ApplyMigrations(context.Background(), db, nil)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("failed to apply migrations: %v", err)
|
t.Fatalf("failed to apply migrations: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -40,7 +40,8 @@ type ParsedURL struct {
|
|||||||
Format ImageFormat
|
Format ImageFormat
|
||||||
}
|
}
|
||||||
|
|
||||||
// ParseImagePath parses the path captured by chi's wildcard: <host>/<path>/<size>.<format>
|
// ParseImagePath parses the path captured by chi's wildcard:
|
||||||
|
// <host>/<path>/<size>.<format>
|
||||||
// This is the primary entry point when using chi routing.
|
// This is the primary entry point when using chi routing.
|
||||||
// Examples:
|
// Examples:
|
||||||
// - cdn.example.com/photos/cat.jpg/800x600.webp
|
// - cdn.example.com/photos/cat.jpg/800x600.webp
|
||||||
@@ -76,7 +77,8 @@ func ParseImageURL(urlPath string) (*ParsedURL, error) {
|
|||||||
// parseImageComponents parses <host>/<path>/<size>.<format> structure.
|
// parseImageComponents parses <host>/<path>/<size>.<format> structure.
|
||||||
func parseImageComponents(remainder string) (*ParsedURL, error) {
|
func parseImageComponents(remainder string) (*ParsedURL, error) {
|
||||||
// Check for path traversal before any other processing
|
// Check for path traversal before any other processing
|
||||||
if err := checkPathTraversal(remainder); err != nil {
|
err := checkPathTraversal(remainder)
|
||||||
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,6 +104,7 @@ func parseImageComponents(remainder string) (*ParsedURL, error) {
|
|||||||
// Split host from path
|
// Split host from path
|
||||||
// The first segment is the host, everything after is the path
|
// The first segment is the host, everything after is the path
|
||||||
firstSlash := strings.Index(hostAndPath, "/")
|
firstSlash := strings.Index(hostAndPath, "/")
|
||||||
|
|
||||||
var host, path, query string
|
var host, path, query string
|
||||||
|
|
||||||
if firstSlash == -1 {
|
if firstSlash == -1 {
|
||||||
@@ -181,8 +184,7 @@ func checkPathTraversal(path string) error {
|
|||||||
|
|
||||||
// Also check for ".." as a path segment in the original path
|
// Also check for ".." as a path segment in the original path
|
||||||
// This catches cases where the path hasn't been normalized
|
// This catches cases where the path hasn't been normalized
|
||||||
segments := strings.Split(path, "/")
|
for seg := range strings.SplitSeq(path, "/") {
|
||||||
for _, seg := range segments {
|
|
||||||
// URL decode the segment
|
// URL decode the segment
|
||||||
decodedSeg, _ := url.PathUnescape(seg)
|
decodedSeg, _ := url.PathUnescape(seg)
|
||||||
decodedSeg = strings.ReplaceAll(decodedSeg, "\\", "/")
|
decodedSeg = strings.ReplaceAll(decodedSeg, "\\", "/")
|
||||||
@@ -202,8 +204,10 @@ func parseSizeFormat(s string) (Size, ImageFormat, error) {
|
|||||||
return Size{}, "", ErrInvalidSize
|
return Size{}, "", ErrInvalidSize
|
||||||
}
|
}
|
||||||
|
|
||||||
var size Size
|
var (
|
||||||
var formatStr string
|
size Size
|
||||||
|
formatStr string
|
||||||
|
)
|
||||||
|
|
||||||
if matches[4] == "orig" {
|
if matches[4] == "orig" {
|
||||||
// "orig.format" pattern
|
// "orig.format" pattern
|
||||||
|
|||||||
@@ -1,93 +1,124 @@
|
|||||||
package imgcache
|
package imgcache
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// assertParsedURL compares all fields of a parsed URL against the
|
||||||
|
// expected value.
|
||||||
|
func assertParsedURL(t *testing.T, got, want *ParsedURL) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
if got.Host != want.Host {
|
||||||
|
t.Errorf("Host = %q, want %q", got.Host, want.Host)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.Path != want.Path {
|
||||||
|
t.Errorf("Path = %q, want %q", got.Path, want.Path)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.Query != want.Query {
|
||||||
|
t.Errorf("Query = %q, want %q", got.Query, want.Query)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.Size != want.Size {
|
||||||
|
t.Errorf("Size = %v, want %v", got.Size, want.Size)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got.Format != want.Format {
|
||||||
|
t.Errorf("Format = %q, want %q", got.Format, want.Format)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestParseImageURL(t *testing.T) {
|
func TestParseImageURL(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
input string
|
input string
|
||||||
want *ParsedURL
|
want *ParsedURL
|
||||||
wantErr error
|
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "basic path with size",
|
name: "basic path with size",
|
||||||
input: "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp",
|
input: "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN, Path: testPathCat,
|
||||||
Path: "/photos/cat.jpg",
|
Size: Size{Width: 800, Height: 600}, Format: FormatWebP,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 800, Height: 600},
|
|
||||||
Format: FormatWebP,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "original size with 0x0",
|
name: "original size with 0x0",
|
||||||
input: "/v1/image/cdn.example.com/photos/cat.jpg/0x0.jpeg",
|
input: "/v1/image/cdn.example.com/photos/cat.jpg/0x0.jpeg",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN, Path: testPathCat,
|
||||||
Path: "/photos/cat.jpg",
|
Size: Size{Width: 0, Height: 0}, Format: FormatJPEG,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 0, Height: 0},
|
|
||||||
Format: FormatJPEG,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "original size with orig keyword",
|
name: "original size with orig keyword",
|
||||||
input: "/v1/image/cdn.example.com/photos/cat.jpg/orig.png",
|
input: "/v1/image/cdn.example.com/photos/cat.jpg/orig.png",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN, Path: testPathCat,
|
||||||
Path: "/photos/cat.jpg",
|
Size: Size{Width: 0, Height: 0}, Format: FormatPNG,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 0, Height: 0},
|
|
||||||
Format: FormatPNG,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "path with query string",
|
name: "path with query string",
|
||||||
input: "/v1/image/cdn.example.com/photos/cat.jpg?arg1=val1&arg2=val2/800x600.webp",
|
input: "/v1/image/cdn.example.com/photos/cat.jpg?arg1=val1&arg2=val2/800x600.webp",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN, Path: testPathCat, Query: "arg1=val1&arg2=val2",
|
||||||
Path: "/photos/cat.jpg",
|
Size: Size{Width: 800, Height: 600}, Format: FormatWebP,
|
||||||
Query: "arg1=val1&arg2=val2",
|
|
||||||
Size: Size{Width: 800, Height: 600},
|
|
||||||
Format: FormatWebP,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "deep nested path",
|
name: "deep nested path",
|
||||||
input: "/v1/image/cdn.example.com/a/b/c/d/image.jpg/1920x1080.avif",
|
input: "/v1/image/cdn.example.com/a/b/c/d/image.jpg/1920x1080.avif",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN, Path: "/a/b/c/d/image.jpg",
|
||||||
Path: "/a/b/c/d/image.jpg",
|
Size: Size{Width: 1920, Height: 1080}, Format: FormatAVIF,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 1920, Height: 1080},
|
|
||||||
Format: FormatAVIF,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "jpg alias for jpeg",
|
name: "jpg alias for jpeg",
|
||||||
input: "/v1/image/example.com/img.png/100x100.jpg",
|
input: "/v1/image/example.com/img.png/100x100.jpg",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "example.com",
|
Host: testHostExample, Path: "/img.png",
|
||||||
Path: "/img.png",
|
Size: Size{Width: 100, Height: 100}, Format: FormatJPEG,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 100, Height: 100},
|
|
||||||
Format: FormatJPEG,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "gif format",
|
name: "gif format",
|
||||||
input: "/v1/image/example.com/animated.gif/200x200.gif",
|
input: "/v1/image/example.com/animated.gif/200x200.gif",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "example.com",
|
Host: testHostExample, Path: "/animated.gif",
|
||||||
Path: "/animated.gif",
|
Size: Size{Width: 200, Height: 200}, Format: FormatGIF,
|
||||||
Query: "",
|
|
||||||
Size: Size{Width: 200, Height: 200},
|
|
||||||
Format: FormatGIF,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
got, err := ParseImageURL(tt.input)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("ParseImageURL() unexpected error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assertParsedURL(t, got, tt.want)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseImageURL_Errors(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantErr error
|
||||||
|
}{
|
||||||
{
|
{
|
||||||
name: "missing prefix",
|
name: "missing prefix",
|
||||||
input: "/image/cdn.example.com/photo.jpg/800x600.webp",
|
input: "/image/cdn.example.com/photo.jpg/800x600.webp",
|
||||||
@@ -122,47 +153,23 @@ func TestParseImageURL(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
got, err := ParseImageURL(tt.input)
|
t.Parallel()
|
||||||
|
|
||||||
if tt.wantErr != nil {
|
_, err := ParseImageURL(tt.input)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Errorf("ParseImageURL() error = nil, wantErr %v", tt.wantErr)
|
t.Fatalf("ParseImageURL() error = nil, wantErr %v", tt.wantErr)
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if !errorIs(err, tt.wantErr) {
|
|
||||||
t.Errorf("ParseImageURL() error = %v, wantErr %v", err, tt.wantErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if !errorIs(err, tt.wantErr) {
|
||||||
t.Errorf("ParseImageURL() unexpected error = %v", err)
|
t.Errorf("ParseImageURL() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if got.Host != tt.want.Host {
|
|
||||||
t.Errorf("Host = %q, want %q", got.Host, tt.want.Host)
|
|
||||||
}
|
|
||||||
if got.Path != tt.want.Path {
|
|
||||||
t.Errorf("Path = %q, want %q", got.Path, tt.want.Path)
|
|
||||||
}
|
|
||||||
if got.Query != tt.want.Query {
|
|
||||||
t.Errorf("Query = %q, want %q", got.Query, tt.want.Query)
|
|
||||||
}
|
|
||||||
if got.Size != tt.want.Size {
|
|
||||||
t.Errorf("Size = %v, want %v", got.Size, tt.want.Size)
|
|
||||||
}
|
|
||||||
if got.Format != tt.want.Format {
|
|
||||||
t.Errorf("Format = %q, want %q", got.Format, tt.want.Format)
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParseImagePath(t *testing.T) {
|
func TestParseImagePath(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// ParseImagePath is for chi wildcard capture (no /v1/image/ prefix)
|
// ParseImagePath is for chi wildcard capture (no /v1/image/ prefix)
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -174,8 +181,8 @@ func TestParseImagePath(t *testing.T) {
|
|||||||
name: "chi wildcard capture",
|
name: "chi wildcard capture",
|
||||||
input: "cdn.example.com/photos/cat.jpg/800x600.webp",
|
input: "cdn.example.com/photos/cat.jpg/800x600.webp",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN,
|
||||||
Path: "/photos/cat.jpg",
|
Path: testPathCat,
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
},
|
},
|
||||||
@@ -184,8 +191,8 @@ func TestParseImagePath(t *testing.T) {
|
|||||||
name: "with leading slash from chi",
|
name: "with leading slash from chi",
|
||||||
input: "/cdn.example.com/photos/cat.jpg/800x600.webp",
|
input: "/cdn.example.com/photos/cat.jpg/800x600.webp",
|
||||||
want: &ParsedURL{
|
want: &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN,
|
||||||
Path: "/photos/cat.jpg",
|
Path: testPathCat,
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
},
|
},
|
||||||
@@ -194,35 +201,30 @@ func TestParseImagePath(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got, err := ParseImagePath(tt.input)
|
got, err := ParseImagePath(tt.input)
|
||||||
if (err != nil) != tt.wantErr {
|
if (err != nil) != tt.wantErr {
|
||||||
t.Errorf("ParseImagePath() error = %v, wantErr %v", err, tt.wantErr)
|
t.Errorf("ParseImagePath() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if got.Host != tt.want.Host {
|
|
||||||
t.Errorf("Host = %q, want %q", got.Host, tt.want.Host)
|
assertParsedURL(t, got, tt.want)
|
||||||
}
|
|
||||||
if got.Path != tt.want.Path {
|
|
||||||
t.Errorf("Path = %q, want %q", got.Path, tt.want.Path)
|
|
||||||
}
|
|
||||||
if got.Size != tt.want.Size {
|
|
||||||
t.Errorf("Size = %v, want %v", got.Size, tt.want.Size)
|
|
||||||
}
|
|
||||||
if got.Format != tt.want.Format {
|
|
||||||
t.Errorf("Format = %q, want %q", got.Format, tt.want.Format)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParsedURL_ToImageRequest(t *testing.T) {
|
func TestParsedURL_ToImageRequest(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
parsed := &ParsedURL{
|
parsed := &ParsedURL{
|
||||||
Host: "cdn.example.com",
|
Host: testHostCDN,
|
||||||
Path: "/photos/cat.jpg",
|
Path: testPathCat,
|
||||||
Query: "version=2",
|
Query: "version=2",
|
||||||
Size: Size{Width: 800, Height: 600},
|
Size: Size{Width: 800, Height: 600},
|
||||||
Format: FormatWebP,
|
Format: FormatWebP,
|
||||||
@@ -233,21 +235,27 @@ func TestParsedURL_ToImageRequest(t *testing.T) {
|
|||||||
if req.SourceHost != parsed.Host {
|
if req.SourceHost != parsed.Host {
|
||||||
t.Errorf("SourceHost = %q, want %q", req.SourceHost, parsed.Host)
|
t.Errorf("SourceHost = %q, want %q", req.SourceHost, parsed.Host)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.SourcePath != parsed.Path {
|
if req.SourcePath != parsed.Path {
|
||||||
t.Errorf("SourcePath = %q, want %q", req.SourcePath, parsed.Path)
|
t.Errorf("SourcePath = %q, want %q", req.SourcePath, parsed.Path)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.SourceQuery != parsed.Query {
|
if req.SourceQuery != parsed.Query {
|
||||||
t.Errorf("SourceQuery = %q, want %q", req.SourceQuery, parsed.Query)
|
t.Errorf("SourceQuery = %q, want %q", req.SourceQuery, parsed.Query)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Size != parsed.Size {
|
if req.Size != parsed.Size {
|
||||||
t.Errorf("Size = %v, want %v", req.Size, parsed.Size)
|
t.Errorf("Size = %v, want %v", req.Size, parsed.Size)
|
||||||
}
|
}
|
||||||
|
|
||||||
if req.Format != parsed.Format {
|
if req.Format != parsed.Format {
|
||||||
t.Errorf("Format = %q, want %q", req.Format, parsed.Format)
|
t.Errorf("Format = %q, want %q", req.Format, parsed.Format)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParseImageURL_PathTraversal(t *testing.T) {
|
func TestParseImageURL_PathTraversal(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// All path traversal attempts should be rejected
|
// All path traversal attempts should be rejected
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -293,12 +301,14 @@ func TestParseImageURL_PathTraversal(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
_, err := ParseImageURL(tt.input)
|
_, err := ParseImageURL(tt.input)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("ParseImageURL() should reject path traversal attempts")
|
t.Error("ParseImageURL() should reject path traversal attempts")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrPathTraversal {
|
if !errors.Is(err, ErrPathTraversal) {
|
||||||
t.Errorf("ParseImageURL() error = %v, want ErrPathTraversal", err)
|
t.Errorf("ParseImageURL() error = %v, want ErrPathTraversal", err)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -306,6 +316,8 @@ func TestParseImageURL_PathTraversal(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestParseImagePath_PathTraversal(t *testing.T) {
|
func TestParseImagePath_PathTraversal(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Test path traversal via ParseImagePath (chi wildcard)
|
// Test path traversal via ParseImagePath (chi wildcard)
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -323,12 +335,14 @@ func TestParseImagePath_PathTraversal(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
_, err := ParseImagePath(tt.input)
|
_, err := ParseImagePath(tt.input)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("ParseImagePath() should reject path traversal attempts")
|
t.Error("ParseImagePath() should reject path traversal attempts")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrPathTraversal {
|
if !errors.Is(err, ErrPathTraversal) {
|
||||||
t.Errorf("ParseImagePath() error = %v, want ErrPathTraversal", err)
|
t.Errorf("ParseImagePath() error = %v, want ErrPathTraversal", err)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -337,7 +351,7 @@ func TestParseImagePath_PathTraversal(t *testing.T) {
|
|||||||
|
|
||||||
// errorIs checks if err matches target (handles wrapped errors).
|
// errorIs checks if err matches target (handles wrapped errors).
|
||||||
func errorIs(err, target error) bool {
|
func errorIs(err, target error) bool {
|
||||||
if err == target {
|
if errors.Is(err, target) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
// Check if error message contains target message for wrapped errors
|
// Check if error message contains target message for wrapped errors
|
||||||
@@ -15,6 +15,7 @@ import (
|
|||||||
// Params defines dependencies for Logger.
|
// Params defines dependencies for Logger.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Globals *globals.Globals
|
Globals *globals.Globals
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -46,6 +46,10 @@ const (
|
|||||||
// MinMagicBytes is the minimum number of bytes needed to detect format.
|
// MinMagicBytes is the minimum number of bytes needed to detect format.
|
||||||
const MinMagicBytes = 12
|
const MinMagicBytes = 12
|
||||||
|
|
||||||
|
// mimeOctetStream is the fallback MIME type for formats without a
|
||||||
|
// specific MIME type.
|
||||||
|
const mimeOctetStream = "application/octet-stream"
|
||||||
|
|
||||||
// Magic byte signatures for supported formats.
|
// Magic byte signatures for supported formats.
|
||||||
// These are effectively constants but Go doesn't support const slices.
|
// These are effectively constants but Go doesn't support const slices.
|
||||||
//
|
//
|
||||||
@@ -190,14 +194,17 @@ func IsSupportedMIMEType(mimeType string) bool {
|
|||||||
func PeekAndValidate(r io.Reader, declaredType string) (io.Reader, error) {
|
func PeekAndValidate(r io.Reader, declaredType string) (io.Reader, error) {
|
||||||
// Read minimum bytes for detection
|
// Read minimum bytes for detection
|
||||||
buf := make([]byte, MinMagicBytes)
|
buf := make([]byte, MinMagicBytes)
|
||||||
|
|
||||||
n, err := io.ReadFull(r, buf)
|
n, err := io.ReadFull(r, buf)
|
||||||
if err != nil && err != io.ErrUnexpectedEOF {
|
if err != nil && !errors.Is(err, io.ErrUnexpectedEOF) {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
buf = buf[:n]
|
buf = buf[:n]
|
||||||
|
|
||||||
// Validate magic bytes
|
// Validate magic bytes
|
||||||
if err := ValidateMagicBytes(buf, declaredType); err != nil {
|
err = ValidateMagicBytes(buf, declaredType)
|
||||||
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -219,6 +226,9 @@ func MIMEToImageFormat(mimeType string) (ImageFormat, bool) {
|
|||||||
return FormatGIF, true
|
return FormatGIF, true
|
||||||
case MIMETypeAVIF:
|
case MIMETypeAVIF:
|
||||||
return FormatAVIF, true
|
return FormatAVIF, true
|
||||||
|
case MIMETypeSVG:
|
||||||
|
// SVG has no corresponding output format.
|
||||||
|
return "", false
|
||||||
default:
|
default:
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
@@ -237,7 +247,10 @@ func ImageFormatToMIME(format ImageFormat) string {
|
|||||||
return string(MIMETypeGIF)
|
return string(MIMETypeGIF)
|
||||||
case FormatAVIF:
|
case FormatAVIF:
|
||||||
return string(MIMETypeAVIF)
|
return string(MIMETypeAVIF)
|
||||||
|
case FormatOriginal:
|
||||||
|
// Original format passes content through unchanged.
|
||||||
|
return mimeOctetStream
|
||||||
default:
|
default:
|
||||||
return "application/octet-stream"
|
return mimeOctetStream
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,121 +2,90 @@ package magic
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Shared test fixture strings.
|
||||||
|
const (
|
||||||
|
testNameEmpty = "empty"
|
||||||
|
testMIMEJPEG = "image/jpeg"
|
||||||
|
testMIMEJPEGParams = "image/jpeg; charset=utf-8"
|
||||||
|
testMIMEPNG = "image/png"
|
||||||
|
testMIMEWebP = "image/webp"
|
||||||
|
testMIMEGIF = "image/gif"
|
||||||
|
testMIMEAVIF = "image/avif"
|
||||||
|
)
|
||||||
|
|
||||||
|
// pad appends zero bytes so data is comfortably above MinMagicBytes.
|
||||||
|
func pad(b ...byte) []byte {
|
||||||
|
return append(b, make([]byte, 100)...)
|
||||||
|
}
|
||||||
|
|
||||||
func TestDetectFormat(t *testing.T) {
|
func TestDetectFormat(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
jpeg := pad(0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01)
|
||||||
|
png := pad(0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D)
|
||||||
|
gif87a := pad(0x47, 0x49, 0x46, 0x38, 0x37, 0x61, 0, 0, 0, 0, 0, 0)
|
||||||
|
gif89a := pad(0x47, 0x49, 0x46, 0x38, 0x39, 0x61, 0, 0, 0, 0, 0, 0)
|
||||||
|
// RIFF + size placeholder + WEBP
|
||||||
|
webp := pad(0x52, 0x49, 0x46, 0x46, 0, 0, 0, 0, 0x57, 0x45, 0x42, 0x50)
|
||||||
|
// box size + ftyp + brand
|
||||||
|
avif := pad(0x00, 0x00, 0x00, 0x1C, 0x66, 0x74, 0x79, 0x70, 0x61, 0x76, 0x69, 0x66)
|
||||||
|
avis := pad(0x00, 0x00, 0x00, 0x1C, 0x66, 0x74, 0x79, 0x70, 0x61, 0x76, 0x69, 0x73)
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
data []byte
|
data []byte
|
||||||
wantMIME MIMEType
|
wantMIME MIMEType
|
||||||
wantErr error
|
wantErr error
|
||||||
}{
|
}{
|
||||||
{
|
{name: "JPEG", data: jpeg, wantMIME: MIMETypeJPEG},
|
||||||
name: "JPEG",
|
{name: "PNG", data: png, wantMIME: MIMETypePNG},
|
||||||
data: append([]byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01}, make([]byte, 100)...),
|
{name: "GIF87a", data: gif87a, wantMIME: MIMETypeGIF},
|
||||||
wantMIME: MIMETypeJPEG,
|
{name: "GIF89a", data: gif89a, wantMIME: MIMETypeGIF},
|
||||||
wantErr: nil,
|
{name: "WebP", data: webp, wantMIME: MIMETypeWebP},
|
||||||
},
|
{name: "AVIF", data: avif, wantMIME: MIMETypeAVIF},
|
||||||
{
|
{name: "AVIF sequence", data: avis, wantMIME: MIMETypeAVIF},
|
||||||
name: "PNG",
|
|
||||||
data: append([]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypePNG,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "GIF87a",
|
|
||||||
data: append([]byte{0x47, 0x49, 0x46, 0x38, 0x37, 0x61, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypeGIF,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "GIF89a",
|
|
||||||
data: append([]byte{0x47, 0x49, 0x46, 0x38, 0x39, 0x61, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypeGIF,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "WebP",
|
|
||||||
data: append([]byte{
|
|
||||||
0x52, 0x49, 0x46, 0x46, // RIFF
|
|
||||||
0x00, 0x00, 0x00, 0x00, // file size (placeholder)
|
|
||||||
0x57, 0x45, 0x42, 0x50, // WEBP
|
|
||||||
}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypeWebP,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "AVIF",
|
|
||||||
data: append([]byte{
|
|
||||||
0x00, 0x00, 0x00, 0x1C, // box size
|
|
||||||
0x66, 0x74, 0x79, 0x70, // ftyp
|
|
||||||
0x61, 0x76, 0x69, 0x66, // avif brand
|
|
||||||
}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypeAVIF,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "AVIF sequence",
|
|
||||||
data: append([]byte{
|
|
||||||
0x00, 0x00, 0x00, 0x1C, // box size
|
|
||||||
0x66, 0x74, 0x79, 0x70, // ftyp
|
|
||||||
0x61, 0x76, 0x69, 0x73, // avis brand
|
|
||||||
}, make([]byte, 100)...),
|
|
||||||
wantMIME: MIMETypeAVIF,
|
|
||||||
wantErr: nil,
|
|
||||||
},
|
|
||||||
{
|
{
|
||||||
name: "SVG with XML declaration",
|
name: "SVG with XML declaration",
|
||||||
data: []byte(`<?xml version="1.0"?><svg></svg>`),
|
data: []byte(`<?xml version="1.0"?><svg></svg>`),
|
||||||
wantMIME: MIMETypeSVG,
|
wantMIME: MIMETypeSVG,
|
||||||
wantErr: nil,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "SVG without declaration",
|
name: "SVG without declaration",
|
||||||
data: []byte(`<svg xmlns="http://www.w3.org/2000/svg"></svg>`),
|
data: []byte(`<svg xmlns="http://www.w3.org/2000/svg"></svg>`),
|
||||||
wantMIME: MIMETypeSVG,
|
wantMIME: MIMETypeSVG,
|
||||||
wantErr: nil,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "SVG with whitespace",
|
name: "SVG with whitespace",
|
||||||
data: []byte(` <?xml version="1.0"?><svg></svg>`),
|
data: []byte(` <?xml version="1.0"?><svg></svg>`),
|
||||||
wantMIME: MIMETypeSVG,
|
wantMIME: MIMETypeSVG,
|
||||||
wantErr: nil,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "SVG with BOM",
|
name: "SVG with BOM",
|
||||||
data: append([]byte{0xEF, 0xBB, 0xBF}, []byte(`<svg></svg>`)...),
|
data: append([]byte{0xEF, 0xBB, 0xBF}, []byte(`<svg></svg>`)...),
|
||||||
wantMIME: MIMETypeSVG,
|
wantMIME: MIMETypeSVG,
|
||||||
wantErr: nil,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "unknown format",
|
name: "unknown format",
|
||||||
data: []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00},
|
data: make([]byte, MinMagicBytes),
|
||||||
wantMIME: "",
|
wantErr: ErrUnknownFormat,
|
||||||
wantErr: ErrUnknownFormat,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "too short",
|
|
||||||
data: []byte{0xFF, 0xD8},
|
|
||||||
wantMIME: "",
|
|
||||||
wantErr: ErrNotEnoughData,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "empty",
|
|
||||||
data: []byte{},
|
|
||||||
wantMIME: "",
|
|
||||||
wantErr: ErrNotEnoughData,
|
|
||||||
},
|
},
|
||||||
|
{name: "too short", data: []byte{0xFF, 0xD8}, wantErr: ErrNotEnoughData},
|
||||||
|
{name: testNameEmpty, data: []byte{}, wantErr: ErrNotEnoughData},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
got, err := DetectFormat(tt.data)
|
t.Parallel()
|
||||||
|
|
||||||
if err != tt.wantErr {
|
got, err := DetectFormat(tt.data)
|
||||||
|
if !errors.Is(err, tt.wantErr) {
|
||||||
t.Errorf("DetectFormat() error = %v, wantErr %v", err, tt.wantErr)
|
t.Errorf("DetectFormat() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
|
||||||
return
|
return
|
||||||
@@ -130,8 +99,10 @@ func TestDetectFormat(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateMagicBytes(t *testing.T) {
|
func TestValidateMagicBytes(t *testing.T) {
|
||||||
jpegData := append([]byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01}, make([]byte, 100)...)
|
t.Parallel()
|
||||||
pngData := append([]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D}, make([]byte, 100)...)
|
|
||||||
|
jpegData := pad(0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01)
|
||||||
|
pngData := pad(0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D)
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -142,40 +113,42 @@ func TestValidateMagicBytes(t *testing.T) {
|
|||||||
{
|
{
|
||||||
name: "matching JPEG",
|
name: "matching JPEG",
|
||||||
data: jpegData,
|
data: jpegData,
|
||||||
declaredType: "image/jpeg",
|
declaredType: testMIMEJPEG,
|
||||||
wantErr: nil,
|
wantErr: nil,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "matching JPEG with params",
|
name: "matching JPEG with params",
|
||||||
data: jpegData,
|
data: jpegData,
|
||||||
declaredType: "image/jpeg; charset=utf-8",
|
declaredType: testMIMEJPEGParams,
|
||||||
wantErr: nil,
|
wantErr: nil,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "matching PNG",
|
name: "matching PNG",
|
||||||
data: pngData,
|
data: pngData,
|
||||||
declaredType: "image/png",
|
declaredType: testMIMEPNG,
|
||||||
wantErr: nil,
|
wantErr: nil,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "mismatched type",
|
name: "mismatched type",
|
||||||
data: jpegData,
|
data: jpegData,
|
||||||
declaredType: "image/png",
|
declaredType: testMIMEPNG,
|
||||||
wantErr: ErrMagicByteMismatch,
|
wantErr: ErrMagicByteMismatch,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "unknown data",
|
name: "unknown data",
|
||||||
data: []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00},
|
data: make([]byte, MinMagicBytes),
|
||||||
declaredType: "image/jpeg",
|
declaredType: testMIMEJPEG,
|
||||||
wantErr: ErrUnknownFormat,
|
wantErr: ErrUnknownFormat,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
err := ValidateMagicBytes(tt.data, tt.declaredType)
|
err := ValidateMagicBytes(tt.data, tt.declaredType)
|
||||||
|
|
||||||
if err != tt.wantErr {
|
if !errors.Is(err, tt.wantErr) {
|
||||||
t.Errorf("ValidateMagicBytes() error = %v, wantErr %v", err, tt.wantErr)
|
t.Errorf("ValidateMagicBytes() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -183,27 +156,31 @@ func TestValidateMagicBytes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIsSupportedMIMEType(t *testing.T) {
|
func TestIsSupportedMIMEType(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
mimeType string
|
mimeType string
|
||||||
want bool
|
want bool
|
||||||
}{
|
}{
|
||||||
{"image/jpeg", true},
|
{testMIMEJPEG, true},
|
||||||
{"image/png", true},
|
{testMIMEPNG, true},
|
||||||
{"image/webp", true},
|
{testMIMEWebP, true},
|
||||||
{"image/gif", true},
|
{testMIMEGIF, true},
|
||||||
{"image/avif", true},
|
{testMIMEAVIF, true},
|
||||||
{"image/svg+xml", true},
|
{"image/svg+xml", true},
|
||||||
{"IMAGE/JPEG", true},
|
{"IMAGE/JPEG", true},
|
||||||
{"image/jpeg; charset=utf-8", true},
|
{testMIMEJPEGParams, true},
|
||||||
{"image/tiff", false},
|
{"image/tiff", false},
|
||||||
{"image/bmp", false},
|
{"image/bmp", false},
|
||||||
{"application/octet-stream", false},
|
{mimeOctetStream, false},
|
||||||
{"text/plain", false},
|
{"text/plain", false},
|
||||||
{"", false},
|
{"", false},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.mimeType, func(t *testing.T) {
|
t.Run(tt.mimeType, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
if got := IsSupportedMIMEType(tt.mimeType); got != tt.want {
|
if got := IsSupportedMIMEType(tt.mimeType); got != tt.want {
|
||||||
t.Errorf("IsSupportedMIMEType(%q) = %v, want %v", tt.mimeType, got, tt.want)
|
t.Errorf("IsSupportedMIMEType(%q) = %v, want %v", tt.mimeType, got, tt.want)
|
||||||
}
|
}
|
||||||
@@ -212,8 +189,16 @@ func TestIsSupportedMIMEType(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestPeekAndValidate(t *testing.T) {
|
func TestPeekAndValidate(t *testing.T) {
|
||||||
jpegData := append([]byte{0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01}, []byte("rest of jpeg data")...)
|
t.Parallel()
|
||||||
pngData := append([]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D}, []byte("rest of png data")...)
|
|
||||||
|
jpegMagic := []byte{
|
||||||
|
0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46, 0x49, 0x46, 0x00, 0x01,
|
||||||
|
}
|
||||||
|
pngMagic := []byte{
|
||||||
|
0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D,
|
||||||
|
}
|
||||||
|
jpegData := slices.Concat(jpegMagic, []byte("rest of jpeg data"))
|
||||||
|
pngData := slices.Concat(pngMagic, []byte("rest of png data"))
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
@@ -225,30 +210,32 @@ func TestPeekAndValidate(t *testing.T) {
|
|||||||
{
|
{
|
||||||
name: "valid JPEG",
|
name: "valid JPEG",
|
||||||
data: jpegData,
|
data: jpegData,
|
||||||
declaredType: "image/jpeg",
|
declaredType: testMIMEJPEG,
|
||||||
wantErr: false,
|
wantErr: false,
|
||||||
wantData: jpegData,
|
wantData: jpegData,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "valid PNG",
|
name: "valid PNG",
|
||||||
data: pngData,
|
data: pngData,
|
||||||
declaredType: "image/png",
|
declaredType: testMIMEPNG,
|
||||||
wantErr: false,
|
wantErr: false,
|
||||||
wantData: pngData,
|
wantData: pngData,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "mismatched type",
|
name: "mismatched type",
|
||||||
data: jpegData,
|
data: jpegData,
|
||||||
declaredType: "image/png",
|
declaredType: testMIMEPNG,
|
||||||
wantErr: true,
|
wantErr: true,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
r := bytes.NewReader(tt.data)
|
t.Parallel()
|
||||||
result, err := PeekAndValidate(r, tt.declaredType)
|
|
||||||
|
|
||||||
|
r := bytes.NewReader(tt.data)
|
||||||
|
|
||||||
|
result, err := PeekAndValidate(r, tt.declaredType)
|
||||||
if tt.wantErr {
|
if tt.wantErr {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("PeekAndValidate() expected error, got nil")
|
t.Error("PeekAndValidate() expected error, got nil")
|
||||||
@@ -272,23 +259,28 @@ func TestPeekAndValidate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !bytes.Equal(got, tt.wantData) {
|
if !bytes.Equal(got, tt.wantData) {
|
||||||
t.Errorf("PeekAndValidate() data mismatch: got %d bytes, want %d bytes", len(got), len(tt.wantData))
|
t.Errorf(
|
||||||
|
"PeekAndValidate() data mismatch: got %d bytes, want %d bytes",
|
||||||
|
len(got), len(tt.wantData),
|
||||||
|
)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestMIMEToImageFormat(t *testing.T) {
|
func TestMIMEToImageFormat(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
mimeType string
|
mimeType string
|
||||||
wantFormat ImageFormat
|
wantFormat ImageFormat
|
||||||
wantOk bool
|
wantOk bool
|
||||||
}{
|
}{
|
||||||
{"image/jpeg", FormatJPEG, true},
|
{testMIMEJPEG, FormatJPEG, true},
|
||||||
{"image/png", FormatPNG, true},
|
{testMIMEPNG, FormatPNG, true},
|
||||||
{"image/webp", FormatWebP, true},
|
{testMIMEWebP, FormatWebP, true},
|
||||||
{"image/gif", FormatGIF, true},
|
{testMIMEGIF, FormatGIF, true},
|
||||||
{"image/avif", FormatAVIF, true},
|
{testMIMEAVIF, FormatAVIF, true},
|
||||||
{"image/svg+xml", "", false}, // SVG doesn't convert to ImageFormat
|
{"image/svg+xml", "", false}, // SVG doesn't convert to ImageFormat
|
||||||
{"image/tiff", "", false},
|
{"image/tiff", "", false},
|
||||||
{"text/plain", "", false},
|
{"text/plain", "", false},
|
||||||
@@ -296,6 +288,8 @@ func TestMIMEToImageFormat(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.mimeType, func(t *testing.T) {
|
t.Run(tt.mimeType, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got, ok := MIMEToImageFormat(tt.mimeType)
|
got, ok := MIMEToImageFormat(tt.mimeType)
|
||||||
|
|
||||||
if ok != tt.wantOk {
|
if ok != tt.wantOk {
|
||||||
@@ -310,21 +304,25 @@ func TestMIMEToImageFormat(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestImageFormatToMIME(t *testing.T) {
|
func TestImageFormatToMIME(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
format ImageFormat
|
format ImageFormat
|
||||||
wantMIME string
|
wantMIME string
|
||||||
}{
|
}{
|
||||||
{FormatJPEG, "image/jpeg"},
|
{FormatJPEG, testMIMEJPEG},
|
||||||
{FormatPNG, "image/png"},
|
{FormatPNG, testMIMEPNG},
|
||||||
{FormatWebP, "image/webp"},
|
{FormatWebP, testMIMEWebP},
|
||||||
{FormatGIF, "image/gif"},
|
{FormatGIF, testMIMEGIF},
|
||||||
{FormatAVIF, "image/avif"},
|
{FormatAVIF, testMIMEAVIF},
|
||||||
{FormatOriginal, "application/octet-stream"},
|
{FormatOriginal, mimeOctetStream},
|
||||||
{"unknown", "application/octet-stream"},
|
{"unknown", mimeOctetStream},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(string(tt.format), func(t *testing.T) {
|
t.Run(string(tt.format), func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := ImageFormatToMIME(tt.format)
|
got := ImageFormatToMIME(tt.format)
|
||||||
|
|
||||||
if got != tt.wantMIME {
|
if got != tt.wantMIME {
|
||||||
@@ -335,19 +333,23 @@ func TestImageFormatToMIME(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestNormalizeMIMEType(t *testing.T) {
|
func TestNormalizeMIMEType(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
input string
|
input string
|
||||||
want string
|
want string
|
||||||
}{
|
}{
|
||||||
{"image/jpeg", "image/jpeg"},
|
{testMIMEJPEG, testMIMEJPEG},
|
||||||
{"IMAGE/JPEG", "image/jpeg"},
|
{"IMAGE/JPEG", testMIMEJPEG},
|
||||||
{"image/jpeg; charset=utf-8", "image/jpeg"},
|
{testMIMEJPEGParams, testMIMEJPEG},
|
||||||
{" image/jpeg ", "image/jpeg"},
|
{" image/jpeg ", testMIMEJPEG},
|
||||||
{"image/jpeg; boundary=something", "image/jpeg"},
|
{"image/jpeg; boundary=something", testMIMEJPEG},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.input, func(t *testing.T) {
|
t.Run(tt.input, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := normalizeMIMEType(tt.input)
|
got := normalizeMIMEType(tt.input)
|
||||||
|
|
||||||
if got != tt.want {
|
if got != tt.want {
|
||||||
@@ -358,6 +360,8 @@ func TestNormalizeMIMEType(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDetectSVG(t *testing.T) {
|
func TestDetectSVG(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
data string
|
data string
|
||||||
@@ -365,17 +369,24 @@ func TestDetectSVG(t *testing.T) {
|
|||||||
}{
|
}{
|
||||||
{"xml declaration", `<?xml version="1.0"?><svg></svg>`, true},
|
{"xml declaration", `<?xml version="1.0"?><svg></svg>`, true},
|
||||||
{"svg element", `<svg xmlns="http://www.w3.org/2000/svg"></svg>`, true},
|
{"svg element", `<svg xmlns="http://www.w3.org/2000/svg"></svg>`, true},
|
||||||
{"doctype", `<!DOCTYPE svg PUBLIC "-//W3C//DTD SVG 1.1//EN" "http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd">`, true},
|
{
|
||||||
|
"doctype",
|
||||||
|
`<!DOCTYPE svg PUBLIC "-//W3C//DTD SVG 1.1//EN" ` +
|
||||||
|
`"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd">`,
|
||||||
|
true,
|
||||||
|
},
|
||||||
{"with whitespace", `
|
{"with whitespace", `
|
||||||
<?xml version="1.0"?><svg></svg>`, true},
|
<?xml version="1.0"?><svg></svg>`, true},
|
||||||
{"uppercase", `<SVG></SVG>`, true},
|
{"uppercase", `<SVG></SVG>`, true},
|
||||||
{"not svg", `<html></html>`, false},
|
{"not svg", `<html></html>`, false},
|
||||||
{"random text", `hello world`, false},
|
{"random text", `hello world`, false},
|
||||||
{"empty", ``, false},
|
{testNameEmpty, ``, false},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := detectSVG([]byte(tt.data))
|
got := detectSVG([]byte(tt.data))
|
||||||
|
|
||||||
if got != tt.want {
|
if got != tt.want {
|
||||||
@@ -386,6 +397,8 @@ func TestDetectSVG(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestSkipBOM(t *testing.T) {
|
func TestSkipBOM(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
data []byte
|
data []byte
|
||||||
@@ -393,13 +406,15 @@ func TestSkipBOM(t *testing.T) {
|
|||||||
}{
|
}{
|
||||||
{"with BOM", []byte{0xEF, 0xBB, 0xBF, 'h', 'e', 'l', 'l', 'o'}, []byte("hello")},
|
{"with BOM", []byte{0xEF, 0xBB, 0xBF, 'h', 'e', 'l', 'l', 'o'}, []byte("hello")},
|
||||||
{"without BOM", []byte("hello"), []byte("hello")},
|
{"without BOM", []byte("hello"), []byte("hello")},
|
||||||
{"empty", []byte{}, []byte{}},
|
{testNameEmpty, []byte{}, []byte{}},
|
||||||
{"only BOM", []byte{0xEF, 0xBB, 0xBF}, []byte{}},
|
{"only BOM", []byte{0xEF, 0xBB, 0xBF}, []byte{}},
|
||||||
{"partial BOM", []byte{0xEF, 0xBB}, []byte{0xEF, 0xBB}},
|
{"partial BOM", []byte{0xEF, 0xBB}, []byte{0xEF, 0xBB}},
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := skipBOM(tt.data)
|
got := skipBOM(tt.data)
|
||||||
|
|
||||||
if !bytes.Equal(got, tt.want) {
|
if !bytes.Equal(got, tt.want) {
|
||||||
@@ -410,6 +425,8 @@ func TestSkipBOM(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRealWorldSVGPatterns(t *testing.T) {
|
func TestRealWorldSVGPatterns(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Test various real-world SVG patterns
|
// Test various real-world SVG patterns
|
||||||
svgPatterns := []string{
|
svgPatterns := []string{
|
||||||
`<?xml version="1.0" encoding="UTF-8"?>
|
`<?xml version="1.0" encoding="UTF-8"?>
|
||||||
@@ -419,7 +436,8 @@ func TestRealWorldSVGPatterns(t *testing.T) {
|
|||||||
`<svg viewBox="0 0 24 24" fill="none" xmlns="http://www.w3.org/2000/svg">
|
`<svg viewBox="0 0 24 24" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||||
<path d="M12 2L2 7l10 5 10-5-10-5z"/>
|
<path d="M12 2L2 7l10 5 10-5-10-5z"/>
|
||||||
</svg>`,
|
</svg>`,
|
||||||
`<!DOCTYPE svg PUBLIC "-//W3C//DTD SVG 1.1//EN" "http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd">
|
`<!DOCTYPE svg PUBLIC "-//W3C//DTD SVG 1.1//EN" ` +
|
||||||
|
`"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd">` + `
|
||||||
<svg xmlns="http://www.w3.org/2000/svg">
|
<svg xmlns="http://www.w3.org/2000/svg">
|
||||||
</svg>`,
|
</svg>`,
|
||||||
}
|
}
|
||||||
@@ -445,6 +463,8 @@ func TestRealWorldSVGPatterns(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDetectFormatRIFFNotWebP(t *testing.T) {
|
func TestDetectFormatRIFFNotWebP(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// RIFF container but not WebP (e.g., WAV file)
|
// RIFF container but not WebP (e.g., WAV file)
|
||||||
wavData := []byte{
|
wavData := []byte{
|
||||||
0x52, 0x49, 0x46, 0x46, // RIFF
|
0x52, 0x49, 0x46, 0x46, // RIFF
|
||||||
@@ -453,12 +473,14 @@ func TestDetectFormatRIFFNotWebP(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
_, err := DetectFormat(wavData)
|
_, err := DetectFormat(wavData)
|
||||||
if err != ErrUnknownFormat {
|
if !errors.Is(err, ErrUnknownFormat) {
|
||||||
t.Errorf("DetectFormat(WAV) error = %v, want %v", err, ErrUnknownFormat)
|
t.Errorf("DetectFormat(WAV) error = %v, want %v", err, ErrUnknownFormat)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDetectFormatFtypNotAVIF(t *testing.T) {
|
func TestDetectFormatFtypNotAVIF(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// ftyp container but not AVIF (e.g., MP4)
|
// ftyp container but not AVIF (e.g., MP4)
|
||||||
mp4Data := []byte{
|
mp4Data := []byte{
|
||||||
0x00, 0x00, 0x00, 0x1C, // box size
|
0x00, 0x00, 0x00, 0x1C, // box size
|
||||||
@@ -467,20 +489,24 @@ func TestDetectFormatFtypNotAVIF(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
_, err := DetectFormat(mp4Data)
|
_, err := DetectFormat(mp4Data)
|
||||||
if err != ErrUnknownFormat {
|
if !errors.Is(err, ErrUnknownFormat) {
|
||||||
t.Errorf("DetectFormat(MP4) error = %v, want %v", err, ErrUnknownFormat)
|
t.Errorf("DetectFormat(MP4) error = %v, want %v", err, ErrUnknownFormat)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPeekAndValidatePreservesReader(t *testing.T) {
|
func TestPeekAndValidatePreservesReader(t *testing.T) {
|
||||||
// Ensure that after PeekAndValidate, we can read the complete original content
|
t.Parallel()
|
||||||
|
|
||||||
|
// Ensure that after PeekAndValidate, we can read the complete
|
||||||
|
// original content
|
||||||
originalContent := append(
|
originalContent := append(
|
||||||
[]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D},
|
[]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D},
|
||||||
[]byte(strings.Repeat("PNG IDAT chunk data here ", 100))...,
|
[]byte(strings.Repeat("PNG IDAT chunk data here ", 100))...,
|
||||||
)
|
)
|
||||||
|
|
||||||
r := bytes.NewReader(originalContent)
|
r := bytes.NewReader(originalContent)
|
||||||
validated, err := PeekAndValidate(r, "image/png")
|
|
||||||
|
validated, err := PeekAndValidate(r, testMIMEPNG)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("PeekAndValidate() error = %v", err)
|
t.Fatalf("PeekAndValidate() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -492,6 +518,9 @@ func TestPeekAndValidatePreservesReader(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !bytes.Equal(got, originalContent) {
|
if !bytes.Equal(got, originalContent) {
|
||||||
t.Errorf("Content mismatch: got %d bytes, want %d bytes", len(got), len(originalContent))
|
t.Errorf(
|
||||||
|
"Content mismatch: got %d bytes, want %d bytes",
|
||||||
|
len(got), len(originalContent),
|
||||||
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -24,6 +24,7 @@ const CORSMaxAgeSeconds = 86400
|
|||||||
// Params defines dependencies for Middleware.
|
// Params defines dependencies for Middleware.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
}
|
}
|
||||||
@@ -49,6 +50,7 @@ func ipFromHostPort(hp string) string {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(h) > 0 && h[0] == '[' {
|
if len(h) > 0 && h[0] == '[' {
|
||||||
return h[1 : len(h)-1]
|
return h[1 : len(h)-1]
|
||||||
}
|
}
|
||||||
@@ -58,6 +60,7 @@ func ipFromHostPort(hp string) string {
|
|||||||
|
|
||||||
type loggingResponseWriter struct {
|
type loggingResponseWriter struct {
|
||||||
http.ResponseWriter
|
http.ResponseWriter
|
||||||
|
|
||||||
statusCode int
|
statusCode int
|
||||||
bytesWritten int64
|
bytesWritten int64
|
||||||
}
|
}
|
||||||
@@ -85,6 +88,7 @@ func (s *Middleware) Logging() func(http.Handler) http.Handler {
|
|||||||
start := time.Now()
|
start := time.Now()
|
||||||
lrw := newLoggingResponseWriter(w)
|
lrw := newLoggingResponseWriter(w)
|
||||||
ctx := r.Context()
|
ctx := r.Context()
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
latency := time.Since(start)
|
latency := time.Since(start)
|
||||||
reqID, _ := ctx.Value(middleware.RequestIDKey).(string)
|
reqID, _ := ctx.Value(middleware.RequestIDKey).(string)
|
||||||
|
|||||||
@@ -10,6 +10,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestSecurityHeaders(t *testing.T) {
|
func TestSecurityHeaders(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
// Create middleware instance
|
// Create middleware instance
|
||||||
cfg := &config.Config{}
|
cfg := &config.Config{}
|
||||||
mw := &Middleware{
|
mw := &Middleware{
|
||||||
@@ -26,7 +28,7 @@ func TestSecurityHeaders(t *testing.T) {
|
|||||||
handler := mw.SecurityHeaders()(testHandler)
|
handler := mw.SecurityHeaders()(testHandler)
|
||||||
|
|
||||||
// Make a test request
|
// Make a test request
|
||||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/test", nil)
|
||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
handler.ServeHTTP(rec, req)
|
handler.ServeHTTP(rec, req)
|
||||||
@@ -44,6 +46,8 @@ func TestSecurityHeaders(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.header, func(t *testing.T) {
|
t.Run(tt.header, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
got := rec.Header().Get(tt.header)
|
got := rec.Header().Get(tt.header)
|
||||||
if got != tt.want {
|
if got != tt.want {
|
||||||
t.Errorf("%s = %q, want %q", tt.header, got, tt.want)
|
t.Errorf("%s = %q, want %q", tt.header, got, tt.want)
|
||||||
@@ -53,6 +57,8 @@ func TestSecurityHeaders(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestSecurityHeaders_PreservesExistingHeaders(t *testing.T) {
|
func TestSecurityHeaders_PreservesExistingHeaders(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
cfg := &config.Config{}
|
cfg := &config.Config{}
|
||||||
mw := &Middleware{
|
mw := &Middleware{
|
||||||
log: slog.Default(),
|
log: slog.Default(),
|
||||||
@@ -68,7 +74,7 @@ func TestSecurityHeaders_PreservesExistingHeaders(t *testing.T) {
|
|||||||
|
|
||||||
handler := mw.SecurityHeaders()(testHandler)
|
handler := mw.SecurityHeaders()(testHandler)
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/test", nil)
|
||||||
rec := httptest.NewRecorder()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
handler.ServeHTTP(rec, req)
|
handler.ServeHTTP(rec, req)
|
||||||
@@ -34,7 +34,8 @@ func DeriveKey(masterKey []byte, salt string) ([KeySize]byte, error) {
|
|||||||
|
|
||||||
hkdfReader := hkdf.New(sha256.New, masterKey, []byte(salt), nil)
|
hkdfReader := hkdf.New(sha256.New, masterKey, []byte(salt), nil)
|
||||||
|
|
||||||
if _, err := io.ReadFull(hkdfReader, key[:]); err != nil {
|
_, err := io.ReadFull(hkdfReader, key[:])
|
||||||
|
if err != nil {
|
||||||
return key, ErrKeyDerivation
|
return key, ErrKeyDerivation
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -46,7 +47,9 @@ func DeriveKey(masterKey []byte, salt string) ([KeySize]byte, error) {
|
|||||||
func Encrypt(key [KeySize]byte, plaintext []byte) (string, error) {
|
func Encrypt(key [KeySize]byte, plaintext []byte) (string, error) {
|
||||||
// Generate random nonce
|
// Generate random nonce
|
||||||
var nonce [NonceSize]byte
|
var nonce [NonceSize]byte
|
||||||
if _, err := rand.Read(nonce[:]); err != nil {
|
|
||||||
|
_, err := rand.Read(nonce[:])
|
||||||
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,20 +1,25 @@
|
|||||||
package seal
|
package seal_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/seal"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestDeriveKey_Consistent(t *testing.T) {
|
func TestDeriveKey_Consistent(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
masterKey := []byte("test-master-key-12345")
|
masterKey := []byte("test-master-key-12345")
|
||||||
salt := "test-salt-v1"
|
salt := "test-salt-v1"
|
||||||
|
|
||||||
key1, err := DeriveKey(masterKey, salt)
|
key1, err := seal.DeriveKey(masterKey, salt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
key2, err := DeriveKey(masterKey, salt)
|
key2, err := seal.DeriveKey(masterKey, salt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -25,14 +30,16 @@ func TestDeriveKey_Consistent(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDeriveKey_DifferentSalts(t *testing.T) {
|
func TestDeriveKey_DifferentSalts(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
masterKey := []byte("test-master-key-12345")
|
masterKey := []byte("test-master-key-12345")
|
||||||
|
|
||||||
key1, err := DeriveKey(masterKey, "salt-1")
|
key1, err := seal.DeriveKey(masterKey, "salt-1")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
key2, err := DeriveKey(masterKey, "salt-2")
|
key2, err := seal.DeriveKey(masterKey, "salt-2")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -43,14 +50,16 @@ func TestDeriveKey_DifferentSalts(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDeriveKey_DifferentMasterKeys(t *testing.T) {
|
func TestDeriveKey_DifferentMasterKeys(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
salt := "test-salt"
|
salt := "test-salt"
|
||||||
|
|
||||||
key1, err := DeriveKey([]byte("master-key-1"), salt)
|
key1, err := seal.DeriveKey([]byte("master-key-1"), salt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
key2, err := DeriveKey([]byte("master-key-2"), salt)
|
key2, err := seal.DeriveKey([]byte("master-key-2"), salt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -61,19 +70,21 @@ func TestDeriveKey_DifferentMasterKeys(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestEncryptDecrypt_RoundTrip(t *testing.T) {
|
func TestEncryptDecrypt_RoundTrip(t *testing.T) {
|
||||||
key, err := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
|
key, err := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("DeriveKey() error = %v", err)
|
t.Fatalf("DeriveKey() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
plaintext := []byte("hello, world! this is a test message.")
|
plaintext := []byte("hello, world! this is a test message.")
|
||||||
|
|
||||||
ciphertext, err := Encrypt(key, plaintext)
|
ciphertext, err := seal.Encrypt(key, plaintext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Encrypt() error = %v", err)
|
t.Fatalf("Encrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
decrypted, err := Decrypt(key, ciphertext)
|
decrypted, err := seal.Decrypt(key, ciphertext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Decrypt() error = %v", err)
|
t.Fatalf("Decrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -84,15 +95,17 @@ func TestEncryptDecrypt_RoundTrip(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestEncryptDecrypt_EmptyPlaintext(t *testing.T) {
|
func TestEncryptDecrypt_EmptyPlaintext(t *testing.T) {
|
||||||
key, _ := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
|
key, _ := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
plaintext := []byte{}
|
plaintext := []byte{}
|
||||||
|
|
||||||
ciphertext, err := Encrypt(key, plaintext)
|
ciphertext, err := seal.Encrypt(key, plaintext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Encrypt() error = %v", err)
|
t.Fatalf("Encrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
decrypted, err := Decrypt(key, ciphertext)
|
decrypted, err := seal.Decrypt(key, ciphertext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Decrypt() error = %v", err)
|
t.Fatalf("Decrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -103,31 +116,35 @@ func TestEncryptDecrypt_EmptyPlaintext(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDecrypt_WrongKey(t *testing.T) {
|
func TestDecrypt_WrongKey(t *testing.T) {
|
||||||
key1, _ := DeriveKey([]byte("key-1"), "salt")
|
t.Parallel()
|
||||||
key2, _ := DeriveKey([]byte("key-2"), "salt")
|
|
||||||
|
key1, _ := seal.DeriveKey([]byte("key-1"), "salt")
|
||||||
|
key2, _ := seal.DeriveKey([]byte("key-2"), "salt")
|
||||||
|
|
||||||
plaintext := []byte("secret message")
|
plaintext := []byte("secret message")
|
||||||
|
|
||||||
ciphertext, err := Encrypt(key1, plaintext)
|
ciphertext, err := seal.Encrypt(key1, plaintext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Encrypt() error = %v", err)
|
t.Fatalf("Encrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = Decrypt(key2, ciphertext)
|
_, err = seal.Decrypt(key2, ciphertext)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("Decrypt() should fail with wrong key")
|
t.Error("Decrypt() should fail with wrong key")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrDecryptionFailed {
|
if !errors.Is(err, seal.ErrDecryptionFailed) {
|
||||||
t.Errorf("Decrypt() error = %v, want %v", err, ErrDecryptionFailed)
|
t.Errorf("Decrypt() error = %v, want %v", err, seal.ErrDecryptionFailed)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDecrypt_TamperedCiphertext(t *testing.T) {
|
func TestDecrypt_TamperedCiphertext(t *testing.T) {
|
||||||
key, _ := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
|
key, _ := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
plaintext := []byte("secret message")
|
plaintext := []byte("secret message")
|
||||||
|
|
||||||
ciphertext, err := Encrypt(key, plaintext)
|
ciphertext, err := seal.Encrypt(key, plaintext)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Encrypt() error = %v", err)
|
t.Fatalf("Encrypt() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -138,45 +155,51 @@ func TestDecrypt_TamperedCiphertext(t *testing.T) {
|
|||||||
tampered[10] ^= 0x01
|
tampered[10] ^= 0x01
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = Decrypt(key, string(tampered))
|
_, err = seal.Decrypt(key, string(tampered))
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("Decrypt() should fail with tampered ciphertext")
|
t.Error("Decrypt() should fail with tampered ciphertext")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDecrypt_InvalidBase64(t *testing.T) {
|
func TestDecrypt_InvalidBase64(t *testing.T) {
|
||||||
key, _ := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
_, err := Decrypt(key, "not-valid-base64!!!")
|
key, _ := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
|
|
||||||
|
_, err := seal.Decrypt(key, "not-valid-base64!!!")
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("Decrypt() should fail with invalid base64")
|
t.Error("Decrypt() should fail with invalid base64")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrInvalidPayload {
|
if !errors.Is(err, seal.ErrInvalidPayload) {
|
||||||
t.Errorf("Decrypt() error = %v, want %v", err, ErrInvalidPayload)
|
t.Errorf("Decrypt() error = %v, want %v", err, seal.ErrInvalidPayload)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDecrypt_TooShort(t *testing.T) {
|
func TestDecrypt_TooShort(t *testing.T) {
|
||||||
key, _ := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
|
key, _ := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
|
|
||||||
// Create a base64 string that's too short to contain nonce + auth tag
|
// Create a base64 string that's too short to contain nonce + auth tag
|
||||||
_, err := Decrypt(key, "dG9vLXNob3J0")
|
_, err := seal.Decrypt(key, "dG9vLXNob3J0")
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("Decrypt() should fail with too-short ciphertext")
|
t.Error("Decrypt() should fail with too-short ciphertext")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrInvalidPayload {
|
if !errors.Is(err, seal.ErrInvalidPayload) {
|
||||||
t.Errorf("Decrypt() error = %v, want %v", err, ErrInvalidPayload)
|
t.Errorf("Decrypt() error = %v, want %v", err, seal.ErrInvalidPayload)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEncrypt_ProducesDifferentCiphertexts(t *testing.T) {
|
func TestEncrypt_ProducesDifferentCiphertexts(t *testing.T) {
|
||||||
key, _ := DeriveKey([]byte("test-key"), "test-salt")
|
t.Parallel()
|
||||||
|
|
||||||
|
key, _ := seal.DeriveKey([]byte("test-key"), "test-salt")
|
||||||
plaintext := []byte("same message")
|
plaintext := []byte("same message")
|
||||||
|
|
||||||
ciphertext1, _ := Encrypt(key, plaintext)
|
ciphertext1, _ := seal.Encrypt(key, plaintext)
|
||||||
ciphertext2, _ := Encrypt(key, plaintext)
|
ciphertext2, _ := seal.Encrypt(key, plaintext)
|
||||||
|
|
||||||
if ciphertext1 == ciphertext2 {
|
if ciphertext1 == ciphertext2 {
|
||||||
t.Error("Encrypt() should produce different ciphertexts due to random nonce")
|
t.Error("Encrypt() should produce different ciphertexts due to random nonce")
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package server
|
package server
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
@@ -26,8 +27,11 @@ func (s *Server) serveUntilShutdown() {
|
|||||||
s.SetupRoutes()
|
s.SetupRoutes()
|
||||||
|
|
||||||
s.log.Info("http begin listen", "listenaddr", listenAddr)
|
s.log.Info("http begin listen", "listenaddr", listenAddr)
|
||||||
if err := s.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
|
||||||
|
err := s.httpServer.ListenAndServe()
|
||||||
|
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
s.log.Error("listen error", "error", err)
|
s.log.Error("listen error", "error", err)
|
||||||
|
|
||||||
if s.cancelFunc != nil {
|
if s.cancelFunc != nil {
|
||||||
s.cancelFunc()
|
s.cancelFunc()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -56,7 +56,8 @@ func (s *Server) SetupRoutes() {
|
|||||||
s.router.Head("/v1/image/*", s.h.HandleImage())
|
s.router.Head("/v1/image/*", s.h.HandleImage())
|
||||||
|
|
||||||
// Encrypted image URL route
|
// Encrypted image URL route
|
||||||
// The trailing filename (e.g., /img.jpg) is ignored but helps browsers with content type
|
// The trailing filename (e.g., /img.jpg) is ignored but helps
|
||||||
|
// browsers with content type
|
||||||
s.router.Get("/v1/e/{token}/*", s.h.HandleImageEnc())
|
s.router.Get("/v1/e/{token}/*", s.h.HandleImageEnc())
|
||||||
|
|
||||||
// Metrics endpoint with auth
|
// Metrics endpoint with auth
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ const (
|
|||||||
// Params defines dependencies for Server.
|
// Params defines dependencies for Server.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
fx.In
|
fx.In
|
||||||
|
|
||||||
Logger *logger.Logger
|
Logger *logger.Logger
|
||||||
Globals *globals.Globals
|
Globals *globals.Globals
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
@@ -47,7 +48,6 @@ type Server struct {
|
|||||||
startupTime time.Time
|
startupTime time.Time
|
||||||
exitCode int
|
exitCode int
|
||||||
sentryEnabled bool
|
sentryEnabled bool
|
||||||
ctx context.Context
|
|
||||||
cancelFunc context.CancelFunc
|
cancelFunc context.CancelFunc
|
||||||
httpServer *http.Server
|
httpServer *http.Server
|
||||||
router *chi.Mux
|
router *chi.Mux
|
||||||
@@ -64,9 +64,9 @@ func New(lc fx.Lifecycle, params Params) (*Server, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
lc.Append(fx.Hook{
|
lc.Append(fx.Hook{
|
||||||
OnStart: func(_ context.Context) error {
|
OnStart: func(ctx context.Context) error {
|
||||||
s.startupTime = time.Now()
|
s.startupTime = time.Now()
|
||||||
go s.Run()
|
go s.Run(context.WithoutCancel(ctx))
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
@@ -83,9 +83,14 @@ func New(lc fx.Lifecycle, params Params) (*Server, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Run starts the server.
|
// Run starts the server.
|
||||||
func (s *Server) Run() {
|
func (s *Server) Run(ctx context.Context) {
|
||||||
s.enableSentry()
|
s.enableSentry()
|
||||||
s.serve()
|
s.serve(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaintenanceMode returns whether maintenance mode is enabled.
|
||||||
|
func (s *Server) MaintenanceMode() bool {
|
||||||
|
return s.config.MaintenanceMode
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) enableSentry() {
|
func (s *Server) enableSentry() {
|
||||||
@@ -103,19 +108,24 @@ func (s *Server) enableSentry() {
|
|||||||
s.log.Error("sentry init failure", "error", err)
|
s.log.Error("sentry init failure", "error", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
s.log.Info("sentry error reporting activated")
|
s.log.Info("sentry error reporting activated")
|
||||||
s.sentryEnabled = true
|
s.sentryEnabled = true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) serve() int {
|
func (s *Server) serve(ctx context.Context) int {
|
||||||
s.ctx, s.cancelFunc = context.WithCancel(context.Background())
|
ctx, cancelFunc := context.WithCancel(ctx)
|
||||||
|
s.cancelFunc = cancelFunc
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
c := make(chan os.Signal, 1)
|
c := make(chan os.Signal, 1)
|
||||||
|
|
||||||
signal.Ignore(syscall.SIGPIPE)
|
signal.Ignore(syscall.SIGPIPE)
|
||||||
signal.Notify(c, os.Interrupt, syscall.SIGTERM)
|
signal.Notify(c, os.Interrupt, syscall.SIGTERM)
|
||||||
|
|
||||||
sig := <-c
|
sig := <-c
|
||||||
s.log.Info("signal received", "signal", sig)
|
s.log.Info("signal received", "signal", sig)
|
||||||
|
|
||||||
if s.cancelFunc != nil {
|
if s.cancelFunc != nil {
|
||||||
s.cancelFunc()
|
s.cancelFunc()
|
||||||
}
|
}
|
||||||
@@ -123,19 +133,22 @@ func (s *Server) serve() int {
|
|||||||
|
|
||||||
go s.serveUntilShutdown()
|
go s.serveUntilShutdown()
|
||||||
|
|
||||||
<-s.ctx.Done()
|
<-ctx.Done()
|
||||||
s.cleanShutdown()
|
s.cleanShutdown(ctx)
|
||||||
|
|
||||||
return s.exitCode
|
return s.exitCode
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) cleanShutdown() {
|
func (s *Server) cleanShutdown(ctx context.Context) {
|
||||||
s.exitCode = 0
|
s.exitCode = 0
|
||||||
ctxShutdown, shutdownCancel := context.WithTimeout(context.Background(), ShutdownTimeout)
|
|
||||||
|
ctxShutdown, shutdownCancel := context.WithTimeout(
|
||||||
|
context.WithoutCancel(ctx), ShutdownTimeout)
|
||||||
defer shutdownCancel()
|
defer shutdownCancel()
|
||||||
|
|
||||||
if s.httpServer != nil {
|
if s.httpServer != nil {
|
||||||
if err := s.httpServer.Shutdown(ctxShutdown); err != nil {
|
err := s.httpServer.Shutdown(ctxShutdown)
|
||||||
|
if err != nil {
|
||||||
s.log.Error("server clean shutdown failed", "error", err)
|
s.log.Error("server clean shutdown failed", "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -144,8 +157,3 @@ func (s *Server) cleanShutdown() {
|
|||||||
sentry.Flush(SentryFlushTimeout)
|
sentry.Flush(SentryFlushTimeout)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// MaintenanceMode returns whether maintenance mode is enabled.
|
|
||||||
func (s *Server) MaintenanceMode() bool {
|
|
||||||
return s.config.MaintenanceMode
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -107,7 +107,9 @@ func (m *Manager) ValidateSession(r *http.Request) (*Data, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var data Data
|
var data Data
|
||||||
if err := m.sc.Decode(CookieName, cookie.Value, &data); err != nil {
|
|
||||||
|
err = m.sc.Decode(CookieName, cookie.Value, &data)
|
||||||
|
if err != nil {
|
||||||
return nil, ErrInvalidSession
|
return nil, ErrInvalidSession
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
package session
|
package session_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
// TestSessionCookieAttributesAlwaysSecure verifies that every cookie
|
// TestSessionCookieAttributesAlwaysSecure verifies that every cookie
|
||||||
@@ -16,7 +18,9 @@ import (
|
|||||||
// This covers both cookie-writing paths: CreateSession (the login
|
// This covers both cookie-writing paths: CreateSession (the login
|
||||||
// set-cookie path) and ClearSession (the logout delete-cookie path).
|
// set-cookie path) and ClearSession (the logout delete-cookie path).
|
||||||
func TestSessionCookieAttributesAlwaysSecure(t *testing.T) {
|
func TestSessionCookieAttributesAlwaysSecure(t *testing.T) {
|
||||||
mgr, err := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
mgr, err := session.NewManager("test-signing-key-12345")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewManager() error = %v", err)
|
t.Fatalf("NewManager() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -29,7 +33,9 @@ func TestSessionCookieAttributesAlwaysSecure(t *testing.T) {
|
|||||||
name: "CreateSession",
|
name: "CreateSession",
|
||||||
setCookie: func(t *testing.T, w http.ResponseWriter) {
|
setCookie: func(t *testing.T, w http.ResponseWriter) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
if err := mgr.CreateSession(w); err != nil {
|
|
||||||
|
err := mgr.CreateSession(w)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("CreateSession() error = %v", err)
|
t.Fatalf("CreateSession() error = %v", err)
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
@@ -45,12 +51,15 @@ func TestSessionCookieAttributesAlwaysSecure(t *testing.T) {
|
|||||||
|
|
||||||
for _, writePath := range writePaths {
|
for _, writePath := range writePaths {
|
||||||
t.Run(writePath.name, func(t *testing.T) {
|
t.Run(writePath.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
writePath.setCookie(t, w)
|
writePath.setCookie(t, w)
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range w.Result().Cookies() {
|
for _, c := range w.Result().Cookies() {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
@@ -58,7 +67,7 @@ func TestSessionCookieAttributesAlwaysSecure(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if sessionCookie == nil {
|
if sessionCookie == nil {
|
||||||
t.Fatalf("no cookie named %q was set", CookieName)
|
t.Fatalf("no cookie named %q was set", session.CookieName)
|
||||||
}
|
}
|
||||||
|
|
||||||
t.Logf("cookie attributes: HttpOnly=%v Secure=%v SameSite=%v",
|
t.Logf("cookie attributes: HttpOnly=%v Secure=%v SameSite=%v",
|
||||||
|
|||||||
@@ -1,45 +1,55 @@
|
|||||||
package session
|
package session_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestManager_CreateAndValidate(t *testing.T) {
|
func TestManager_CreateAndValidate(t *testing.T) {
|
||||||
mgr, err := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
mgr, err := session.NewManager("test-signing-key-12345")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewManager() error = %v", err)
|
t.Fatalf("NewManager() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create a session
|
// Create a session
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
if err := mgr.CreateSession(w); err != nil {
|
|
||||||
|
err = mgr.CreateSession(w)
|
||||||
|
if err != nil {
|
||||||
t.Fatalf("CreateSession() error = %v", err)
|
t.Fatalf("CreateSession() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Extract the cookie from response
|
// Extract the cookie from response
|
||||||
resp := w.Result()
|
resp := w.Result()
|
||||||
|
|
||||||
cookies := resp.Cookies()
|
cookies := resp.Cookies()
|
||||||
if len(cookies) == 0 {
|
if len(cookies) == 0 {
|
||||||
t.Fatal("CreateSession() did not set a cookie")
|
t.Fatal("CreateSession() did not set a cookie")
|
||||||
}
|
}
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range cookies {
|
for _, c := range cookies {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if sessionCookie == nil {
|
if sessionCookie == nil {
|
||||||
t.Fatalf("CreateSession() did not set cookie named %q", CookieName)
|
t.Fatalf("CreateSession() did not set cookie named %q", session.CookieName)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate the session
|
// Validate the session
|
||||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
req.AddCookie(sessionCookie)
|
req.AddCookie(sessionCookie)
|
||||||
|
|
||||||
data, err := mgr.ValidateSession(req)
|
data, err := mgr.ValidateSession(req)
|
||||||
@@ -57,27 +67,34 @@ func TestManager_CreateAndValidate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_ValidateSession_NoCookie(t *testing.T) {
|
func TestManager_ValidateSession_NoCookie(t *testing.T) {
|
||||||
mgr, _ := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
mgr, _ := session.NewManager("test-signing-key-12345")
|
||||||
|
|
||||||
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
|
|
||||||
_, err := mgr.ValidateSession(req)
|
_, err := mgr.ValidateSession(req)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Error("ValidateSession() should fail with no cookie")
|
t.Error("ValidateSession() should fail with no cookie")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrNoSession {
|
if !errors.Is(err, session.ErrNoSession) {
|
||||||
t.Errorf("ValidateSession() error = %v, want %v", err, ErrNoSession)
|
t.Errorf("ValidateSession() error = %v, want %v", err, session.ErrNoSession)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_ValidateSession_TamperedCookie(t *testing.T) {
|
func TestManager_ValidateSession_TamperedCookie(t *testing.T) {
|
||||||
mgr, _ := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
mgr, _ := session.NewManager("test-signing-key-12345")
|
||||||
|
|
||||||
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
req.AddCookie(&http.Cookie{
|
req.AddCookie(&http.Cookie{
|
||||||
Name: CookieName,
|
Name: session.CookieName,
|
||||||
Value: "tampered-invalid-cookie-value",
|
Value: "tampered-invalid-cookie-value",
|
||||||
|
Secure: true,
|
||||||
|
HttpOnly: true,
|
||||||
|
SameSite: http.SameSiteStrictMode,
|
||||||
})
|
})
|
||||||
|
|
||||||
_, err := mgr.ValidateSession(req)
|
_, err := mgr.ValidateSession(req)
|
||||||
@@ -85,30 +102,35 @@ func TestManager_ValidateSession_TamperedCookie(t *testing.T) {
|
|||||||
t.Error("ValidateSession() should fail with tampered cookie")
|
t.Error("ValidateSession() should fail with tampered cookie")
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != ErrInvalidSession {
|
if !errors.Is(err, session.ErrInvalidSession) {
|
||||||
t.Errorf("ValidateSession() error = %v, want %v", err, ErrInvalidSession)
|
t.Errorf("ValidateSession() error = %v, want %v", err, session.ErrInvalidSession)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_ValidateSession_WrongKey(t *testing.T) {
|
func TestManager_ValidateSession_WrongKey(t *testing.T) {
|
||||||
mgr1, _ := NewManager("signing-key-1")
|
t.Parallel()
|
||||||
mgr2, _ := NewManager("signing-key-2")
|
|
||||||
|
mgr1, _ := session.NewManager("signing-key-1")
|
||||||
|
mgr2, _ := session.NewManager("signing-key-2")
|
||||||
|
|
||||||
// Create session with mgr1
|
// Create session with mgr1
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
_ = mgr1.CreateSession(w)
|
_ = mgr1.CreateSession(w)
|
||||||
|
|
||||||
resp := w.Result()
|
resp := w.Result()
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range resp.Cookies() {
|
for _, c := range resp.Cookies() {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try to validate with mgr2 (different key)
|
// Try to validate with mgr2 (different key)
|
||||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
req.AddCookie(sessionCookie)
|
req.AddCookie(sessionCookie)
|
||||||
|
|
||||||
_, err := mgr2.ValidateSession(req)
|
_, err := mgr2.ValidateSession(req)
|
||||||
@@ -118,7 +140,9 @@ func TestManager_ValidateSession_WrongKey(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_ClearSession(t *testing.T) {
|
func TestManager_ClearSession(t *testing.T) {
|
||||||
mgr, _ := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
mgr, _ := session.NewManager("test-signing-key-12345")
|
||||||
|
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
mgr.ClearSession(w)
|
mgr.ClearSession(w)
|
||||||
@@ -127,9 +151,11 @@ func TestManager_ClearSession(t *testing.T) {
|
|||||||
cookies := resp.Cookies()
|
cookies := resp.Cookies()
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range cookies {
|
for _, c := range cookies {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -144,10 +170,12 @@ func TestManager_ClearSession(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_IsAuthenticated(t *testing.T) {
|
func TestManager_IsAuthenticated(t *testing.T) {
|
||||||
mgr, _ := NewManager("test-signing-key-12345")
|
t.Parallel()
|
||||||
|
|
||||||
|
mgr, _ := session.NewManager("test-signing-key-12345")
|
||||||
|
|
||||||
// No session - should return false
|
// No session - should return false
|
||||||
req := httptest.NewRequest(http.MethodGet, "/", nil)
|
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
if mgr.IsAuthenticated(req) {
|
if mgr.IsAuthenticated(req) {
|
||||||
t.Error("IsAuthenticated() should return false with no session")
|
t.Error("IsAuthenticated() should return false with no session")
|
||||||
}
|
}
|
||||||
@@ -157,16 +185,19 @@ func TestManager_IsAuthenticated(t *testing.T) {
|
|||||||
_ = mgr.CreateSession(w)
|
_ = mgr.CreateSession(w)
|
||||||
|
|
||||||
resp := w.Result()
|
resp := w.Result()
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range resp.Cookies() {
|
for _, c := range resp.Cookies() {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// With valid session - should return true
|
// With valid session - should return true
|
||||||
req = httptest.NewRequest(http.MethodGet, "/", nil)
|
req = httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
|
||||||
req.AddCookie(sessionCookie)
|
req.AddCookie(sessionCookie)
|
||||||
|
|
||||||
if !mgr.IsAuthenticated(req) {
|
if !mgr.IsAuthenticated(req) {
|
||||||
@@ -175,16 +206,21 @@ func TestManager_IsAuthenticated(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestManager_CookieAttributes(t *testing.T) {
|
func TestManager_CookieAttributes(t *testing.T) {
|
||||||
mgr, _ := NewManager("test-key")
|
t.Parallel()
|
||||||
|
|
||||||
|
mgr, _ := session.NewManager("test-key")
|
||||||
|
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
_ = mgr.CreateSession(w)
|
_ = mgr.CreateSession(w)
|
||||||
|
|
||||||
resp := w.Result()
|
resp := w.Result()
|
||||||
|
|
||||||
var sessionCookie *http.Cookie
|
var sessionCookie *http.Cookie
|
||||||
|
|
||||||
for _, c := range resp.Cookies() {
|
for _, c := range resp.Cookies() {
|
||||||
if c.Name == CookieName {
|
if c.Name == session.CookieName {
|
||||||
sessionCookie = c
|
sessionCookie = c
|
||||||
|
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -198,6 +234,7 @@ func TestManager_CookieAttributes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if sessionCookie.SameSite != http.SameSiteStrictMode {
|
if sessionCookie.SameSite != http.SameSiteStrictMode {
|
||||||
t.Errorf("Cookie SameSite = %v, want %v", sessionCookie.SameSite, http.SameSiteStrictMode)
|
t.Errorf("Cookie SameSite = %v, want %v",
|
||||||
|
sessionCookie.SameSite, http.SameSiteStrictMode)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
package signature
|
package signature_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/signature"
|
||||||
)
|
)
|
||||||
|
|
||||||
// goldenExpiresUnix is the fixed expiration timestamp used by all golden
|
// goldenExpiresUnix is the fixed expiration timestamp used by all golden
|
||||||
@@ -12,10 +14,74 @@ const goldenExpiresUnix int64 = 1704067200
|
|||||||
// goldenSigningKey is the fixed signing key used by all golden vectors.
|
// goldenSigningKey is the fixed signing key used by all golden vectors.
|
||||||
const goldenSigningKey = "golden-test-key"
|
const goldenSigningKey = "golden-test-key"
|
||||||
|
|
||||||
|
type goldenVector struct {
|
||||||
|
name string
|
||||||
|
req signature.Request
|
||||||
|
// wantSignature is the exact base64url (RFC 4648 URL-safe,
|
||||||
|
// padded) HMAC-SHA256 signature for the request with Expires
|
||||||
|
// set to goldenExpiresUnix.
|
||||||
|
wantSignature string
|
||||||
|
// wantSignedPath is the exact path returned by
|
||||||
|
// GenerateSignedURL for the request. The signature and
|
||||||
|
// expiration are returned separately by GenerateSignedURL and
|
||||||
|
// are not embedded in the path.
|
||||||
|
wantSignedPath string
|
||||||
|
}
|
||||||
|
|
||||||
|
// goldenVectors returns the known-answer vectors. The expected values
|
||||||
|
// were computed once and are hardcoded here.
|
||||||
|
func goldenVectors() []goldenVector {
|
||||||
|
return []goldenVector{
|
||||||
|
{
|
||||||
|
name: "resized without query",
|
||||||
|
req: signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
|
SourceQuery: "",
|
||||||
|
Width: 800,
|
||||||
|
Height: 600,
|
||||||
|
Format: testFormatWebP,
|
||||||
|
},
|
||||||
|
// Signed data: "cdn.example.com:/photos/cat.jpg::800:600:webp:1704067200"
|
||||||
|
wantSignature: "x5PfPp8QSDo0cJT96od-AEgrQyOVLfqifH5sst61_-w=",
|
||||||
|
wantSignedPath: testSignedPath,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "resized with query string",
|
||||||
|
req: signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
|
SourceQuery: "token=abc&v=2",
|
||||||
|
Width: 800,
|
||||||
|
Height: 600,
|
||||||
|
Format: testFormatWebP,
|
||||||
|
},
|
||||||
|
// Signed data:
|
||||||
|
// "cdn.example.com:/photos/cat.jpg:token=abc&v=2:800:600:webp:1704067200"
|
||||||
|
wantSignature: "394_Vf9TdQFkpQ3XKFDQSyxgqKq8N7mApf2S4QaHqyo=",
|
||||||
|
wantSignedPath: "/v1/image/cdn.example.com/photos/cat.jpg" +
|
||||||
|
"%3Ftoken=abc&v=2/800x600.webp",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "original size without query",
|
||||||
|
req: signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
|
SourceQuery: "",
|
||||||
|
Width: 0,
|
||||||
|
Height: 0,
|
||||||
|
Format: testFormatPNG,
|
||||||
|
},
|
||||||
|
// Signed data: "cdn.example.com:/photos/cat.jpg::0:0:png:1704067200"
|
||||||
|
wantSignature: "7Be7oteeQwvnSPU4bchyQ4ZGYGsAGBKpeEtuQ02ox60=",
|
||||||
|
wantSignedPath: "/v1/image/cdn.example.com/photos/cat.jpg/orig.png",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestSigner_GoldenVectors pins the exact HMAC-SHA256 signature output and
|
// TestSigner_GoldenVectors pins the exact HMAC-SHA256 signature output and
|
||||||
// the exact generated signed URL path for fully-specified requests with a
|
// the exact generated signed URL path for fully-specified requests with a
|
||||||
// hardcoded signing key. The expected values were computed once and are
|
// hardcoded signing key.
|
||||||
// hardcoded here as known answers.
|
|
||||||
//
|
//
|
||||||
// If any of these assertions fail, the signed byte format
|
// If any of these assertions fail, the signed byte format
|
||||||
// ("host:path:query:width:height:format:expiration"), the base64url
|
// ("host:path:query:width:height:format:expiration"), the base64url
|
||||||
@@ -24,67 +90,14 @@ const goldenSigningKey = "golden-test-key"
|
|||||||
// deliberately: update these constants only as part of an intentional,
|
// deliberately: update these constants only as part of an intentional,
|
||||||
// documented signature format migration.
|
// documented signature format migration.
|
||||||
func TestSigner_GoldenVectors(t *testing.T) {
|
func TestSigner_GoldenVectors(t *testing.T) {
|
||||||
signer := New(goldenSigningKey)
|
t.Parallel()
|
||||||
|
|
||||||
vectors := []struct {
|
signer := signature.New(goldenSigningKey)
|
||||||
name string
|
|
||||||
req Request
|
|
||||||
// wantSignature is the exact base64url (RFC 4648 URL-safe,
|
|
||||||
// padded) HMAC-SHA256 signature for the request with Expires
|
|
||||||
// set to goldenExpiresUnix.
|
|
||||||
wantSignature string
|
|
||||||
// wantSignedPath is the exact path returned by
|
|
||||||
// GenerateSignedURL for the request. The signature and
|
|
||||||
// expiration are returned separately by GenerateSignedURL and
|
|
||||||
// are not embedded in the path.
|
|
||||||
wantSignedPath string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "resized without query",
|
|
||||||
req: Request{
|
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
SourceQuery: "",
|
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
},
|
|
||||||
// Signed data: "cdn.example.com:/photos/cat.jpg::800:600:webp:1704067200"
|
|
||||||
wantSignature: "x5PfPp8QSDo0cJT96od-AEgrQyOVLfqifH5sst61_-w=",
|
|
||||||
wantSignedPath: "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "resized with query string",
|
|
||||||
req: Request{
|
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
SourceQuery: "token=abc&v=2",
|
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
},
|
|
||||||
// Signed data: "cdn.example.com:/photos/cat.jpg:token=abc&v=2:800:600:webp:1704067200"
|
|
||||||
wantSignature: "394_Vf9TdQFkpQ3XKFDQSyxgqKq8N7mApf2S4QaHqyo=",
|
|
||||||
wantSignedPath: "/v1/image/cdn.example.com/photos/cat.jpg%3Ftoken=abc&v=2/800x600.webp",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "original size without query",
|
|
||||||
req: Request{
|
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
SourceQuery: "",
|
|
||||||
Width: 0,
|
|
||||||
Height: 0,
|
|
||||||
Format: "png",
|
|
||||||
},
|
|
||||||
// Signed data: "cdn.example.com:/photos/cat.jpg::0:0:png:1704067200"
|
|
||||||
wantSignature: "7Be7oteeQwvnSPU4bchyQ4ZGYGsAGBKpeEtuQ02ox60=",
|
|
||||||
wantSignedPath: "/v1/image/cdn.example.com/photos/cat.jpg/orig.png",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range vectors {
|
for _, tt := range goldenVectors() {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
signReq := tt.req
|
signReq := tt.req
|
||||||
signReq.Expires = time.Unix(goldenExpiresUnix, 0)
|
signReq.Expires = time.Unix(goldenExpiresUnix, 0)
|
||||||
|
|
||||||
@@ -95,9 +108,10 @@ func TestSigner_GoldenVectors(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
urlReq := tt.req
|
urlReq := tt.req
|
||||||
|
|
||||||
gotPath, _, _ := signer.GenerateSignedURL(&urlReq, time.Hour)
|
gotPath, _, _ := signer.GenerateSignedURL(&urlReq, time.Hour)
|
||||||
if gotPath != tt.wantSignedPath {
|
if gotPath != tt.wantSignedPath {
|
||||||
t.Errorf("GenerateSignedURL() path = %q, want %q (signed URL layout changed?)",
|
t.Errorf("GenerateSignedURL() path = %q, want %q (layout changed?)",
|
||||||
gotPath, tt.wantSignedPath)
|
gotPath, tt.wantSignedPath)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -93,31 +93,17 @@ func (s *Signer) Verify(req *Request) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildSignatureData creates the string to be signed.
|
|
||||||
// Format: "host:path:query:width:height:format:expiration"
|
|
||||||
// All components are used verbatim (exact match). No normalization,
|
|
||||||
// suffix matching, or wildcard expansion is performed.
|
|
||||||
func (s *Signer) buildSignatureData(req *Request) string {
|
|
||||||
return fmt.Sprintf("%s:%s:%s:%d:%d:%s:%d",
|
|
||||||
req.SourceHost,
|
|
||||||
req.SourcePath,
|
|
||||||
req.SourceQuery,
|
|
||||||
req.Width,
|
|
||||||
req.Height,
|
|
||||||
req.Format,
|
|
||||||
req.Expires.Unix(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GenerateSignedURL creates a complete URL with signature and expiration.
|
// GenerateSignedURL creates a complete URL with signature and expiration.
|
||||||
// Returns the path portion that should be appended to the base URL.
|
// Returns the path portion that should be appended to the base URL.
|
||||||
func (s *Signer) GenerateSignedURL(req *Request, ttl time.Duration) (path string, sig string, exp int64) {
|
func (s *Signer) GenerateSignedURL(
|
||||||
|
req *Request, ttl time.Duration,
|
||||||
|
) (string, string, int64) {
|
||||||
// Set expiration
|
// Set expiration
|
||||||
req.Expires = time.Now().Add(ttl)
|
req.Expires = time.Now().Add(ttl)
|
||||||
exp = req.Expires.Unix()
|
exp := req.Expires.Unix()
|
||||||
|
|
||||||
// Generate signature
|
// Generate signature
|
||||||
sig = s.Sign(req)
|
sig := s.Sign(req)
|
||||||
req.Signature = sig
|
req.Signature = sig
|
||||||
|
|
||||||
// Build the size component
|
// Build the size component
|
||||||
@@ -134,6 +120,7 @@ func (s *Signer) GenerateSignedURL(req *Request, ttl time.Duration) (path string
|
|||||||
// it from the last-slash split. The "?" inside a path segment is
|
// it from the last-slash split. The "?" inside a path segment is
|
||||||
// percent-encoded by clients but chi delivers it decoded, which is
|
// percent-encoded by clients but chi delivers it decoded, which is
|
||||||
// exactly what the URL parser expects.
|
// exactly what the URL parser expects.
|
||||||
|
var path string
|
||||||
if req.SourceQuery != "" {
|
if req.SourceQuery != "" {
|
||||||
path = fmt.Sprintf("/v1/image/%s%s%%3F%s/%s.%s",
|
path = fmt.Sprintf("/v1/image/%s%s%%3F%s/%s.%s",
|
||||||
req.SourceHost,
|
req.SourceHost,
|
||||||
@@ -154,12 +141,26 @@ func (s *Signer) GenerateSignedURL(req *Request, ttl time.Duration) (path string
|
|||||||
return path, sig, exp
|
return path, sig, exp
|
||||||
}
|
}
|
||||||
|
|
||||||
// ParseParams extracts signature and expiration from query parameters.
|
// buildSignatureData creates the string to be signed.
|
||||||
func ParseParams(sig, expStr string) (parsed string, expires time.Time, err error) {
|
// Format: "host:path:query:width:height:format:expiration"
|
||||||
parsed = sig
|
// All components are used verbatim (exact match). No normalization,
|
||||||
|
// suffix matching, or wildcard expansion is performed.
|
||||||
|
func (s *Signer) buildSignatureData(req *Request) string {
|
||||||
|
return fmt.Sprintf("%s:%s:%s:%d:%d:%s:%d",
|
||||||
|
req.SourceHost,
|
||||||
|
req.SourcePath,
|
||||||
|
req.SourceQuery,
|
||||||
|
req.Width,
|
||||||
|
req.Height,
|
||||||
|
req.Format,
|
||||||
|
req.Expires.Unix(),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseParams extracts signature and expiration from query parameters.
|
||||||
|
func ParseParams(sig, expStr string) (string, time.Time, error) {
|
||||||
if expStr == "" {
|
if expStr == "" {
|
||||||
return parsed, time.Time{}, nil
|
return sig, time.Time{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
expUnix, err := strconv.ParseInt(expStr, 10, 64)
|
expUnix, err := strconv.ParseInt(expStr, 10, 64)
|
||||||
@@ -167,7 +168,5 @@ func ParseParams(sig, expStr string) (parsed string, expires time.Time, err erro
|
|||||||
return "", time.Time{}, fmt.Errorf("invalid expiration: %w", err)
|
return "", time.Time{}, fmt.Errorf("invalid expiration: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
expires = time.Unix(expUnix, 0)
|
return sig, time.Unix(expUnix, 0), nil
|
||||||
|
|
||||||
return parsed, expires, nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,21 +1,36 @@
|
|||||||
package signature
|
package signature_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"sneak.berlin/go/pixa/internal/signature"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Shared fixture values used across the signature tests.
|
||||||
|
const (
|
||||||
|
testHost = "cdn.example.com"
|
||||||
|
testPath = "/photos/cat.jpg"
|
||||||
|
testFormatWebP = "webp"
|
||||||
|
testFormatPNG = "png"
|
||||||
|
testSignedPath = "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp"
|
||||||
|
testSig = "abc123"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestSigner_Sign(t *testing.T) {
|
func TestSigner_Sign(t *testing.T) {
|
||||||
signer := New("test-secret-key")
|
t.Parallel()
|
||||||
|
|
||||||
req := &Request{
|
signer := signature.New("test-secret-key")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
Expires: time.Unix(1704067200, 0), // Fixed timestamp for reproducibility
|
Expires: time.Unix(1704067200, 0), // Fixed timestamp for reproducibility
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -24,7 +39,8 @@ func TestSigner_Sign(t *testing.T) {
|
|||||||
|
|
||||||
// Same input should produce same signature
|
// Same input should produce same signature
|
||||||
if sig1 != sig2 {
|
if sig1 != sig2 {
|
||||||
t.Errorf("Sign() produced different signatures for same input: %q vs %q", sig1, sig2)
|
t.Errorf("Sign() produced different signatures for same input: %q vs %q",
|
||||||
|
sig1, sig2)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Signature should be non-empty
|
// Signature should be non-empty
|
||||||
@@ -33,13 +49,13 @@ func TestSigner_Sign(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Different input should produce different signature
|
// Different input should produce different signature
|
||||||
req2 := &Request{
|
req2 := &signature.Request{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHost,
|
||||||
SourcePath: "/photos/dog.jpg", // Different path
|
SourcePath: "/photos/dog.jpg", // Different path
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
Expires: time.Unix(1704067200, 0),
|
Expires: time.Unix(1704067200, 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -49,25 +65,31 @@ func TestSigner_Sign(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSigner_Verify(t *testing.T) {
|
// validVerifyRequest returns a fully-populated request that verifies
|
||||||
signer := New("test-secret-key")
|
// successfully once signed.
|
||||||
|
func validVerifyRequest() *signature.Request {
|
||||||
|
return &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
|
Width: 800,
|
||||||
|
Height: 600,
|
||||||
|
Format: testFormatWebP,
|
||||||
|
Expires: time.Now().Add(1 * time.Hour),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
tests := []struct {
|
type verifyCase struct {
|
||||||
name string
|
name string
|
||||||
setup func() *Request
|
setup func() *signature.Request
|
||||||
wantErr error
|
wantErr error
|
||||||
}{
|
}
|
||||||
|
|
||||||
|
func verifyCases(signer *signature.Signer) []verifyCase {
|
||||||
|
return []verifyCase{
|
||||||
{
|
{
|
||||||
name: "valid signature",
|
name: "valid signature",
|
||||||
setup: func() *Request {
|
setup: func() *signature.Request {
|
||||||
req := &Request{
|
req := validVerifyRequest()
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
Expires: time.Now().Add(1 * time.Hour),
|
|
||||||
}
|
|
||||||
req.Signature = signer.Sign(req)
|
req.Signature = signer.Sign(req)
|
||||||
|
|
||||||
return req
|
return req
|
||||||
@@ -76,74 +98,59 @@ func TestSigner_Verify(t *testing.T) {
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "expired signature",
|
name: "expired signature",
|
||||||
setup: func() *Request {
|
setup: func() *signature.Request {
|
||||||
req := &Request{
|
req := validVerifyRequest()
|
||||||
SourceHost: "cdn.example.com",
|
req.Expires = time.Now().Add(-1 * time.Hour)
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
Expires: time.Now().Add(-1 * time.Hour), // Expired
|
|
||||||
}
|
|
||||||
req.Signature = signer.Sign(req)
|
req.Signature = signer.Sign(req)
|
||||||
|
|
||||||
return req
|
return req
|
||||||
},
|
},
|
||||||
wantErr: ErrExpired,
|
wantErr: signature.ErrExpired,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "invalid signature",
|
name: "invalid signature",
|
||||||
setup: func() *Request {
|
setup: func() *signature.Request {
|
||||||
return &Request{
|
req := validVerifyRequest()
|
||||||
SourceHost: "cdn.example.com",
|
req.Signature = "invalid-signature"
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
Width: 800,
|
return req
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
Expires: time.Now().Add(1 * time.Hour),
|
|
||||||
Signature: "invalid-signature",
|
|
||||||
}
|
|
||||||
},
|
},
|
||||||
wantErr: ErrInvalid,
|
wantErr: signature.ErrInvalid,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "missing expiration",
|
name: "missing expiration",
|
||||||
setup: func() *Request {
|
setup: func() *signature.Request {
|
||||||
return &Request{
|
req := validVerifyRequest()
|
||||||
SourceHost: "cdn.example.com",
|
req.Expires = time.Time{}
|
||||||
SourcePath: "/photos/cat.jpg",
|
req.Signature = "some-signature"
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
return req
|
||||||
Format: "webp",
|
|
||||||
Signature: "some-signature",
|
|
||||||
// Expires is zero
|
|
||||||
}
|
|
||||||
},
|
},
|
||||||
wantErr: ErrMissingExpiration,
|
wantErr: signature.ErrMissingExpiration,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "tampered request",
|
name: "tampered request",
|
||||||
setup: func() *Request {
|
setup: func() *signature.Request {
|
||||||
req := &Request{
|
req := validVerifyRequest()
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
Width: 800,
|
|
||||||
Height: 600,
|
|
||||||
Format: "webp",
|
|
||||||
Expires: time.Now().Add(1 * time.Hour),
|
|
||||||
}
|
|
||||||
req.Signature = signer.Sign(req)
|
req.Signature = signer.Sign(req)
|
||||||
// Tamper with the request
|
|
||||||
req.SourcePath = "/photos/secret.jpg"
|
req.SourcePath = "/photos/secret.jpg"
|
||||||
|
|
||||||
return req
|
return req
|
||||||
},
|
},
|
||||||
wantErr: ErrInvalid,
|
wantErr: signature.ErrInvalid,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
for _, tt := range tests {
|
func TestSigner_Verify(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
signer := signature.New("test-secret-key")
|
||||||
|
|
||||||
|
for _, tt := range verifyCases(signer) {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := tt.setup()
|
req := tt.setup()
|
||||||
err := signer.Verify(req)
|
err := signer.Verify(req)
|
||||||
|
|
||||||
@@ -151,30 +158,102 @@ func TestSigner_Verify(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Errorf("Verify() unexpected error = %v", err)
|
t.Errorf("Verify() unexpected error = %v", err)
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
if err != tt.wantErr {
|
return
|
||||||
t.Errorf("Verify() error = %v, wantErr %v", err, tt.wantErr)
|
}
|
||||||
}
|
|
||||||
|
if !errors.Is(err, tt.wantErr) {
|
||||||
|
t.Errorf("Verify() error = %v, wantErr %v", err, tt.wantErr)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type tamperCase struct {
|
||||||
|
name string
|
||||||
|
tamper func(r *signature.Request)
|
||||||
|
}
|
||||||
|
|
||||||
|
// exactMatchTamperCases mutates one signed component per case; every
|
||||||
|
// mutation must cause verification to fail with ErrInvalid.
|
||||||
|
func exactMatchTamperCases() []tamperCase {
|
||||||
|
return []tamperCase{
|
||||||
|
{
|
||||||
|
name: "parent domain does not match subdomain",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceHost = "example.com" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "subdomain does not match parent domain",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceHost = "images.cdn.example.com" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sibling subdomain does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceHost = "images.example.com" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "host with suffix appended does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceHost = testHost + ".evil.com" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "host with prefix does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceHost = "evilcdn.example.com" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different path does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourcePath = "/photos/dog.jpg" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "path suffix does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourcePath = testPath + "/extra" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "path prefix does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourcePath = "/other" + testPath },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different query does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceQuery = "token=xyz" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "added query does not match empty query",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceQuery = "extra=1" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "removed query does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.SourceQuery = "" },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different width does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.Width = 801 },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different height does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.Height = 601 },
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "different format does not match",
|
||||||
|
tamper: func(r *signature.Request) { r.Format = testFormatPNG },
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestSigner_Verify_ExactMatchOnly verifies that signatures enforce exact
|
// TestSigner_Verify_ExactMatchOnly verifies that signatures enforce exact
|
||||||
// matching on every URL component. No suffix matching, wildcard matching,
|
// matching on every URL component. No suffix matching, wildcard matching,
|
||||||
// or partial matching is supported.
|
// or partial matching is supported.
|
||||||
func TestSigner_Verify_ExactMatchOnly(t *testing.T) {
|
func TestSigner_Verify_ExactMatchOnly(t *testing.T) {
|
||||||
signer := New("test-secret-key")
|
t.Parallel()
|
||||||
|
|
||||||
|
signer := signature.New("test-secret-key")
|
||||||
|
|
||||||
// Base request that we'll sign, then tamper with individual fields.
|
// Base request that we'll sign, then tamper with individual fields.
|
||||||
baseReq := func() *Request {
|
baseReq := func() *signature.Request {
|
||||||
req := &Request{
|
req := &signature.Request{
|
||||||
SourceHost: "cdn.example.com",
|
SourceHost: testHost,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPath,
|
||||||
SourceQuery: "token=abc",
|
SourceQuery: "token=abc",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
Expires: time.Now().Add(1 * time.Hour),
|
Expires: time.Now().Add(1 * time.Hour),
|
||||||
}
|
}
|
||||||
req.Signature = signer.Sign(req)
|
req.Signature = signer.Sign(req)
|
||||||
@@ -182,117 +261,28 @@ func TestSigner_Verify_ExactMatchOnly(t *testing.T) {
|
|||||||
return req
|
return req
|
||||||
}
|
}
|
||||||
|
|
||||||
tests := []struct {
|
for _, tt := range exactMatchTamperCases() {
|
||||||
name string
|
|
||||||
tamper func(req *Request)
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
name: "parent domain does not match subdomain",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
// Signed for cdn.example.com, try example.com
|
|
||||||
req.SourceHost = "example.com"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "subdomain does not match parent domain",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
// Signed for cdn.example.com, try images.cdn.example.com
|
|
||||||
req.SourceHost = "images.cdn.example.com"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "sibling subdomain does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
// Signed for cdn.example.com, try images.example.com
|
|
||||||
req.SourceHost = "images.example.com"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "host with suffix appended does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
// Signed for cdn.example.com, try cdn.example.com.evil.com
|
|
||||||
req.SourceHost = "cdn.example.com.evil.com"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "host with prefix does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
// Signed for cdn.example.com, try evilcdn.example.com
|
|
||||||
req.SourceHost = "evilcdn.example.com"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different path does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourcePath = "/photos/dog.jpg"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "path suffix does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourcePath = "/photos/cat.jpg/extra"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "path prefix does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourcePath = "/other/photos/cat.jpg"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different query does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourceQuery = "token=xyz"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "added query does not match empty query",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourceQuery = "extra=1"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "removed query does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.SourceQuery = ""
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different width does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.Width = 801
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different height does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.Height = 601
|
|
||||||
},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
name: "different format does not match",
|
|
||||||
tamper: func(req *Request) {
|
|
||||||
req.Format = "png"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := baseReq()
|
req := baseReq()
|
||||||
tt.tamper(req)
|
tt.tamper(req)
|
||||||
|
|
||||||
err := signer.Verify(req)
|
err := signer.Verify(req)
|
||||||
if err != ErrInvalid {
|
if !errors.Is(err, signature.ErrInvalid) {
|
||||||
t.Errorf("Verify() = %v, want %v", err, ErrInvalid)
|
t.Errorf("Verify() = %v, want %v", err, signature.ErrInvalid)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify the unmodified base request still passes
|
// Verify the unmodified base request still passes
|
||||||
t.Run("unmodified request passes", func(t *testing.T) {
|
t.Run("unmodified request passes", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
req := baseReq()
|
req := baseReq()
|
||||||
if err := signer.Verify(req); err != nil {
|
|
||||||
|
err := signer.Verify(req)
|
||||||
|
if err != nil {
|
||||||
t.Errorf("Verify() unmodified request failed: %v", err)
|
t.Errorf("Verify() unmodified request failed: %v", err)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -302,10 +292,12 @@ func TestSigner_Verify_ExactMatchOnly(t *testing.T) {
|
|||||||
// string in the signature data, producing different signatures for
|
// string in the signature data, producing different signatures for
|
||||||
// suffix-related hosts.
|
// suffix-related hosts.
|
||||||
func TestSigner_Sign_ExactHostInData(t *testing.T) {
|
func TestSigner_Sign_ExactHostInData(t *testing.T) {
|
||||||
signer := New("test-secret-key")
|
t.Parallel()
|
||||||
|
|
||||||
|
signer := signature.New("test-secret-key")
|
||||||
|
|
||||||
hosts := []string{
|
hosts := []string{
|
||||||
"cdn.example.com",
|
testHost,
|
||||||
"example.com",
|
"example.com",
|
||||||
"images.example.com",
|
"images.example.com",
|
||||||
"images.cdn.example.com",
|
"images.cdn.example.com",
|
||||||
@@ -315,13 +307,13 @@ func TestSigner_Sign_ExactHostInData(t *testing.T) {
|
|||||||
sigs := make(map[string]string)
|
sigs := make(map[string]string)
|
||||||
|
|
||||||
for _, host := range hosts {
|
for _, host := range hosts {
|
||||||
req := &Request{
|
req := &signature.Request{
|
||||||
SourceHost: host,
|
SourceHost: host,
|
||||||
SourcePath: "/photos/cat.jpg",
|
SourcePath: testPath,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
Expires: time.Unix(1704067200, 0),
|
Expires: time.Unix(1704067200, 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -335,15 +327,17 @@ func TestSigner_Sign_ExactHostInData(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestSigner_DifferentKeys(t *testing.T) {
|
func TestSigner_DifferentKeys(t *testing.T) {
|
||||||
signer1 := New("secret-key-1")
|
t.Parallel()
|
||||||
signer2 := New("secret-key-2")
|
|
||||||
|
|
||||||
req := &Request{
|
signer1 := signature.New("secret-key-1")
|
||||||
SourceHost: "cdn.example.com",
|
signer2 := signature.New("secret-key-2")
|
||||||
SourcePath: "/photos/cat.jpg",
|
|
||||||
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
Expires: time.Now().Add(1 * time.Hour),
|
Expires: time.Now().Add(1 * time.Hour),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -351,35 +345,38 @@ func TestSigner_DifferentKeys(t *testing.T) {
|
|||||||
req.Signature = signer1.Sign(req)
|
req.Signature = signer1.Sign(req)
|
||||||
|
|
||||||
// Verify with key 1 should succeed
|
// Verify with key 1 should succeed
|
||||||
if err := signer1.Verify(req); err != nil {
|
err := signer1.Verify(req)
|
||||||
|
if err != nil {
|
||||||
t.Errorf("Verify() with same key failed: %v", err)
|
t.Errorf("Verify() with same key failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify with key 2 should fail
|
// Verify with key 2 should fail
|
||||||
if err := signer2.Verify(req); err != ErrInvalid {
|
err = signer2.Verify(req)
|
||||||
|
if !errors.Is(err, signature.ErrInvalid) {
|
||||||
t.Errorf("Verify() with different key should fail, got: %v", err)
|
t.Errorf("Verify() with different key should fail, got: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerateSignedURL(t *testing.T) {
|
func TestGenerateSignedURL(t *testing.T) {
|
||||||
signer := New("test-secret-key")
|
t.Parallel()
|
||||||
|
|
||||||
req := &Request{
|
signer := signature.New("test-secret-key")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
SourceQuery: "",
|
SourceQuery: "",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
}
|
}
|
||||||
|
|
||||||
ttl := 1 * time.Hour
|
ttl := 1 * time.Hour
|
||||||
path, sig, exp := signer.GenerateSignedURL(req, ttl)
|
path, sig, exp := signer.GenerateSignedURL(req, ttl)
|
||||||
|
|
||||||
// Path should be correct format
|
// Path should be correct format
|
||||||
expectedPath := "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp"
|
if path != testSignedPath {
|
||||||
if path != expectedPath {
|
t.Errorf("GenerateSignedURL() path = %q, want %q", path, testSignedPath)
|
||||||
t.Errorf("GenerateSignedURL() path = %q, want %q", path, expectedPath)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Signature should be non-empty
|
// Signature should be non-empty
|
||||||
@@ -389,6 +386,7 @@ func TestGenerateSignedURL(t *testing.T) {
|
|||||||
|
|
||||||
// Expiration should be approximately now + TTL
|
// Expiration should be approximately now + TTL
|
||||||
expTime := time.Unix(exp, 0)
|
expTime := time.Unix(exp, 0)
|
||||||
|
|
||||||
expectedExp := time.Now().Add(ttl)
|
expectedExp := time.Now().Add(ttl)
|
||||||
if expTime.Sub(expectedExp) > time.Second {
|
if expTime.Sub(expectedExp) > time.Second {
|
||||||
t.Errorf("GenerateSignedURL() exp time off by too much")
|
t.Errorf("GenerateSignedURL() exp time off by too much")
|
||||||
@@ -401,14 +399,16 @@ func TestGenerateSignedURL(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerateSignedURL_OrigSize(t *testing.T) {
|
func TestGenerateSignedURL_OrigSize(t *testing.T) {
|
||||||
signer := New("test-secret-key")
|
t.Parallel()
|
||||||
|
|
||||||
req := &Request{
|
signer := signature.New("test-secret-key")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
Width: 0, // Original size
|
Width: 0, // Original size
|
||||||
Height: 0,
|
Height: 0,
|
||||||
Format: "png",
|
Format: testFormatPNG,
|
||||||
}
|
}
|
||||||
|
|
||||||
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
||||||
@@ -420,21 +420,24 @@ func TestGenerateSignedURL_OrigSize(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerateSignedURL_WithQueryString(t *testing.T) {
|
func TestGenerateSignedURL_WithQueryString(t *testing.T) {
|
||||||
signer := New("test-secret-key-for-testing!")
|
t.Parallel()
|
||||||
|
|
||||||
req := &Request{
|
signer := signature.New("test-secret-key-for-testing!")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
SourceQuery: "token=abc&v=2",
|
SourceQuery: "token=abc&v=2",
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
}
|
}
|
||||||
|
|
||||||
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
||||||
|
|
||||||
// The path must NOT contain a bare "?" that would be interpreted as a query string delimiter.
|
// The path must NOT contain a bare "?" that would be interpreted as
|
||||||
// The size segment must appear as the last path component.
|
// a query string delimiter. The size segment must appear as the last
|
||||||
|
// path component.
|
||||||
if strings.Contains(path, "?token=abc") {
|
if strings.Contains(path, "?token=abc") {
|
||||||
t.Errorf("GenerateSignedURL() produced bare query string in path: %q", path)
|
t.Errorf("GenerateSignedURL() produced bare query string in path: %q", path)
|
||||||
}
|
}
|
||||||
@@ -451,25 +454,28 @@ func TestGenerateSignedURL_WithQueryString(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerateSignedURL_WithoutQueryString(t *testing.T) {
|
func TestGenerateSignedURL_WithoutQueryString(t *testing.T) {
|
||||||
signer := New("test-secret-key-for-testing!")
|
t.Parallel()
|
||||||
|
|
||||||
req := &Request{
|
signer := signature.New("test-secret-key-for-testing!")
|
||||||
SourceHost: "cdn.example.com",
|
|
||||||
SourcePath: "/photos/cat.jpg",
|
req := &signature.Request{
|
||||||
|
SourceHost: testHost,
|
||||||
|
SourcePath: testPath,
|
||||||
Width: 800,
|
Width: 800,
|
||||||
Height: 600,
|
Height: 600,
|
||||||
Format: "webp",
|
Format: testFormatWebP,
|
||||||
}
|
}
|
||||||
|
|
||||||
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
path, _, _ := signer.GenerateSignedURL(req, time.Hour)
|
||||||
|
|
||||||
expected := "/v1/image/cdn.example.com/photos/cat.jpg/800x600.webp"
|
if path != testSignedPath {
|
||||||
if path != expected {
|
t.Errorf("GenerateSignedURL() path = %q, want %q", path, testSignedPath)
|
||||||
t.Errorf("GenerateSignedURL() path = %q, want %q", path, expected)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestParseParams(t *testing.T) {
|
func TestParseParams(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
sig string
|
sig string
|
||||||
@@ -480,21 +486,21 @@ func TestParseParams(t *testing.T) {
|
|||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "valid params",
|
name: "valid params",
|
||||||
sig: "abc123",
|
sig: testSig,
|
||||||
expStr: "1704067200",
|
expStr: "1704067200",
|
||||||
wantSig: "abc123",
|
wantSig: testSig,
|
||||||
wantErr: false,
|
wantErr: false,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "empty expiration",
|
name: "empty expiration",
|
||||||
sig: "abc123",
|
sig: testSig,
|
||||||
expStr: "",
|
expStr: "",
|
||||||
wantSig: "abc123",
|
wantSig: testSig,
|
||||||
wantErr: false,
|
wantErr: false,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "invalid expiration",
|
name: "invalid expiration",
|
||||||
sig: "abc123",
|
sig: testSig,
|
||||||
expStr: "not-a-number",
|
expStr: "not-a-number",
|
||||||
wantErr: true,
|
wantErr: true,
|
||||||
},
|
},
|
||||||
@@ -502,7 +508,9 @@ func TestParseParams(t *testing.T) {
|
|||||||
|
|
||||||
for _, tt := range tests {
|
for _, tt := range tests {
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
sig, exp, err := ParseParams(tt.sig, tt.expStr)
|
t.Parallel()
|
||||||
|
|
||||||
|
sig, exp, err := signature.ParseParams(tt.sig, tt.expStr)
|
||||||
|
|
||||||
if tt.wantErr {
|
if tt.wantErr {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
|||||||
@@ -11,11 +11,11 @@ set -eu
|
|||||||
|
|
||||||
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
|
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
|
||||||
|
|
||||||
# Pinned versions, 2026-07-07. Never "latest"; exact versions only.
|
# Pinned versions, 2026-08-07. Never "latest"; exact versions only.
|
||||||
GOLANGCI_LINT_VERSION="2.10.1"
|
GOLANGCI_LINT_VERSION="2.12.2"
|
||||||
# sha256 of golangci-lint-2.10.1-linux-<arch>.tar.gz release archives
|
# sha256 of golangci-lint-2.12.2-linux-<arch>.tar.gz release archives
|
||||||
GOLANGCI_LINT_SHA256_AMD64="dfa775874cf0561b404a02a8f4481fc69b28091da95aa697259820d429b09c99"
|
GOLANGCI_LINT_SHA256_AMD64="8df580d2670fed8fa984aac0507099af8df275e665215f5c7a2ae3943893a553"
|
||||||
GOLANGCI_LINT_SHA256_ARM64="6652b42ae02915eb2f9cb2a2e0cac99514c8eded8388d88ae3e06e1a52c00de8"
|
GOLANGCI_LINT_SHA256_ARM64="44cd40a8c76c86755375adfeea52cfd3533cb43d7bd647771e0ae065e166df3a"
|
||||||
|
|
||||||
PKGMGR=""
|
PKGMGR=""
|
||||||
SUDO=""
|
SUDO=""
|
||||||
|
|||||||
Reference in New Issue
Block a user