fix(canary): clear pre-audit lint debt + migrate golangci config to v2

Pre-audit gate (`golangci-lint run`) at the start of Phase 2 surfaced 15
issues that should have been zeroed before Phase 1 rollup. Per the
fix-in-phase rule (no backlog rot), clearing everything before Phase 2
audit agents run on a green tree.

Config:
- .golangci.yml: migrate `issues.exclude-rules` and `issues.exclude-dirs`
  to v2 syntax (`linters.exclusions.{rules,paths}`); test-file funlen/
  dupl/goconst exclusion now actually applies under golangci-lint v2.10
- .golangci.yml: add G706 to gosec excludes — false-positive log-injection
  reports on slog structured-logging call sites (slog separates message
  from kv args, immune to log-line injection by construction)

errcheck (5):
- cmd/canary/main.go: `_ = telemetry.Shutdown(...)` → log on error
- internal/token/repository.go: `defer stmt.Close()` → log on close error
- internal/event/repository.go: same as token
- internal/testutil/postgres.go: `_ = pgContainer.Terminate(...)` and
  `_ = db.Close()` → t.Logf on cleanup error (use distinct err names to
  avoid govet shadow)

funlen (1):
- cmd/canary/main.go: split `run` into `run` + `initTelemetry` +
  `mountRouter` + `gracefulShutdown` (was 55 statements, now under
  the 50 cap; helpers are individually well under)

govet shadow (1):
- cmd/canary/main.go: migrations check switched from `if err :=` to
  `if err =` (reuses outer err whose value was already consumed) — the
  inner short-decl was lexically shadowing without intent
- cmd/canary/main.go: select-case `err` renamed to `startErr` to avoid
  the same shadow pattern

golines (6, all auto-fixed by `golangci-lint run --fix`):
- main.go, config.go, telemetry.go, event/repository.go, token/dto.go,
  token/repository.go — long lines wrapped to 80-col

Verified post-fix: build OK, vet OK, unit tests PASS under -race,
integration tests PASS under -tags=integration -race (12s testcontainers),
golangci-lint reports 0 issues.
This commit is contained in:
CarterPerez-dev 2026-05-12 02:49:56 -04:00
parent bb95f74579
commit c39b56b9af
13 changed files with 351 additions and 167 deletions

View File

@ -72,23 +72,26 @@ linters:
gosec: gosec:
excludes: excludes:
- G104 - G104
- G706
sloglint: sloglint:
no-mixed-args: true no-mixed-args: true
kv-only: true kv-only: true
context: all context: all
exclusions:
paths:
- vendor
- testdata
rules:
- path: _test\.go
linters:
- funlen
- dupl
- goconst
issues: issues:
max-same-issues: 50 max-same-issues: 50
exclude-dirs:
- vendor
- testdata
exclude-rules:
- path: _test\.go
linters:
- funlen
- dupl
- goconst
formatters: formatters:

View File

