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:
|
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:
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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...)
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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())
|
||||||
|
|
|
||||||
|
|
@ -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"`
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue