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:
parent
bb95f74579
commit
c39b56b9af
|
|
@ -72,23 +72,26 @@ linters:
|
|||
gosec:
|
||||
excludes:
|
||||
- G104
|
||||
- G706
|
||||
|
||||
sloglint:
|
||||
no-mixed-args: true
|
||||
kv-only: true
|
||||
context: all
|
||||
|
||||
exclusions:
|
||||
paths:
|
||||
- vendor
|
||||
- testdata
|
||||
rules:
|
||||
- path: _test\.go
|
||||
linters:
|
||||
- funlen
|
||||
- dupl
|
||||
- goconst
|
||||
|
||||
issues:
|
||||
max-same-issues: 50
|
||||
exclude-dirs:
|
||||
- vendor
|
||||
- testdata
|
||||
exclude-rules:
|
||||
- path: _test\.go
|
||||
linters:
|
||||
- funlen
|
||||
- dupl
|
||||
- goconst
|
||||
|
||||
|
||||
formatters:
|
||||
|
|
|
|||
|
|
@ -24,7 +24,10 @@ import (
|
|||
"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() {
|
||||
configPath := flag.String("config", "config.yaml", "path to config file")
|
||||
|
|
@ -37,7 +40,11 @@ func main() {
|
|||
}
|
||||
|
||||
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()
|
||||
|
||||
cfg, err := config.Load(configPath)
|
||||
|
|
@ -52,14 +59,7 @@ func run(configPath string) error {
|
|||
"environment", cfg.App.Environment,
|
||||
)
|
||||
|
||||
var telemetry *core.Telemetry
|
||||
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
|
||||
}
|
||||
}
|
||||
telemetry := initTelemetry(ctx, cfg, logger)
|
||||
|
||||
db, err := core.NewDatabase(ctx, cfg.Database)
|
||||
if err != nil {
|
||||
|
|
@ -67,15 +67,13 @@ func run(configPath string) error {
|
|||
}
|
||||
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)
|
||||
}
|
||||
logger.Info("migrations applied")
|
||||
|
||||
tokenRepo := token.NewRepository(db.DB)
|
||||
eventRepo := event.NewRepository(db.DB)
|
||||
_ = tokenRepo
|
||||
_ = eventRepo
|
||||
_ = token.NewRepository(db.DB)
|
||||
_ = event.NewRepository(db.DB)
|
||||
|
||||
rdb, err := core.NewRedis(ctx, cfg.Redis)
|
||||
if err != nil {
|
||||
|
|
@ -84,7 +82,43 @@ func run(configPath string) error {
|
|||
logger.Info("redis connected")
|
||||
|
||||
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{
|
||||
ServerConfig: cfg.Server,
|
||||
HealthHandler: healthH,
|
||||
|
|
@ -96,7 +130,10 @@ func run(configPath string) error {
|
|||
r.Use(middleware.Logger(logger))
|
||||
r.Use(
|
||||
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,
|
||||
}).Handler,
|
||||
)
|
||||
|
|
@ -104,29 +141,31 @@ func run(configPath string) error {
|
|||
r.Use(middleware.CORS(cfg.CORS))
|
||||
|
||||
healthH.RegisterRoutes(r)
|
||||
r.Route("/api", func(_ chi.Router) {})
|
||||
return srv
|
||||
}
|
||||
|
||||
r.Route("/api", func(_ chi.Router) {
|
||||
})
|
||||
|
||||
errChan := make(chan error, 1)
|
||||
go func() { errChan <- srv.Start() }()
|
||||
|
||||
select {
|
||||
case err := <-errChan:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
logger.Info("shutdown signal received")
|
||||
}
|
||||
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(),
|
||||
cfg.Server.ShutdownTimeout+drainDelay+5*time.Second)
|
||||
func gracefulShutdown(
|
||||
cfg *config.Config,
|
||||
logger *slog.Logger,
|
||||
srv *server.Server,
|
||||
telemetry *core.Telemetry,
|
||||
rdb *core.Redis,
|
||||
db *core.Database,
|
||||
) error {
|
||||
shutdownCtx, cancel := context.WithTimeout(
|
||||
context.Background(),
|
||||
cfg.Server.ShutdownTimeout+drainDelay+shutdownGraceExtra,
|
||||
)
|
||||
defer cancel()
|
||||
|
||||
if err := srv.Shutdown(shutdownCtx, drainDelay); err != nil {
|
||||
logger.Error("server shutdown error", "error", err)
|
||||
}
|
||||
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 {
|
||||
logger.Error("redis close error", "error", err)
|
||||
|
|
|
|||
|
|
@ -98,13 +98,19 @@ func Load(configPath string) (*Config, error) {
|
|||
}
|
||||
|
||||
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)
|
||||
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)
|
||||
return
|
||||
}
|
||||
|
|
|
|||
|
|
@ -51,7 +51,12 @@ func NewTelemetry(
|
|||
otlptracegrpc.WithTLSCredentials(insecure.NewCredentials()),
|
||||
)
|
||||
} else {
|
||||
opts = append(opts, otlptracegrpc.WithTLSCredentials(credentials.NewClientTLSFromCert(nil, "")))
|
||||
opts = append(
|
||||
opts,
|
||||
otlptracegrpc.WithTLSCredentials(
|
||||
credentials.NewClientTLSFromCert(nil, ""),
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
exporter, err := otlptracegrpc.New(ctx, opts...)
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import (
|
|||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/jmoiron/sqlx"
|
||||
|
|
@ -50,7 +51,12 @@ func (r *Repository) Insert(ctx context.Context, e *Event) error {
|
|||
if err != nil {
|
||||
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 {
|
||||
return fmt.Errorf("insert event: %w", err)
|
||||
|
|
@ -90,7 +96,14 @@ func (r *Repository) ListByToken(
|
|||
LIMIT $3`
|
||||
|
||||
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 {
|
||||
return ListResult{}, fmt.Errorf("list events: %w", err)
|
||||
}
|
||||
|
|
@ -112,7 +125,10 @@ func (r *Repository) ListByToken(
|
|||
}, 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
|
||||
err := r.db.GetContext(ctx, &n,
|
||||
`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(
|
||||
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 {
|
||||
q := `
|
||||
UPDATE events
|
||||
|
|
@ -136,8 +155,14 @@ UPDATE events
|
|||
ORDER BY id DESC
|
||||
LIMIT 1
|
||||
)`
|
||||
res, err := r.db.ExecContext(ctx, q,
|
||||
tokenID, sourceIP, []byte(fingerprint), fmt.Sprintf("%d milliseconds", window.Milliseconds()))
|
||||
res, err := r.db.ExecContext(
|
||||
ctx,
|
||||
q,
|
||||
tokenID,
|
||||
sourceIP,
|
||||
[]byte(fingerprint),
|
||||
fmt.Sprintf("%d milliseconds", window.Milliseconds()),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("attach fingerprint: %w", err)
|
||||
}
|
||||
|
|
@ -172,7 +197,10 @@ UPDATE events
|
|||
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 {
|
||||
return 0, errors.New("perTokenLimit must be positive")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,7 +36,11 @@ func NewTestDB(t *testing.T) *sql.DB {
|
|||
require.NoError(t, err)
|
||||
|
||||
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")
|
||||
|
|
@ -46,7 +50,9 @@ func NewTestDB(t *testing.T) *sql.DB {
|
|||
require.NoError(t, err)
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = db.Close()
|
||||
if closeErr := db.Close(); closeErr != nil {
|
||||
t.Logf("db close: %v", closeErr)
|
||||
}
|
||||
})
|
||||
|
||||
require.NoError(t, db.Ping())
|
||||
|
|
|
|||
|
|
@ -9,13 +9,13 @@ import (
|
|||
)
|
||||
|
||||
type CreateRequest struct {
|
||||
Type Type `json:"type" validate:"required,oneof=webbug slowredirect docx pdf kubeconfig envfile mysql"`
|
||||
Memo string `json:"memo" validate:"max=256"`
|
||||
Filename string `json:"filename" validate:"max=128"`
|
||||
AlertChannel AlertChannel `json:"alert_channel" validate:"required,oneof=telegram webhook"`
|
||||
TelegramBot string `json:"telegram_bot" 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"`
|
||||
Type Type `json:"type" validate:"required,oneof=webbug slowredirect docx pdf kubeconfig envfile mysql"`
|
||||
Memo string `json:"memo" validate:"max=256"`
|
||||
Filename string `json:"filename" validate:"max=128"`
|
||||
AlertChannel AlertChannel `json:"alert_channel" validate:"required,oneof=telegram webhook"`
|
||||
TelegramBot string `json:"telegram_bot" 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"`
|
||||
Metadata json.RawMessage `json:"metadata"`
|
||||
TurnstileResp string `json:"cf_turnstile_response" validate:"required"`
|
||||
}
|
||||
|
|
|
|||
|
|
@ -40,6 +40,14 @@ type TriggerResponse struct {
|
|||
|
||||
type Generator interface {
|
||||
Type() token.Type
|
||||
Generate(ctx context.Context, t *token.Token, baseURL string) (Artifact, error)
|
||||
Trigger(ctx context.Context, t *token.Token, r *http.Request) (*event.Event, *TriggerResponse, error)
|
||||
Generate(
|
||||
ctx context.Context,
|
||||
t *token.Token,
|
||||
baseURL string,
|
||||
) (Artifact, error)
|
||||
Trigger(
|
||||
ctx context.Context,
|
||||
t *token.Token,
|
||||
r *http.Request,
|
||||
) (*event.Event, *TriggerResponse, error)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -29,12 +29,18 @@ func TestTransparentGIF_Length(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestTransparentGIF_MagicBytes(t *testing.T) {
|
||||
require.True(t,
|
||||
require.True(
|
||||
t,
|
||||
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) {
|
||||
|
|
|
|||
|
|
@ -41,12 +41,21 @@ func TestBuild_PendingTypesNotYetRegistered(t *testing.T) {
|
|||
}
|
||||
for _, tt := range pending {
|
||||
_, ok := reg[tt]
|
||||
require.False(t, ok,
|
||||
"type %q is not yet registered (subsequent phases will add it); registry must not claim it", tt)
|
||||
require.False(
|
||||
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) {
|
||||
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",
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -34,14 +34,22 @@ func New() *Generator { return &Generator{} }
|
|||
|
||||
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{
|
||||
Kind: generators.KindURL,
|
||||
URL: strings.TrimRight(baseURL, "/") + triggerPathPrefix + t.ID,
|
||||
}, 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 := ""
|
||||
if t != nil {
|
||||
tokenID = t.ID
|
||||
|
|
|
|||
|
|
@ -70,7 +70,11 @@ func TestGenerate_ReturnsURLArtifact(t *testing.T) {
|
|||
for _, tc := range cases {
|
||||
tc := tc
|
||||
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.Equal(t, generators.KindURL, art.Kind)
|
||||
require.Equal(t, tc.wantURL, art.URL)
|
||||
|
|
@ -82,98 +86,115 @@ func TestTrigger_RecordsEventWithRequestMetadata(t *testing.T) {
|
|||
g := webbug.New()
|
||||
tok := newWebbugToken("token1")
|
||||
|
||||
t.Run("captures token id, source ip, user agent, referer", func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
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")
|
||||
t.Run(
|
||||
"captures token id, source ip, user agent, referer",
|
||||
func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
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)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, evt)
|
||||
require.Equal(t, "token1", evt.TokenID)
|
||||
require.Equal(t, "203.0.113.50", evt.SourceIP)
|
||||
require.NotNil(t, evt.UserAgent)
|
||||
require.Equal(t, "Mozilla/5.0 (X11; Linux x86_64)", *evt.UserAgent)
|
||||
require.NotNil(t, evt.Referer)
|
||||
require.Equal(t, "https://victim.example.com/inbox", *evt.Referer)
|
||||
})
|
||||
evt, _, err := g.Trigger(context.Background(), tok, r)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, evt)
|
||||
require.Equal(t, "token1", evt.TokenID)
|
||||
require.Equal(t, "203.0.113.50", evt.SourceIP)
|
||||
require.NotNil(t, evt.UserAgent)
|
||||
require.Equal(t, "Mozilla/5.0 (X11; Linux x86_64)", *evt.UserAgent)
|
||||
require.NotNil(t, 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) {
|
||||
cases := []struct {
|
||||
name 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",
|
||||
"X-Forwarded-For": "198.51.100.1, 198.51.100.2",
|
||||
"X-Real-IP": "192.0.2.99",
|
||||
t.Run(
|
||||
"source ip precedence CF > XFF (last) > XRI > RemoteAddr",
|
||||
func(t *testing.T) {
|
||||
cases := []struct {
|
||||
name 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",
|
||||
"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",
|
||||
"X-Real-IP": "192.0.2.99",
|
||||
{
|
||||
name: "XFF rightmost wins over XRI when no CF",
|
||||
headers: map[string]string{
|
||||
"X-Forwarded-For": "198.51.100.1, 198.51.100.7",
|
||||
"X-Real-IP": "192.0.2.99",
|
||||
},
|
||||
remote: "127.0.0.1:9999",
|
||||
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{
|
||||
"X-Real-IP": "192.0.2.99",
|
||||
},
|
||||
remote: "127.0.0.1:9999",
|
||||
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",
|
||||
wantIP: "127.0.0.1:9999",
|
||||
},
|
||||
{
|
||||
name: "CF value is trimmed of whitespace",
|
||||
headers: map[string]string{
|
||||
"CF-Connecting-IP": " 203.0.113.10 ",
|
||||
{
|
||||
name: "RemoteAddr when no proxy headers",
|
||||
headers: nil,
|
||||
remote: "127.0.0.1:9999",
|
||||
wantIP: "127.0.0.1:9999",
|
||||
},
|
||||
remote: "127.0.0.1:9999",
|
||||
wantIP: "203.0.113.10",
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
for k, v := range tc.headers {
|
||||
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)
|
||||
})
|
||||
}
|
||||
})
|
||||
{
|
||||
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",
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
for k, v := range tc.headers {
|
||||
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) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
r.Header.Del("User-Agent")
|
||||
r.Header.Del("Referer")
|
||||
r.Header.Set("CF-Connecting-IP", "203.0.113.5")
|
||||
t.Run(
|
||||
"missing user agent and referer record as nil pointers",
|
||||
func(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/c/token1", nil)
|
||||
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)
|
||||
require.NoError(t, err)
|
||||
require.Nil(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")
|
||||
})
|
||||
evt, _, err := g.Trigger(context.Background(), tok, r)
|
||||
require.NoError(t, err)
|
||||
require.Nil(
|
||||
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) {
|
||||
|
|
@ -199,14 +220,27 @@ func TestTrigger_TokenNotFound_StillReturnsGIF(t *testing.T) {
|
|||
r.Header.Set("User-Agent", "curl/8.0.0")
|
||||
|
||||
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.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
require.Equal(t, pixel.ContentType, resp.ContentType)
|
||||
require.Equal(t, pixel.TransparentGIF, resp.Body)
|
||||
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.Equal(t, "203.0.113.100", evt.SourceIP, "source IP still captured for forensics")
|
||||
require.Empty(
|
||||
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.Equal(t, "curl/8.0.0", *evt.UserAgent)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import (
|
|||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/jmoiron/sqlx"
|
||||
)
|
||||
|
|
@ -41,7 +42,12 @@ func (r *Repository) Insert(ctx context.Context, t *Token) error {
|
|||
if err != nil {
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
q := `SELECT ` + selectColumns + ` FROM tokens WHERE manage_id = $1`
|
||||
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
|
||||
}
|
||||
|
||||
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,
|
||||
`DELETE FROM tokens WHERE manage_id = $1`, manageID)
|
||||
if err != nil {
|
||||
|
|
@ -97,7 +109,10 @@ func (r *Repository) DeleteByManageID(ctx context.Context, manageID string) erro
|
|||
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, `
|
||||
UPDATE tokens
|
||||
SET trigger_count = trigger_count + 1,
|
||||
|
|
@ -109,7 +124,11 @@ func (r *Repository) IncrementTriggerCount(ctx context.Context, id string) error
|
|||
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,
|
||||
`UPDATE tokens SET enabled = $2 WHERE id = $1`, id, enabled)
|
||||
if err != nil {
|
||||
|
|
@ -130,7 +149,10 @@ type ListOptions struct {
|
|||
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 {
|
||||
opts.Limit = defaultListLimit
|
||||
}
|
||||
|
|
@ -138,7 +160,13 @@ func (r *Repository) ListAll(ctx context.Context, opts ListOptions) ([]Token, er
|
|||
ORDER BY created_at DESC
|
||||
LIMIT $1 OFFSET $2`
|
||||
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 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) {
|
||||
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 n, nil
|
||||
|
|
|
|||
Loading…
Reference in New Issue