@ -24,7 +24,10 @@ import (
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token" "github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token"
) )
const drainDelay = 5 * time.Second const (
drainDelay = 5 * time.Second
shutdownGraceExtra = 5 * time.Second
)
func main() { func main() {
configPath := flag.String("config", "config.yaml", "path to config file") configPath := flag.String("config", "config.yaml", "path to config file")
@ -37,7 +40,11 @@ func main() {
} }
func run(configPath string) error { func run(configPath string) error {
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(
context.Background(),
syscall.SIGINT,
syscall.SIGTERM,
)
defer stop() defer stop()
cfg, err := config.Load(configPath) cfg, err := config.Load(configPath)
@ -52,14 +59,7 @@ func run(configPath string) error {
"environment", cfg.App.Environment, "environment", cfg.App.Environment,
) )
var telemetry *core.Telemetry telemetry := initTelemetry(ctx, cfg, logger)
if cfg.Otel.Enabled {
if t, telErr := core.NewTelemetry(ctx, cfg.Otel, cfg.App); telErr != nil {
logger.Warn("telemetry init failed", "error", telErr)
} else {
telemetry = t
}
}
db, err := core.NewDatabase(ctx, cfg.Database) db, err := core.NewDatabase(ctx, cfg.Database)
if err != nil { if err != nil {
@ -67,15 +67,13 @@ func run(configPath string) error {
} }
logger.Info("database connected") logger.Info("database connected")
if err := core.RunMigrations(db.SQLDB()); err != nil { if err = core.RunMigrations(db.SQLDB()); err != nil {
return fmt.Errorf("run migrations: %w", err) return fmt.Errorf("run migrations: %w", err)
} }
logger.Info("migrations applied") logger.Info("migrations applied")
tokenRepo := token.NewRepository(db.DB) _ = token.NewRepository(db.DB)
eventRepo := event.NewRepository(db.DB) _ = event.NewRepository(db.DB)
_ = tokenRepo
_ = eventRepo
rdb, err := core.NewRedis(ctx, cfg.Redis) rdb, err := core.NewRedis(ctx, cfg.Redis)
if err != nil { if err != nil {
@ -84,7 +82,43 @@ func run(configPath string) error {
logger.Info("redis connected") logger.Info("redis connected")
healthH := health.NewHandler(db, rdb) healthH := health.NewHandler(db, rdb)
srv := mountRouter(cfg, logger, rdb, healthH)
errChan := make(chan error, 1)
go func() { errChan <- srv.Start() }()
select {
case startErr := <-errChan:
return startErr
case <-ctx.Done():
logger.Info("shutdown signal received")
}
return gracefulShutdown(cfg, logger, srv, telemetry, rdb, db)
}
func initTelemetry(
ctx context.Context,
cfg *config.Config,
logger *slog.Logger,
) *core.Telemetry {
if !cfg.Otel.Enabled {
return nil
}
t, err := core.NewTelemetry(ctx, cfg.Otel, cfg.App)
if err != nil {
logger.Warn("telemetry init failed", "error", err)
return nil
}
return t
}
func mountRouter(
cfg *config.Config,
logger *slog.Logger,
rdb *core.Redis,
healthH *health.Handler,
) *server.Server {
srv := server.New(server.Config{ srv := server.New(server.Config{
ServerConfig: cfg.Server, ServerConfig: cfg.Server,
HealthHandler: healthH, HealthHandler: healthH,
@ -96,7 +130,10 @@ func run(configPath string) error {
r.Use(middleware.Logger(logger)) r.Use(middleware.Logger(logger))
r.Use( r.Use(
middleware.NewRateLimiter(rdb.Client, middleware.RateLimitConfig{ middleware.NewRateLimiter(rdb.Client, middleware.RateLimitConfig{
Limit: middleware.PerMinute(cfg.RateLimit.Requests, cfg.RateLimit.Burst), Limit: middleware.PerMinute(
cfg.RateLimit.Requests,
cfg.RateLimit.Burst,
),
FailOpen: true, FailOpen: true,
}).Handler, }).Handler,
) )
@ -104,29 +141,31 @@ func run(configPath string) error {
r.Use(middleware.CORS(cfg.CORS)) r.Use(middleware.CORS(cfg.CORS))
healthH.RegisterRoutes(r) healthH.RegisterRoutes(r)
r.Route("/api", func(_ chi.Router) {})
return srv
}
r.Route("/api", func(_ chi.Router) { func gracefulShutdown(
}) cfg *config.Config,
logger *slog.Logger,
errChan := make(chan error, 1) srv *server.Server,
go func() { errChan <- srv.Start() }() telemetry *core.Telemetry,
rdb *core.Redis,
select { db *core.Database,
case err := <-errChan: ) error {
return err shutdownCtx, cancel := context.WithTimeout(
case <-ctx.Done(): context.Background(),
logger.Info("shutdown signal received") cfg.Server.ShutdownTimeout+drainDelay+shutdownGraceExtra,
} )
shutdownCtx, cancel := context.WithTimeout(context.Background(),
cfg.Server.ShutdownTimeout+drainDelay+5*time.Second)
defer cancel() defer cancel()
if err := srv.Shutdown(shutdownCtx, drainDelay); err != nil { if err := srv.Shutdown(shutdownCtx, drainDelay); err != nil {
logger.Error("server shutdown error", "error", err) logger.Error("server shutdown error", "error", err)
} }
if telemetry != nil { if telemetry != nil {
_ = telemetry.Shutdown(shutdownCtx) if err := telemetry.Shutdown(shutdownCtx); err != nil {
logger.Error("telemetry shutdown error", "error", err)
}
} }
if err := rdb.Close(); err != nil { if err := rdb.Close(); err != nil {
logger.Error("redis close error", "error", err) logger.Error("redis close error", "error", err)

View File

@ -98,13 +98,19 @@ func Load(configPath string) (*Config, error) {
} }
if configPath != "" { if configPath != "" {
if err := k.Load(file.Provider(configPath), yaml.Parser()); err != nil { if err := k.Load(
file.Provider(configPath),
yaml.Parser(),
); err != nil {
loadErr = fmt.Errorf("load config file: %w", err) loadErr = fmt.Errorf("load config file: %w", err)
return return
} }
} }
if err := k.Load(env.Provider("", ".", envKeyReplacer), nil); err != nil { if err := k.Load(
env.Provider("", ".", envKeyReplacer),
nil,
); err != nil {
loadErr = fmt.Errorf("load env vars: %w", err) loadErr = fmt.Errorf("load env vars: %w", err)
return return
} }

View File

@ -51,7 +51,12 @@ func NewTelemetry(
otlptracegrpc.WithTLSCredentials(insecure.NewCredentials()), otlptracegrpc.WithTLSCredentials(insecure.NewCredentials()),
) )
} else { } else {
opts = append(opts, otlptracegrpc.WithTLSCredentials(credentials.NewClientTLSFromCert(nil, ""))) opts = append(
opts,
otlptracegrpc.WithTLSCredentials(
credentials.NewClientTLSFromCert(nil, ""),
),
)
} }
exporter, err := otlptracegrpc.New(ctx, opts...) exporter, err := otlptracegrpc.New(ctx, opts...)

View File

@ -9,6 +9,7 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"log/slog"
"time" "time"
"github.com/jmoiron/sqlx" "github.com/jmoiron/sqlx"
@ -50,7 +51,12 @@ func (r *Repository) Insert(ctx context.Context, e *Event) error {
if err != nil { if err != nil {
return fmt.Errorf("prepare insert event: %w", err) return fmt.Errorf("prepare insert event: %w", err)
} }
defer stmt.Close() defer func() {
if cerr := stmt.Close(); cerr != nil {
slog.WarnContext(ctx, "close prepared stmt",
"op", "insert_event", "error", cerr)
}
}()
if err := stmt.GetContext(ctx, e, e); err != nil { if err := stmt.GetContext(ctx, e, e); err != nil {
return fmt.Errorf("insert event: %w", err) return fmt.Errorf("insert event: %w", err)
@ -90,7 +96,14 @@ func (r *Repository) ListByToken(
LIMIT $3` LIMIT $3`
var events []Event var events []Event
err := r.db.SelectContext(ctx, &events, q, tokenID, opts.Cursor, opts.Limit+1) err := r.db.SelectContext(
ctx,
&events,
q,
tokenID,
opts.Cursor,
opts.Limit+1,
)
if err != nil { if err != nil {
return ListResult{}, fmt.Errorf("list events: %w", err) return ListResult{}, fmt.Errorf("list events: %w", err)
} }
@ -112,7 +125,10 @@ func (r *Repository) ListByToken(
}, nil }, nil
} }
func (r *Repository) CountByToken(ctx context.Context, tokenID string) (int64, error) { func (r *Repository) CountByToken(
ctx context.Context,
tokenID string,
) (int64, error) {
var n int64 var n int64
err := r.db.GetContext(ctx, &n, err := r.db.GetContext(ctx, &n,
`SELECT COUNT(*) FROM events WHERE token_id = $1`, tokenID) `SELECT COUNT(*) FROM events WHERE token_id = $1`, tokenID)
@ -123,7 +139,10 @@ func (r *Repository) CountByToken(ctx context.Context, tokenID string) (int64, e
} }
func (r *Repository) AttachFingerprint( func (r *Repository) AttachFingerprint(
ctx context.Context, tokenID, sourceIP string, fingerprint json.RawMessage, window time.Duration, ctx context.Context,
tokenID, sourceIP string,
fingerprint json.RawMessage,
window time.Duration,
) error { ) error {
q := ` q := `
UPDATE events UPDATE events
@ -136,8 +155,14 @@ UPDATE events
ORDER BY id DESC ORDER BY id DESC
LIMIT 1 LIMIT 1
)` )`
res, err := r.db.ExecContext(ctx, q, res, err := r.db.ExecContext(
tokenID, sourceIP, []byte(fingerprint), fmt.Sprintf("%d milliseconds", window.Milliseconds())) ctx,
q,
tokenID,
sourceIP,
[]byte(fingerprint),
fmt.Sprintf("%d milliseconds", window.Milliseconds()),
)
if err != nil { if err != nil {
return fmt.Errorf("attach fingerprint: %w", err) return fmt.Errorf("attach fingerprint: %w", err)
} }
@ -172,7 +197,10 @@ UPDATE events
return nil return nil
} }
func (r *Repository) PruneToLimit(ctx context.Context, perTokenLimit int) (int64, error) { func (r *Repository) PruneToLimit(
ctx context.Context,
perTokenLimit int,
) (int64, error) {
if perTokenLimit <= 0 { if perTokenLimit <= 0 {
return 0, errors.New("perTokenLimit must be positive") return 0, errors.New("perTokenLimit must be positive")
} }

View File

@ -36,7 +36,11 @@ func NewTestDB(t *testing.T) *sql.DB {
require.NoError(t, err) require.NoError(t, err)
t.Cleanup(func() { t.Cleanup(func() {
_ = pgContainer.Terminate(context.Background()) if termErr := pgContainer.Terminate(
context.Background(),
); termErr != nil {
t.Logf("postgres container terminate: %v", termErr)
}
}) })
connStr, err := pgContainer.ConnectionString(ctx, "sslmode=disable") connStr, err := pgContainer.ConnectionString(ctx, "sslmode=disable")
@ -46,7 +50,9 @@ func NewTestDB(t *testing.T) *sql.DB {
require.NoError(t, err) require.NoError(t, err)
t.Cleanup(func() { t.Cleanup(func() {
_ = db.Close() if closeErr := db.Close(); closeErr != nil {
t.Logf("db close: %v", closeErr)
}
}) })
require.NoError(t, db.Ping()) require.NoError(t, db.Ping())

View File

@ -9,13 +9,13 @@ import (
) )
type CreateRequest struct { type CreateRequest struct {
Type Type `json:"type" validate:"required,oneof=webbug slowredirect docx pdf kubeconfig envfile mysql"` Type Type `json:"type" validate:"required,oneof=webbug slowredirect docx pdf kubeconfig envfile mysql"`
Memo string `json:"memo" validate:"max=256"` Memo string `json:"memo" validate:"max=256"`
Filename string `json:"filename" validate:"max=128"` Filename string `json:"filename" validate:"max=128"`
AlertChannel AlertChannel `json:"alert_channel" validate:"required,oneof=telegram webhook"` AlertChannel AlertChannel `json:"alert_channel" validate:"required,oneof=telegram webhook"`
TelegramBot string `json:"telegram_bot" validate:"required_if=AlertChannel telegram"` TelegramBot string `json:"telegram_bot" validate:"required_if=AlertChannel telegram"`
TelegramChat string `json:"telegram_chat" validate:"required_if=AlertChannel telegram"` TelegramChat string `json:"telegram_chat" validate:"required_if=AlertChannel telegram"`
WebhookURL string `json:"webhook_url" validate:"required_if=AlertChannel webhook,omitempty,url"` WebhookURL string `json:"webhook_url" validate:"required_if=AlertChannel webhook,omitempty,url"`
Metadata json.RawMessage `json:"metadata"` Metadata json.RawMessage `json:"metadata"`
TurnstileResp string `json:"cf_turnstile_response" validate:"required"` TurnstileResp string `json:"cf_turnstile_response" validate:"required"`
} }

View File

@ -40,6 +40,14 @@ type TriggerResponse struct {
type Generator interface { type Generator interface {
Type() token.Type Type() token.Type
Generate(ctx context.Context, t *token.Token, baseURL string) (Artifact, error) Generate(
Trigger(ctx context.Context, t *token.Token, r *http.Request) (*event.Event, *TriggerResponse, error) ctx context.Context,
t *token.Token,
baseURL string,
) (Artifact, error)
Trigger(
ctx context.Context,
t *token.Token,
r *http.Request,
) (*event.Event, *TriggerResponse, error)
} }

View File

@ -29,12 +29,18 @@ func TestTransparentGIF_Length(t *testing.T) {
} }
func TestTransparentGIF_MagicBytes(t *testing.T) { func TestTransparentGIF_MagicBytes(t *testing.T) {
require.True(t, require.True(
t,
bytes.HasPrefix(pixel.TransparentGIF, gif89aMagic), bytes.HasPrefix(pixel.TransparentGIF, gif89aMagic),
"expected GIF89a magic prefix, got % x", pixel.TransparentGIF[:len(gif89aMagic)], "expected GIF89a magic prefix, got % x",
pixel.TransparentGIF[:len(gif89aMagic)],
)
require.Equal(
t,
gifTrailer,
pixel.TransparentGIF[len(pixel.TransparentGIF)-1],
"expected trailing GIF terminator 0x3B",
) )
require.Equal(t, gifTrailer, pixel.TransparentGIF[len(pixel.TransparentGIF)-1],
"expected trailing GIF terminator 0x3B")
} }
func TestTransparentGIF_DecodesAsImageGIF(t *testing.T) { func TestTransparentGIF_DecodesAsImageGIF(t *testing.T) {

View File

@ -41,12 +41,21 @@ func TestBuild_PendingTypesNotYetRegistered(t *testing.T) {
} }
for _, tt := range pending { for _, tt := range pending {
_, ok := reg[tt] _, ok := reg[tt]
require.False(t, ok, require.False(
"type %q is not yet registered (subsequent phases will add it); registry must not claim it", tt) t,
ok,
"type %q is not yet registered (subsequent phases will add it); registry must not claim it",
tt,
)
} }
} }
func TestBuild_OnlyExpectedTypesPresentInPhase2(t *testing.T) { func TestBuild_OnlyExpectedTypesPresentInPhase2(t *testing.T) {
reg := registry.Build(registry.Config{BaseURL: testBaseURL}) reg := registry.Build(registry.Config{BaseURL: testBaseURL})
require.Len(t, reg, 1, "Phase 2 registers exactly one generator (webbug); other phases append") require.Len(
t,
reg,
1,
"Phase 2 registers exactly one generator (webbug); other phases append",
)
} }

View File

@ -34,14 +34,22 @@ func New() *Generator { return &Generator{} }
func (g *Generator) Type() token.Type { return token.TypeWebbug } func (g *Generator) Type() token.Type { return token.TypeWebbug }
func (g *Generator) Generate(_ context.Context, t *token.Token, baseURL string) (generators.Artifact, error) { func (g *Generator) Generate(
_ context.Context,
t *token.Token,
baseURL string,
) (generators.Artifact, error) {
return generators.Artifact{ return generators.Artifact{
Kind: generators.KindURL, Kind: generators.KindURL,
URL: strings.TrimRight(baseURL, "/") + triggerPathPrefix + t.ID, URL: strings.TrimRight(baseURL, "/") + triggerPathPrefix + t.ID,
}, nil }, nil
} }
func (g *Generator) Trigger(_ context.Context, t *token.Token, r *http.Request) (*event.Event, *generators.TriggerResponse, error) { func (g *Generator) Trigger(
_ context.Context,
t *token.Token,
r *http.Request,
) (*event.Event, *generators.TriggerResponse, error) {
tokenID := "" tokenID := ""
if t != nil { if t != nil {
tokenID = t.ID tokenID = t.ID

View File

@ -70,7 +70,11 @@ func TestGenerate_ReturnsURLArtifact(t *testing.T) {
for _, tc := range cases { for _, tc := range cases {
tc := tc tc := tc
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
art, err := g.Generate(context.Background(), newWebbugToken(tc.id), tc.baseURL) art, err := g.Generate(
context.Background(),
newWebbugToken(tc.id),
tc.baseURL,
)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, generators.KindURL, art.Kind) require.Equal(t, generators.KindURL, art.Kind)
require.Equal(t, tc.wantURL, art.URL) require.Equal(t, tc.wantURL, art.URL)
@ -82,98 +86,115 @@ func TestTrigger_RecordsEventWithRequestMetadata(t *testing.T) {
g := webbug.New() g := webbug.New()
tok := newWebbugToken("token1") tok := newWebbugToken("token1")
t.Run("captures token id, source ip, user agent, referer", func(t *testing.T) { t.Run(
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil) "captures token id, source ip, user agent, referer",
r.Header.Set("CF-Connecting-IP", "203.0.113.50") func(t *testing.T) {
r.Header.Set("User-Agent", "Mozilla/5.0 (X11; Linux x86_64)") r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
r.Header.Set("Referer", "https://victim.example.com/inbox") r.Header.Set("CF-Connecting-IP", "203.0.113.50")
r.Header.Set("User-Agent", "Mozilla/5.0 (X11; Linux x86_64)")
r.Header.Set("Referer", "https://victim.example.com/inbox")
evt, _, err := g.Trigger(context.Background(), tok, r) evt, _, err := g.Trigger(context.Background(), tok, r)
require.NoError(t, err) require.NoError(t, err)
require.NotNil(t, evt) require.NotNil(t, evt)
require.Equal(t, "token1", evt.TokenID) require.Equal(t, "token1", evt.TokenID)
require.Equal(t, "203.0.113.50", evt.SourceIP) require.Equal(t, "203.0.113.50", evt.SourceIP)
require.NotNil(t, evt.UserAgent) require.NotNil(t, evt.UserAgent)
require.Equal(t, "Mozilla/5.0 (X11; Linux x86_64)", *evt.UserAgent) require.Equal(t, "Mozilla/5.0 (X11; Linux x86_64)", *evt.UserAgent)
require.NotNil(t, evt.Referer) require.NotNil(t, evt.Referer)
require.Equal(t, "https://victim.example.com/inbox", *evt.Referer) require.Equal(t, "https://victim.example.com/inbox", *evt.Referer)
}) },
)
t.Run("source ip precedence CF > XFF (last) > XRI > RemoteAddr", func(t *testing.T) { t.Run(
cases := []struct { "source ip precedence CF > XFF (last) > XRI > RemoteAddr",
name string func(t *testing.T) {
headers map[string]string cases := []struct {
remote string name string
wantIP string headers map[string]string
}{ remote string
{ wantIP string
name: "CF wins over XFF and XRI", }{
headers: map[string]string{ {
"CF-Connecting-IP": "203.0.113.10", name: "CF wins over XFF and XRI",
"X-Forwarded-For": "198.51.100.1, 198.51.100.2", headers: map[string]string{
"X-Real-IP": "192.0.2.99", "CF-Connecting-IP": "203.0.113.10",
"X-Forwarded-For": "198.51.100.1, 198.51.100.2",
"X-Real-IP": "192.0.2.99",
},
remote: "127.0.0.1:9999",
wantIP: "203.0.113.10",
}, },
remote: "127.0.0.1:9999", {
wantIP: "203.0.113.10", name: "XFF rightmost wins over XRI when no CF",
}, headers: map[string]string{
{ "X-Forwarded-For": "198.51.100.1, 198.51.100.7",
name: "XFF rightmost wins over XRI when no CF", "X-Real-IP": "192.0.2.99",
headers: map[string]string{ },
"X-Forwarded-For": "198.51.100.1, 198.51.100.7", remote: "127.0.0.1:9999",
"X-Real-IP": "192.0.2.99", wantIP: "198.51.100.7",
}, },
remote: "127.0.0.1:9999", {
wantIP: "198.51.100.7", name: "XRI when no CF or XFF",
}, headers: map[string]string{
{ "X-Real-IP": "192.0.2.99",
name: "XRI when no CF or XFF", },
headers: map[string]string{ remote: "127.0.0.1:9999",
"X-Real-IP": "192.0.2.99", wantIP: "192.0.2.99",
}, },
remote: "127.0.0.1:9999", {
wantIP: "192.0.2.99", name: "RemoteAddr when no proxy headers",
}, headers: nil,
{ remote: "127.0.0.1:9999",
name: "RemoteAddr when no proxy headers", wantIP: "127.0.0.1:9999",
headers: nil,
remote: "127.0.0.1:9999",
wantIP: "127.0.0.1:9999",
},
{
name: "CF value is trimmed of whitespace",
headers: map[string]string{
"CF-Connecting-IP": " 203.0.113.10 ",
}, },
remote: "127.0.0.1:9999", {
wantIP: "203.0.113.10", name: "CF value is trimmed of whitespace",
}, headers: map[string]string{
} "CF-Connecting-IP": " 203.0.113.10 ",
for _, tc := range cases { },
tc := tc remote: "127.0.0.1:9999",
t.Run(tc.name, func(t *testing.T) { wantIP: "203.0.113.10",
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil) },
for k, v := range tc.headers { }
r.Header.Set(k, v) for _, tc := range cases {
} tc := tc
r.RemoteAddr = tc.remote t.Run(tc.name, func(t *testing.T) {
evt, _, err := g.Trigger(context.Background(), tok, r) r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
require.NoError(t, err) for k, v := range tc.headers {
require.Equal(t, tc.wantIP, evt.SourceIP) r.Header.Set(k, v)
}) }
} r.RemoteAddr = tc.remote
}) evt, _, err := g.Trigger(context.Background(), tok, r)
require.NoError(t, err)
require.Equal(t, tc.wantIP, evt.SourceIP)
})
}
},
)
t.Run("missing user agent and referer record as nil pointers", func(t *testing.T) { t.Run(
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil) "missing user agent and referer record as nil pointers",
r.Header.Del("User-Agent") func(t *testing.T) {
r.Header.Del("Referer") r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
r.Header.Set("CF-Connecting-IP", "203.0.113.5") r.Header.Del("User-Agent")
r.Header.Del("Referer")
r.Header.Set("CF-Connecting-IP", "203.0.113.5")
evt, _, err := g.Trigger(context.Background(), tok, r) evt, _, err := g.Trigger(context.Background(), tok, r)
require.NoError(t, err) require.NoError(t, err)
require.Nil(t, evt.UserAgent, "absent user agent must map to nil, not empty string") require.Nil(
require.Nil(t, evt.Referer, "absent referer must map to nil, not empty string") t,
}) evt.UserAgent,
"absent user agent must map to nil, not empty string",
)
require.Nil(
t,
evt.Referer,
"absent referer must map to nil, not empty string",
)
},
)
} }
func TestTrigger_ResponseIs43ByteGIF(t *testing.T) { func TestTrigger_ResponseIs43ByteGIF(t *testing.T) {
@ -199,14 +220,27 @@ func TestTrigger_TokenNotFound_StillReturnsGIF(t *testing.T) {
r.Header.Set("User-Agent", "curl/8.0.0") r.Header.Set("User-Agent", "curl/8.0.0")
evt, resp, err := g.Trigger(context.Background(), nil, r) evt, resp, err := g.Trigger(context.Background(), nil, r)
require.NoError(t, err, "nil-token path must not error (spec §8.5 defense-in-depth)") require.NoError(
t,
err,
"nil-token path must not error (spec §8.5 defense-in-depth)",
)
require.NotNil(t, resp, "nil-token path must still return GIF response") require.NotNil(t, resp, "nil-token path must still return GIF response")
require.Equal(t, http.StatusOK, resp.StatusCode) require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, pixel.ContentType, resp.ContentType) require.Equal(t, pixel.ContentType, resp.ContentType)
require.Equal(t, pixel.TransparentGIF, resp.Body) require.Equal(t, pixel.TransparentGIF, resp.Body)
require.NotNil(t, evt, "event value still produced for forensic continuity") require.NotNil(t, evt, "event value still produced for forensic continuity")
require.Empty(t, evt.TokenID, "empty TokenID signals nil token; persistence layer decides what to do") require.Empty(
require.Equal(t, "203.0.113.100", evt.SourceIP, "source IP still captured for forensics") t,
evt.TokenID,
"empty TokenID signals nil token; persistence layer decides what to do",
)
require.Equal(
t,
"203.0.113.100",
evt.SourceIP,
"source IP still captured for forensics",
)
require.NotNil(t, evt.UserAgent) require.NotNil(t, evt.UserAgent)
require.Equal(t, "curl/8.0.0", *evt.UserAgent) require.Equal(t, "curl/8.0.0", *evt.UserAgent)
} }

View File

@ -8,6 +8,7 @@ import (
"database/sql" "database/sql"
"errors" "errors"
"fmt" "fmt"
"log/slog"
"github.com/jmoiron/sqlx" "github.com/jmoiron/sqlx"
) )
@ -41,7 +42,12 @@ func (r *Repository) Insert(ctx context.Context, t *Token) error {
if err != nil { if err != nil {
return fmt.Errorf("prepare insert token: %w", err) return fmt.Errorf("prepare insert token: %w", err)
} }
defer stmt.Close() defer func() {
if cerr := stmt.Close(); cerr != nil {
slog.WarnContext(ctx, "close prepared stmt",
"op", "insert_token", "error", cerr)
}
}()
if err := stmt.GetContext(ctx, t, t); err != nil { if err := stmt.GetContext(ctx, t, t); err != nil {
return fmt.Errorf("insert token: %w", err) return fmt.Errorf("insert token: %w", err)
@ -68,7 +74,10 @@ func (r *Repository) GetByID(ctx context.Context, id string) (*Token, error) {
return &t, nil return &t, nil
} }
func (r *Repository) GetByManageID(ctx context.Context, manageID string) (*Token, error) { func (r *Repository) GetByManageID(
ctx context.Context,
manageID string,
) (*Token, error) {
var t Token var t Token
q := `SELECT ` + selectColumns + ` FROM tokens WHERE manage_id = $1` q := `SELECT ` + selectColumns + ` FROM tokens WHERE manage_id = $1`
err := r.db.GetContext(ctx, &t, q, manageID) err := r.db.GetContext(ctx, &t, q, manageID)
@ -81,7 +90,10 @@ func (r *Repository) GetByManageID(ctx context.Context, manageID string) (*Token
return &t, nil return &t, nil
} }
func (r *Repository) DeleteByManageID(ctx context.Context, manageID string) error { func (r *Repository) DeleteByManageID(
ctx context.Context,
manageID string,
) error {
res, err := r.db.ExecContext(ctx, res, err := r.db.ExecContext(ctx,
`DELETE FROM tokens WHERE manage_id = $1`, manageID) `DELETE FROM tokens WHERE manage_id = $1`, manageID)
if err != nil { if err != nil {
@ -97,7 +109,10 @@ func (r *Repository) DeleteByManageID(ctx context.Context, manageID string) erro
return nil return nil
} }
func (r *Repository) IncrementTriggerCount(ctx context.Context, id string) error { func (r *Repository) IncrementTriggerCount(
ctx context.Context,
id string,
) error {
_, err := r.db.ExecContext(ctx, ` _, err := r.db.ExecContext(ctx, `
UPDATE tokens UPDATE tokens
SET trigger_count = trigger_count + 1, SET trigger_count = trigger_count + 1,
@ -109,7 +124,11 @@ func (r *Repository) IncrementTriggerCount(ctx context.Context, id string) error
return nil return nil
} }
func (r *Repository) SetEnabled(ctx context.Context, id string, enabled bool) error { func (r *Repository) SetEnabled(
ctx context.Context,
id string,
enabled bool,
) error {
res, err := r.db.ExecContext(ctx, res, err := r.db.ExecContext(ctx,
`UPDATE tokens SET enabled = $2 WHERE id = $1`, id, enabled) `UPDATE tokens SET enabled = $2 WHERE id = $1`, id, enabled)
if err != nil { if err != nil {
@ -130,7 +149,10 @@ type ListOptions struct {
Offset int Offset int
} }
func (r *Repository) ListAll(ctx context.Context, opts ListOptions) ([]Token, error) { func (r *Repository) ListAll(
ctx context.Context,
opts ListOptions,
) ([]Token, error) {
if opts.Limit <= 0 { if opts.Limit <= 0 {
opts.Limit = defaultListLimit opts.Limit = defaultListLimit
} }
@ -138,7 +160,13 @@ func (r *Repository) ListAll(ctx context.Context, opts ListOptions) ([]Token, er
ORDER BY created_at DESC ORDER BY created_at DESC
LIMIT $1 OFFSET $2` LIMIT $1 OFFSET $2`
var tokens []Token var tokens []Token
if err := r.db.SelectContext(ctx, &tokens, q, opts.Limit, opts.Offset); err != nil { if err := r.db.SelectContext(
ctx,
&tokens,
q,
opts.Limit,
opts.Offset,
); err != nil {
return nil, fmt.Errorf("list all tokens: %w", err) return nil, fmt.Errorf("list all tokens: %w", err)
} }
return tokens, nil return tokens, nil
@ -146,7 +174,11 @@ func (r *Repository) ListAll(ctx context.Context, opts ListOptions) ([]Token, er
func (r *Repository) CountAll(ctx context.Context) (int64, error) { func (r *Repository) CountAll(ctx context.Context) (int64, error) {
var n int64 var n int64
if err := r.db.GetContext(ctx, &n, `SELECT COUNT(*) FROM tokens`); err != nil { if err := r.db.GetContext(
ctx,
&n,
`SELECT COUNT(*) FROM tokens`,
); err != nil {
return 0, fmt.Errorf("count tokens: %w", err) return 0, fmt.Errorf("count tokens: %w", err)
} }
return n, nil return n, nil