package main import ( "context" "errors" "fmt" "log/slog" "net/http" "os" "os/signal" "syscall" "time" "github.com/jackc/pgx/v5/pgxpool" "github.com/redis/go-redis/v9" "git.dcrubro.com/dcrubro/hammy-backend/internal/api" "git.dcrubro.com/dcrubro/hammy-backend/internal/config" "git.dcrubro.com/dcrubro/hammy-backend/internal/db" "git.dcrubro.com/dcrubro/hammy-backend/internal/ratelimit" ) // version is stamped at build time: // go build -ldflags "-X main.version=$(git describe --tags --always --dirty)" ./cmd/hammyd var version = "dev" func main() { // All real work happens in run() so that deferred cleanup actually runs. // os.Exit skips defers, so calling it anywhere but here leaks the pool and // drops in-flight requests. if err := run(); err != nil { fmt.Fprintf(os.Stderr, "hammyd: %v\n", err) os.Exit(1) } } func run() error { // NotifyContext cancels on SIGINT/SIGTERM, which is what gives the server // a chance to finish in-flight requests instead of being killed mid-write. ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() args := os.Args[1:] cmd := "serve" if len(args) > 0 { cmd, args = args[0], args[1:] } switch cmd { case "serve": return runServe(ctx, args) case "keygen": return runKeygen(ctx, args) case "version": fmt.Println(version) return nil case "help", "-h", "--help": usage() return nil default: usage() return fmt.Errorf("unknown command %q", cmd) } } func usage() { fmt.Fprintf(os.Stderr, `hammyd %s Usage: hammyd serve Run the API server (default) hammyd keygen -email … Issue an API key hammyd version Configuration comes from the environment: HAMMY_DSN Postgres connection string (required) HAMMY_ADDR Listen address (default 127.0.0.1:8080) HAMMY_REDIS_ADDR Redis for shared rate limiting (optional) HAMMY_MAX_CONNS Postgres pool size (default 10) HAMMY_LOG_LEVEL debug, info, warn, error (default info) HAMMY_LOG_JSON true for structured logs HAMMY_ENDUSER_HEADER Header carrying the end-user id (default X-Hammy-User) `, version) } func runServe(ctx context.Context, _ []string) error { cfg, err := config.Load() if err != nil { return err } logger := newLogger(cfg) slog.SetDefault(logger) logger.Info("starting", "version", version, "config", cfg.Redacted()) // ---- Postgres --------------------------------------------------------- poolCfg, err := pgxpool.ParseConfig(cfg.DSN) if err != nil { return fmt.Errorf("parsing DSN: %w", err) } // Postgres forks a backend process per connection, so the pool size is a // real resource on the server, not just a client-side knob. poolCfg.MaxConns = cfg.MaxConns poolCfg.MaxConnLifetime = time.Hour poolCfg.MaxConnIdleTime = 15 * time.Minute pool, err := pgxpool.NewWithConfig(ctx, poolCfg) if err != nil { return fmt.Errorf("connecting to postgres: %w", err) } defer pool.Close() // Fail at startup rather than on the first request. A backend that boots // happily and 500s on everything is much harder to diagnose. pingCtx, cancel := context.WithTimeout(ctx, 10*time.Second) defer cancel() if err := pool.Ping(pingCtx); err != nil { return fmt.Errorf("pinging postgres: %w", err) } logger.Info("postgres connected", "max_conns", cfg.MaxConns) // ---- Redis, or not ---------------------------------------------------- var limiter ratelimit.Limiter if cfg.RedisAddr == "" { // A single self-hosted instance does not need Redis to have working // rate limits. Several instances do: each process would keep its own // counters, so N processes allow N times the limit. logger.Warn("HAMMY_REDIS_ADDR is unset, using an in-process limiter. " + "Correct for one instance, wrong for several.") mem := ratelimit.NewMemory() limiter = mem go sweepPeriodically(ctx, mem, time.Minute, logger) } else { rdb := redis.NewClient(&redis.Options{Addr: cfg.RedisAddr, DB: cfg.RedisDB}) defer rdb.Close() if err := rdb.Ping(ctx).Err(); err != nil { return fmt.Errorf("connecting to redis: %w", err) } logger.Info("redis connected", "addr", cfg.RedisAddr) limiter = ratelimit.NewRedis(ratelimit.GoRedis{Client: rdb}) } // ---- HTTP ------------------------------------------------------------- queries := db.New(pool) router := &api.Router{ Auth: &api.Authenticator{ Keys: api.PgKeyStore{Q: queries}, Limiter: limiter, Logger: logger, EndUserHeader: cfg.EndUserHeader, }, Logger: logger, Version: version, Ready: func() error { // Readiness checks dependencies; liveness deliberately does not. c, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() return pool.Ping(c) }, } srv := &http.Server{ Addr: cfg.Addr, Handler: router.Handler(), // ReadHeaderTimeout is the one that matters: without it, a client can // hold a connection open by dribbling headers forever (Slowloris). ReadHeaderTimeout: cfg.ReadHeaderTimeout, ReadTimeout: 30 * time.Second, WriteTimeout: 60 * time.Second, IdleTimeout: 2 * time.Minute, ErrorLog: slog.NewLogLogger(logger.Handler(), slog.LevelError), } errCh := make(chan error, 1) go func() { logger.Info("listening", "addr", cfg.Addr) // ErrServerClosed is the normal result of Shutdown, not a failure. if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { errCh <- err } }() select { case err := <-errCh: return fmt.Errorf("listening: %w", err) case <-ctx.Done(): logger.Info("shutdown signal received, draining") } // A FRESH context: ctx is already cancelled, and passing it to Shutdown // would abort in-flight requests immediately, which is the opposite of // what a graceful shutdown is for. shutdownCtx, cancelShutdown := context.WithTimeout( context.Background(), cfg.ShutdownTimeout) defer cancelShutdown() if err := srv.Shutdown(shutdownCtx); err != nil { return fmt.Errorf("shutdown: %w", err) } logger.Info("stopped cleanly") return nil } // sweepPeriodically drops expired windows from the in-process limiter. Without // it the map grows for every key and end user ever seen. func sweepPeriodically(ctx context.Context, m *ratelimit.Memory, every time.Duration, logger *slog.Logger) { t := time.NewTicker(every) defer t.Stop() for { select { case <-ctx.Done(): return case <-t.C: if n := m.Sweep(); n > 0 { logger.Debug("swept rate limit windows", "removed", n) } } } } func newLogger(cfg config.Config) *slog.Logger { opts := &slog.HandlerOptions{Level: cfg.LogLevel} if cfg.LogJSON { return slog.New(slog.NewJSONHandler(os.Stdout, opts)) } return slog.New(slog.NewTextHandler(os.Stdout, opts)) }