Seed from go-template-repo, renamed to simplexcalc

The template's files at a77fd30, without its history or LICENSE, after
script/rename simplexcalc.

Model: opus-5-5
This commit is contained in:
clawbot
2026-09-26 21:38:57 +00:00
parent e336f46e34
commit f8ce8cef83
79 changed files with 7131 additions and 0 deletions
+72
View File
@@ -0,0 +1,72 @@
// Package app wires the object graph. It is the one place that knows
// which concrete types satisfy the application's dependencies, so every
// other package can be constructed in a test with substitutes.
package app
import (
"log/slog"
"go.uber.org/fx"
"go.uber.org/fx/fxevent"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/database"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/handlers"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/middleware"
"sneak.berlin/go/simplexcalc/internal/render"
"sneak.berlin/go/simplexcalc/internal/server"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// Module is every provider the service needs. Constructor order is
// irrelevant to fx; the order here is the order a reader wants: build
// metadata, logging, configuration, storage, then the HTTP layer.
//
//nolint:gochecknoglobals // an fx module is a declaration, not mutable state.
var Module = fx.Options(
fx.Provide(
globals.New,
logger.New,
config.New,
database.New,
telemetry.NewSentry,
telemetry.NewMetrics,
render.New,
middleware.New,
handlers.New,
server.New,
),
// fx's own lifecycle events go through the application logger, so
// the process emits one stream in one format. Without this, fx
// prints its own plain-text output to stderr and a log pipeline
// gets two formats from one process.
fx.WithLogger(func(l *logger.Logger) fxevent.Logger {
return &fxevent.SlogLogger{Logger: l.Get()}
}),
)
// Invoke forces the graph to be built. fx constructs lazily: without a
// request for the Server, a perfectly valid App would start, construct
// nothing, and serve nothing.
//
//nolint:gochecknoglobals // as above.
var Invoke = fx.Invoke(func(_ *server.Server, log *logger.Logger, cfg *config.Config) {
if cfg.Debug {
log.EnableDebugLogging()
}
log.Identify()
})
// New builds the fx application for `serve`.
func New(opts ...fx.Option) *fx.App {
return fx.New(append([]fx.Option{Module, Invoke}, opts...)...)
}
// DiscardLogger is a logger that writes nothing, for tests that build
// the graph and do not want its startup output.
func DiscardLogger() *slog.Logger {
return slog.New(slog.DiscardHandler)
}
+390
View File
@@ -0,0 +1,390 @@
// Package config loads runtime configuration from the environment (and
// an optional ./.env file) via viper.
//
// The iron rule of this package: a value that is SET but cannot be
// parsed aborts startup. It is never replaced by the default. An
// operator who writes PORT=eighty has said something specific and
// wrong, and starting anyway on port 8080 turns their mistake into a
// silent misconfiguration that only surfaces much later, somewhere
// else. Defaults apply to values that are ABSENT, and to nothing else.
//
// Every parse failure found in one pass is reported together, so a
// broken deployment takes one restart to diagnose rather than five.
package config
import (
"errors"
"fmt"
"net/url"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/dustin/go-humanize"
"github.com/spf13/viper"
"go.uber.org/fx"
// spooky action at a distance!
// this populates the environment
// from a ./.env file automatically
// for development configuration.
// .env contents should be things like
// `PORT=8080`
// (without the backticks, of course)
_ "github.com/joho/godotenv/autoload"
)
// Environment variable names. Bare names, no prefix: this matches the
// other services and keeps a compose file readable.
const (
EnvPort = "PORT"
EnvDataDir = "DATA_DIR"
EnvDBPath = "DB_PATH"
EnvDebug = "DEBUG"
EnvHSTS = "HSTS"
EnvBaseURL = "BASE_URL"
EnvMaxRequestBody = "MAX_REQUEST_BODY"
EnvRequestTimeout = "REQUEST_TIMEOUT"
EnvShutdownGrace = "SHUTDOWN_GRACE"
EnvSentryDSN = "SENTRY_DSN"
EnvSentryEnv = "SENTRY_ENVIRONMENT"
EnvMetricsUser = "METRICS_USER"
EnvMetricsPassword = "METRICS_PASSWORD"
EnvCSRFKey = "CSRF_KEY"
)
// Defaults for values that are absent. A value that is present and
// unparseable never reaches these.
const (
DefaultPort int64 = 8080
DefaultDataDir = "./data"
DefaultBaseURL = "http://localhost:8080"
DefaultMaxRequestBody int64 = 1 << 20 // 1 MiB
DefaultRequestTimeout = 30 * time.Second
DefaultShutdownGrace = 15 * time.Second
DefaultSentryEnv = "development"
)
// Bounds. A value inside the type but outside the range is as
// misconfigured as one that does not parse, and fails the same way.
const (
minPort int64 = 1
maxPort int64 = 65535
// minRequestBody is a floor below which no useful form submission
// fits; maxRequestBody is a ceiling above which the cap is not
// doing its job.
minRequestBody int64 = 1 << 10 // 1 KiB
maxRequestBody int64 = 1 << 26 // 64 MiB
minTimeout = 1 * time.Second
maxTimeout = 10 * time.Minute
// csrfKeyBytes is what gorilla/csrf requires: exactly 32 bytes,
// supplied as csrfKeyHexChars hex characters.
csrfKeyBytes = 32
csrfKeyHexChars = csrfKeyBytes * 2
)
// ErrInvalidConfig is the sentinel every configuration failure wraps,
// so callers can distinguish "the operator got it wrong" from "the
// machine is broken" without string matching.
var ErrInvalidConfig = errors.New("invalid configuration")
// Config is the parsed, validated runtime configuration. Every field
// is final by the time New returns: nothing re-reads the environment
// later, so there is exactly one moment at which configuration can be
// wrong, and it is before the listener opens.
type Config struct {
Port int
DataDir string
DBPath string
Debug bool
HSTS bool
BaseURL string
MaxRequestBody int64
RequestTimeout time.Duration
ShutdownGrace time.Duration
SentryDSN string
SentryEnvironment string
// MetricsUser and MetricsPassword gate /metrics. Both set or
// neither: half-set is refused rather than resolved, because
// either resolution is dangerous. Treating a missing password as
// empty would publish the metrics endpoint to anyone who guesses
// the username; treating a missing username as "no auth" would
// publish it to everyone, in a deployment whose operator plainly
// intended it to be closed.
MetricsUser string
MetricsPassword string
// CSRFKey is exactly 32 bytes. When CSRF_KEY is absent, a random
// key is generated at startup and a warning is logged: tokens then
// do not survive a restart, which is fine in development and not
// fine behind more than one replica. Absent is a default; present
// and malformed is a startup failure.
CSRFKey []byte
CSRFKeyEphemeral bool
}
// Params defines dependencies for Config.
type Params struct {
fx.In
}
// loader parses one environment into a Config, accumulating every
// failure instead of stopping at the first, so one restart surfaces the
// whole list.
type loader struct {
v *viper.Viper
errs []error
}
func (l *loader) fail(key, raw, why string) {
l.errs = append(l.errs, fmt.Errorf(
"%w: %s=%q is %s", ErrInvalidConfig, key, raw, why,
))
}
// raw returns the trimmed value of key, and whether it was set to
// anything. Whitespace-only counts as absent: it is what an empty
// compose-file entry produces, and no key here has a meaningful blank
// value.
func (l *loader) raw(key string) (string, bool) {
s := strings.TrimSpace(l.v.GetString(key))
return s, s != ""
}
func (l *loader) str(key, def string) string {
if s, ok := l.raw(key); ok {
return s
}
return def
}
func (l *loader) integer(key string, def, minVal, maxVal int64) int64 {
s, ok := l.raw(key)
if !ok {
return def
}
n, err := strconv.ParseInt(s, 10, 64)
if err != nil {
l.fail(key, s, "not an integer")
return def
}
if n < minVal || n > maxVal {
l.fail(key, s, fmt.Sprintf("outside the range %d..%d", minVal, maxVal))
return def
}
return n
}
// boolean accepts what strconv.ParseBool accepts (1/t/T/TRUE/true/True
// and the false equivalents) and refuses everything else. "yes" is a
// parse failure on purpose: guessing at it is how a security header
// ends up off in production.
func (l *loader) boolean(key string, def bool) bool {
s, ok := l.raw(key)
if !ok {
return def
}
b, err := strconv.ParseBool(s)
if err != nil {
l.fail(key, s, "not a boolean (use true or false)")
return def
}
return b
}
func (l *loader) duration(key string, def time.Duration) time.Duration {
s, ok := l.raw(key)
if !ok {
return def
}
d, err := time.ParseDuration(s)
if err != nil {
l.fail(key, s, "not a duration (e.g. 30s, 2m)")
return def
}
if d < minTimeout || d > maxTimeout {
l.fail(key, s, fmt.Sprintf("outside the range %s..%s", minTimeout, maxTimeout))
return def
}
return d
}
// bytesize accepts both a plain integer and a human size ("1MiB",
// "512kB"), which is the form an operator actually writes.
func (l *loader) bytesize(key string, def, minVal, maxVal int64) int64 {
s, ok := l.raw(key)
if !ok {
return def
}
n, err := humanize.ParseBytes(s)
if err != nil {
l.fail(key, s, "not a byte size (e.g. 1048576, 1MiB, 512kB)")
return def
}
// maxVal and minVal are compile-time constants of this package,
// both positive, so these conversions cannot overflow; n is
// range-checked before it is narrowed.
if n > uint64(maxVal) { //nolint:gosec // see above
l.fail(key, s, byteRangeMessage(minVal, maxVal))
return def
}
sz := int64(n) //nolint:gosec // n was just checked against maxVal, a positive int64.
if sz < minVal {
l.fail(key, s, byteRangeMessage(minVal, maxVal))
return def
}
return sz
}
// byteRangeMessage renders the permitted size range the way an operator
// wrote the value they got wrong.
func byteRangeMessage(minVal, maxVal int64) string {
//nolint:gosec // both are positive compile-time constants of this package.
return fmt.Sprintf("outside the range %s..%s",
humanize.IBytes(uint64(minVal)), humanize.IBytes(uint64(maxVal)))
}
// New parses and validates the environment. Returning an error here
// aborts fx startup before anything listens, which is the whole point:
// there is no partially configured running state to reason about.
//
//nolint:revive // lc parameter is required by fx even if unused.
func New(lc fx.Lifecycle, _ Params) (*Config, error) {
v := viper.New()
v.AutomaticEnv()
return load(v)
}
// load is New's body against an explicit viper instance, so tests can
// drive it with a known environment instead of mutating the process's.
func load(v *viper.Viper) (*Config, error) {
l := &loader{v: v}
c := &Config{}
c.Port = int(l.integer(EnvPort, DefaultPort, minPort, maxPort))
c.DataDir = l.str(EnvDataDir, DefaultDataDir)
c.Debug = l.boolean(EnvDebug, false)
c.BaseURL = l.str(EnvBaseURL, DefaultBaseURL)
// HSTS defaults to on unless debugging: pinning a developer's
// browser to HTTPS on localhost is a self-inflicted outage that
// outlives the process.
c.HSTS = l.boolean(EnvHSTS, !c.Debug)
c.DBPath = l.str(EnvDBPath, filepath.Join(c.DataDir, "simplexcalc.db"))
c.MaxRequestBody = l.bytesize(
EnvMaxRequestBody, DefaultMaxRequestBody, minRequestBody, maxRequestBody,
)
c.RequestTimeout = l.duration(EnvRequestTimeout, DefaultRequestTimeout)
c.ShutdownGrace = l.duration(EnvShutdownGrace, DefaultShutdownGrace)
c.SentryDSN = l.str(EnvSentryDSN, "")
c.SentryEnvironment = l.str(EnvSentryEnv, DefaultSentryEnv)
l.checkSentryDSN(c.SentryDSN)
c.MetricsUser = l.str(EnvMetricsUser, "")
c.MetricsPassword = l.str(EnvMetricsPassword, "")
l.checkMetricsAuth(c)
l.loadCSRFKey(c)
if len(l.errs) > 0 {
return nil, errors.Join(l.errs...)
}
return c, nil
}
// checkSentryDSN refuses a DSN that is present and not a URL. An empty
// DSN disables Sentry and is not an error; a typo'd one that silently
// disabled it would be, since the operator would believe errors were
// being reported.
func (l *loader) checkSentryDSN(dsn string) {
if dsn == "" {
return
}
u, err := url.Parse(dsn)
if err != nil || u.Scheme == "" || u.Host == "" {
// The DSN embeds a key; report the failure without it.
l.errs = append(l.errs, fmt.Errorf(
"%w: %s is set but is not a valid DSN URL", ErrInvalidConfig, EnvSentryDSN,
))
}
}
// checkMetricsAuth refuses a half-configured metrics credential. See
// the field comment on Config.MetricsUser for why neither resolution
// is acceptable.
func (l *loader) checkMetricsAuth(c *Config) {
switch {
case c.MetricsUser == "" && c.MetricsPassword == "":
return
case c.MetricsUser == "":
l.errs = append(l.errs, fmt.Errorf(
"%w: %s is set but %s is not; set both or neither",
ErrInvalidConfig, EnvMetricsPassword, EnvMetricsUser,
))
case c.MetricsPassword == "":
l.errs = append(l.errs, fmt.Errorf(
"%w: %s is set but %s is not; set both or neither",
ErrInvalidConfig, EnvMetricsUser, EnvMetricsPassword,
))
}
}
// loadCSRFKey decodes CSRF_KEY, or marks the config for an ephemeral
// key. Generating the random key is deferred to the server, so that
// this function stays pure and testable.
func (l *loader) loadCSRFKey(c *Config) {
s, ok := l.raw(EnvCSRFKey)
if !ok {
c.CSRFKeyEphemeral = true
return
}
key, err := decodeHex(s)
if err != nil {
// The value is a secret: say what is wrong with it, never
// quote it.
l.errs = append(l.errs, fmt.Errorf(
"%w: %s is set but is not %d hex characters",
ErrInvalidConfig, EnvCSRFKey, csrfKeyHexChars,
))
return
}
c.CSRFKey = key
}
+257
View File
@@ -0,0 +1,257 @@
package config_test
import (
"errors"
"strings"
"testing"
"time"
"github.com/spf13/viper"
"sneak.berlin/go/simplexcalc/internal/config"
)
// Credentials used by the metrics-auth cases.
const (
testUser = "scraper"
testPass = "hunter2"
)
// env builds a viper instance holding exactly the given keys, so a test
// describes one environment without touching the process's.
func env(kv map[string]string) *viper.Viper {
v := viper.New()
for k, val := range kv {
v.Set(k, val)
}
return v
}
// TestAbsentValuesTakeDefaults pins the other half of the iron rule: a
// value that is not set does get the default. Without this, a bug that
// rejected everything would pass every test below.
func TestAbsentValuesTakeDefaults(t *testing.T) {
t.Parallel()
c, err := config.Load(env(nil))
if err != nil {
t.Fatalf("empty environment must be valid, got: %v", err)
}
if c.Port != int(config.DefaultPort) {
t.Errorf("Port = %d, want %d", c.Port, config.DefaultPort)
}
if c.MaxRequestBody != config.DefaultMaxRequestBody {
t.Errorf("MaxRequestBody = %d, want %d",
c.MaxRequestBody, config.DefaultMaxRequestBody)
}
if c.RequestTimeout != config.DefaultRequestTimeout {
t.Errorf("RequestTimeout = %s, want %s",
c.RequestTimeout, config.DefaultRequestTimeout)
}
if !c.HSTS {
t.Error("HSTS must default on when DEBUG is not set")
}
if !c.CSRFKeyEphemeral {
t.Error("an absent CSRF_KEY must mark the config for an ephemeral key")
}
}
// TestSetButUnparseableAborts is the central contract of this package.
// Every case is a value an operator plausibly types, and every one of
// them must fail startup rather than be replaced by the default.
func TestSetButUnparseableAborts(t *testing.T) {
t.Parallel()
cases := map[string]map[string]string{
"port is not a number": {config.EnvPort: "eighty"},
"port is zero": {config.EnvPort: "0"},
"port is above the range": {config.EnvPort: "70000"},
"port is a float": {config.EnvPort: "8080.0"},
"debug is yes": {config.EnvDebug: "yes"},
"hsts is on": {config.EnvHSTS: "on"},
"body cap is nonsense": {config.EnvMaxRequestBody: "big"},
"body cap is too large": {config.EnvMaxRequestBody: "1TiB"},
"body cap is too small": {config.EnvMaxRequestBody: "10"},
"timeout has no unit": {config.EnvRequestTimeout: "30"},
"timeout is out of range": {config.EnvRequestTimeout: "1h"},
"grace is nonsense": {config.EnvShutdownGrace: "soon"},
"sentry dsn is not a url": {config.EnvSentryDSN: "not a dsn"},
"csrf key is not hex": {config.EnvCSRFKey: "not-hex-at-all"},
"csrf key is wrong length": {config.EnvCSRFKey: "abcdef"},
}
for name, kv := range cases {
t.Run(name, func(t *testing.T) {
t.Parallel()
c, err := config.Load(env(kv))
if err == nil {
t.Fatalf("wanted a startup failure, got a Config: %+v", c)
}
if !errors.Is(err, config.ErrInvalidConfig) {
t.Errorf("error does not wrap ErrInvalidConfig: %v", err)
}
if c != nil {
t.Error("a failed load must return no Config at all")
}
})
}
}
// TestSecretsAreNotEchoed: a rejected CSRF key must not appear in the
// error, because errors are logged and a log is not a place to put a
// key.
func TestSecretsAreNotEchoed(t *testing.T) {
t.Parallel()
const secret = "00112233445566778899aabbccdd" // valid hex, wrong length
_, err := config.Load(env(map[string]string{config.EnvCSRFKey: secret}))
if err == nil {
t.Fatal("wanted a failure for a short CSRF key")
}
if strings.Contains(err.Error(), secret) {
t.Errorf("the rejected key was echoed in the error: %v", err)
}
}
// TestHalfSetMetricsAuthAborts covers the case the issue calls out
// explicitly: auth config that is half-set must fail loudly, in both
// directions.
func TestHalfSetMetricsAuthAborts(t *testing.T) {
t.Parallel()
cases := map[string]map[string]string{
"user without password": {config.EnvMetricsUser: testUser},
"password without user": {config.EnvMetricsPassword: testPass},
}
for name, kv := range cases {
t.Run(name, func(t *testing.T) {
t.Parallel()
_, err := config.Load(env(kv))
if err == nil {
t.Fatal("half-set metrics credentials must abort startup")
}
if !errors.Is(err, config.ErrInvalidConfig) {
t.Errorf("error does not wrap ErrInvalidConfig: %v", err)
}
})
}
both, err := config.Load(env(map[string]string{
config.EnvMetricsUser: testUser, config.EnvMetricsPassword: testPass,
}))
if err != nil {
t.Fatalf("both credentials set must be valid, got: %v", err)
}
if both.MetricsUser != testUser || both.MetricsPassword != testPass {
t.Error("credentials did not survive parsing")
}
neither, err := config.Load(env(nil))
if err != nil {
t.Fatalf("neither credential set must be valid, got: %v", err)
}
if neither.MetricsUser != "" || neither.MetricsPassword != "" {
t.Error("credentials appeared from nowhere")
}
}
// TestEveryFailureIsReported: one restart should surface the whole list,
// not just the first problem.
func TestEveryFailureIsReported(t *testing.T) {
t.Parallel()
_, err := config.Load(env(map[string]string{
config.EnvPort: "eighty",
config.EnvDebug: "yes",
config.EnvRequestTimeout: "soon",
}))
if err == nil {
t.Fatal("wanted failures")
}
for _, key := range []string{
config.EnvPort, config.EnvDebug, config.EnvRequestTimeout,
} {
if !strings.Contains(err.Error(), key) {
t.Errorf("%s is broken but is not named in the error: %v", key, err)
}
}
}
// TestValidValuesAreUsed proves the parsers accept what they document.
func TestValidValuesAreUsed(t *testing.T) {
t.Parallel()
const key = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"
c, err := config.Load(env(map[string]string{
config.EnvPort: "9000",
config.EnvDebug: "true",
config.EnvMaxRequestBody: "2MiB",
config.EnvRequestTimeout: "45s",
config.EnvShutdownGrace: "5s",
config.EnvDataDir: "/var/lib/example",
config.EnvCSRFKey: key,
}))
if err != nil {
t.Fatalf("valid environment was rejected: %v", err)
}
if c.Port != 9000 {
t.Errorf("Port = %d, want 9000", c.Port)
}
if c.MaxRequestBody != 2<<20 {
t.Errorf("MaxRequestBody = %d, want %d", c.MaxRequestBody, 2<<20)
}
if c.RequestTimeout != 45*time.Second {
t.Errorf("RequestTimeout = %s, want 45s", c.RequestTimeout)
}
if c.HSTS {
t.Error("HSTS must default off when DEBUG is true")
}
if len(c.CSRFKey) != config.CSRFKeyBytes || c.CSRFKeyEphemeral {
t.Errorf("CSRFKey not decoded: len=%d ephemeral=%v",
len(c.CSRFKey), c.CSRFKeyEphemeral)
}
if c.DBPath != "/var/lib/example/simplexcalc.db" {
t.Errorf("DBPath = %q, want it derived from DATA_DIR", c.DBPath)
}
}
// TestExplicitOverridesDerivedDBPath: DB_PATH wins over the DATA_DIR
// derivation, which is the only reason it exists.
func TestExplicitOverridesDerivedDBPath(t *testing.T) {
t.Parallel()
c, err := config.Load(env(map[string]string{
config.EnvDataDir: "/var/lib/example",
config.EnvDBPath: "/srv/other.db",
}))
if err != nil {
t.Fatalf("valid environment was rejected: %v", err)
}
if c.DBPath != "/srv/other.db" {
t.Errorf("DBPath = %q, want /srv/other.db", c.DBPath)
}
}
+17
View File
@@ -0,0 +1,17 @@
package config
// Load is load, exported for the external test package.
//
// The tests live in config_test rather than config so that they drive
// this package the way the rest of the program does — through its
// exported surface — and cannot quietly depend on an internal detail.
// The one thing they need that is not exported is the ability to
// supply an environment instead of reading the process's, which is a
// testing seam and not API.
//
//nolint:gochecknoglobals // a test seam, not mutable state.
var Load = load
// CSRFKeyBytes is the required key length, so the tests can assert on
// it without restating the number.
const CSRFKeyBytes = csrfKeyBytes
+29
View File
@@ -0,0 +1,29 @@
package config
import (
"encoding/hex"
"errors"
"fmt"
)
// errKeyLength is returned for a well-formed hex string of the wrong
// length, so decodeHex has one error type for both ways of being wrong.
var errKeyLength = errors.New("wrong key length")
// decodeHex decodes exactly csrfKeyBytes bytes of hex. It exists as its
// own function so that the length rule and the encoding rule are
// enforced in one place, and so that the caller never has to decide
// what a short-but-valid key means.
func decodeHex(s string) ([]byte, error) {
b, err := hex.DecodeString(s)
if err != nil {
return nil, fmt.Errorf("decoding hex: %w", err)
}
if len(b) != csrfKeyBytes {
return nil, fmt.Errorf("%w: got %d bytes, want %d",
errKeyLength, len(b), csrfKeyBytes)
}
return b, nil
}
+211
View File
@@ -0,0 +1,211 @@
// Package database owns the sqlite connection and the schema. The
// schema is embedded in the binary, so a deployment is one file: there
// is no migrations directory to ship alongside it and no version of it
// that can be out of step with the code that expects it.
package database
import (
"context"
"database/sql"
"embed"
"fmt"
"log/slog"
"os"
"path/filepath"
"time"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/logger"
// modernc.org/sqlite is the pure-Go driver: no cgo, so the binary
// links statically and the container needs no libc.
_ "modernc.org/sqlite"
)
// schemaFS carries the migrations into the binary.
//
//go:embed schema/*.sql
var schemaFS embed.FS
// dirPerm is the mode for the data directory: owner-only, because it
// holds the database.
const dirPerm = 0o700
// pragmas are applied to every connection. WAL is what makes concurrent
// reads not block on a write; busy_timeout is what turns the remaining
// contention into a short wait rather than an immediate SQLITE_BUSY;
// foreign_keys is off by default in sqlite and has to be asked for.
const pragmas = `
PRAGMA journal_mode = WAL;
PRAGMA busy_timeout = 5000;
PRAGMA foreign_keys = ON;
PRAGMA synchronous = NORMAL;
`
// Params defines dependencies for Database.
type Params struct {
fx.In
Config *config.Config
Logger *logger.Logger
}
// Database is the handle to the application's sqlite database.
type Database struct {
db *sql.DB
log *slog.Logger
}
// New opens the database, applies the embedded migrations, and
// registers a close hook. Migrations run during OnStart rather than
// lazily on first use: a schema that cannot be applied is a failure to
// start, and the process says so before it accepts a request.
func New(lc fx.Lifecycle, params Params) (*Database, error) {
d := &Database{log: params.Logger.Get()}
err := os.MkdirAll(filepath.Dir(params.Config.DBPath), dirPerm)
if err != nil {
return nil, fmt.Errorf("creating data directory: %w", err)
}
// New runs during graph construction, which has no request or
// lifecycle context of its own; the pragmas are a handful of
// in-process statements against a file that was just created.
db, err := Open(context.Background(), params.Config.DBPath)
if err != nil {
return nil, err
}
d.db = db
lc.Append(fx.Hook{
OnStart: func(ctx context.Context) error {
return d.Migrate(ctx)
},
OnStop: func(_ context.Context) error {
d.log.Info("closing database")
closeErr := d.db.Close()
if closeErr != nil {
return fmt.Errorf("closing database: %w", closeErr)
}
return nil
},
})
return d, nil
}
// Open opens a sqlite database at path and applies the connection
// pragmas. Exported so tests can open a scratch database without the
// fx graph.
func Open(ctx context.Context, path string) (*sql.DB, error) {
db, err := sql.Open("sqlite", path)
if err != nil {
return nil, fmt.Errorf("opening database %s: %w", path, err)
}
// sqlite tolerates exactly one writer. Holding the pool to a
// single connection makes that limit explicit here rather than
// intermittent under load, and WAL keeps readers off the writer's
// back anyway.
db.SetMaxOpenConns(1)
db.SetConnMaxLifetime(time.Hour)
_, err = db.ExecContext(ctx, pragmas)
if err != nil {
_ = db.Close()
return nil, fmt.Errorf("applying pragmas: %w", err)
}
return db, nil
}
// NewForTest opens a scratch database in dir and migrates it. Test
// helper, exported so that a project seeded from this template can use
// it from any package's tests.
func NewForTest(ctx context.Context, dir string) (*Database, error) {
d := &Database{log: slog.New(slog.DiscardHandler)}
db, err := Open(ctx, filepath.Join(dir, "test.db"))
if err != nil {
return nil, err
}
d.db = db
err = d.Migrate(ctx)
if err != nil {
_ = db.Close()
return nil, err
}
return d, nil
}
// Migrate applies every embedded migration that has not been applied to
// this database yet.
func (d *Database) Migrate(ctx context.Context) error {
set := migrationSet{fsys: schemaFS, dir: "schema"}
err := set.apply(ctx, d.db, d.log)
if err != nil {
return fmt.Errorf("applying migrations: %w", err)
}
return nil
}
// DB exposes the underlying handle for packages that need to query it.
func (d *Database) DB() *sql.DB {
return d.db
}
// AppliedVersions returns the migration versions recorded as applied,
// ascending. The healthcheck reports the highest of them, so an
// operator can see which schema a running instance is on without
// shelling into it.
func (d *Database) AppliedVersions(ctx context.Context) ([]int, error) {
rows, err := d.db.QueryContext(ctx,
"SELECT version FROM schema_migrations ORDER BY version",
)
if err != nil {
return nil, fmt.Errorf("reading applied migrations: %w", err)
}
defer func() { _ = rows.Close() }()
var versions []int
for rows.Next() {
var v int
scanErr := rows.Scan(&v)
if scanErr != nil {
return nil, fmt.Errorf("scanning migration version: %w", scanErr)
}
versions = append(versions, v)
}
err = rows.Err()
if err != nil {
return nil, fmt.Errorf("iterating applied migrations: %w", err)
}
return versions, nil
}
// Close releases the handle. Production uses the fx OnStop hook; tests
// call this.
func (d *Database) Close() error {
err := d.db.Close()
if err != nil {
return fmt.Errorf("closing database: %w", err)
}
return nil
}
+219
View File
@@ -0,0 +1,219 @@
package database_test
import (
"context"
"errors"
"path/filepath"
"testing"
"sneak.berlin/go/simplexcalc/internal/database"
)
// open returns a migrated scratch database in a directory the test
// framework removes afterwards.
func open(t *testing.T) *database.Database {
t.Helper()
db, err := database.NewForTest(t.Context(), t.TempDir())
if err != nil {
t.Fatalf("opening test database: %v", err)
}
t.Cleanup(func() {
closeErr := db.Close()
if closeErr != nil {
t.Errorf("closing test database: %v", closeErr)
}
})
return db
}
// TestMigrationsApplyFromClean is the claim the healthcheck and the
// container both rest on: an empty directory becomes a usable schema
// with no operator step in between.
func TestMigrationsApplyFromClean(t *testing.T) {
t.Parallel()
db := open(t)
versions, err := db.AppliedVersions(t.Context())
if err != nil {
t.Fatalf("reading applied versions: %v", err)
}
// 000 (the ledger) and 001 (widgets), which is every file the
// schema directory currently embeds.
if len(versions) != 2 || versions[0] != 0 || versions[1] != 1 {
t.Fatalf("applied versions = %v, want [0 1]", versions)
}
}
// TestMigrationsAreIdempotent: a restart re-runs Migrate against a
// database that already has the schema, and must change nothing. A
// migration runner that fails here takes the service down on every
// second start.
func TestMigrationsAreIdempotent(t *testing.T) {
t.Parallel()
dir := t.TempDir()
ctx := t.Context()
first, err := database.NewForTest(ctx, dir)
if err != nil {
t.Fatalf("first open: %v", err)
}
_, err = first.CreateWidget(ctx, "survivor", 1)
if err != nil {
t.Fatalf("creating widget: %v", err)
}
err = first.Close()
if err != nil {
t.Fatalf("closing: %v", err)
}
second, err := database.NewForTest(ctx, dir)
if err != nil {
t.Fatalf("reopening and re-migrating: %v", err)
}
defer func() { _ = second.Close() }()
versions, err := second.AppliedVersions(ctx)
if err != nil {
t.Fatalf("reading applied versions: %v", err)
}
if len(versions) != 2 {
t.Errorf("re-running migrations changed the ledger: %v", versions)
}
// The data has to still be there: a migration runner that "fixes"
// an already-migrated database by recreating tables is worse than
// one that fails.
count, err := second.CountWidgets(ctx)
if err != nil {
t.Fatalf("counting: %v", err)
}
if count != 1 {
t.Errorf("widget count = %d after reopen, want 1", count)
}
}
// TestWidgetRoundTrip exercises the query layer against the real
// schema, including the timestamp format shared between Go and the SQL
// DEFAULT.
func TestWidgetRoundTrip(t *testing.T) {
t.Parallel()
db := open(t)
ctx := t.Context()
created, err := db.CreateWidget(ctx, "widget one", 4096)
if err != nil {
t.Fatalf("creating widget: %v", err)
}
if created.ID == "" {
t.Error("created widget has no id")
}
widgets, err := db.ListWidgets(ctx, 10)
if err != nil {
t.Fatalf("listing widgets: %v", err)
}
if len(widgets) != 1 {
t.Fatalf("listed %d widgets, want 1", len(widgets))
}
got := widgets[0]
if got.ID != created.ID || got.Name != "widget one" || got.SizeBytes != 4096 {
t.Errorf("round trip lost data: %+v", got)
}
if got.CreatedAt.IsZero() {
t.Error("created_at did not survive the round trip")
}
}
// TestListWidgetsRespectsLimit: the index query is bounded, and the
// bound has to actually bind.
func TestListWidgetsRespectsLimit(t *testing.T) {
t.Parallel()
db := open(t)
ctx := t.Context()
for range 5 {
_, err := db.CreateWidget(ctx, "w", 1)
if err != nil {
t.Fatalf("creating widget: %v", err)
}
}
widgets, err := db.ListWidgets(ctx, 2)
if err != nil {
t.Fatalf("listing widgets: %v", err)
}
if len(widgets) != 2 {
t.Errorf("limit 2 returned %d rows", len(widgets))
}
}
// TestParseMigrationVersion covers the naming contract the schema
// directory has to keep. A file this rejects is a file that would
// otherwise be silently skipped.
func TestParseMigrationVersion(t *testing.T) {
t.Parallel()
good := map[string]int{
"000.sql": 0,
"001_widgets.sql": 1,
"017_thing.sql": 17,
}
for name, want := range good {
got, err := database.ParseMigrationVersion(name)
if err != nil {
t.Errorf("%s: unexpected error %v", name, err)
continue
}
if got != want {
t.Errorf("%s: version = %d, want %d", name, got, want)
}
}
for _, name := range []string{"widgets.sql", "_001.sql", "v1_widgets.sql"} {
_, err := database.ParseMigrationVersion(name)
if err == nil {
t.Errorf("%s: wanted a rejection, got none", name)
}
}
}
// TestOpenCreatesFile: Open must produce a database at the path it was
// given, not somewhere else.
func TestOpenCreatesFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
db, err := database.Open(t.Context(), filepath.Join(dir, "explicit.db"))
if err != nil {
t.Fatalf("opening: %v", err)
}
defer func() { _ = db.Close() }()
err = db.PingContext(t.Context())
if err != nil && !errors.Is(err, context.Canceled) {
t.Errorf("pinging the opened database: %v", err)
}
}
+199
View File
@@ -0,0 +1,199 @@
package database
import (
"context"
"database/sql"
"errors"
"fmt"
"io/fs"
"log/slog"
"path"
"sort"
"strconv"
"strings"
)
// bootstrapVersion is 000.sql: the migration that creates the ledger
// the others are recorded in.
const bootstrapVersion = 0
// errBadMigrationName is returned for a schema file whose name does not
// start with a version number. It is a build-time mistake, not a
// runtime condition, and it fails startup rather than being skipped —
// a migration silently not applied is the failure mode this whole
// mechanism exists to prevent.
var errBadMigrationName = errors.New(
"migration filename does not start with a version number",
)
// ParseMigrationVersion extracts the leading integer from a migration
// filename: "001_widgets.sql" is version 1. Exported so that a project
// seeded from this template can validate its own schema directory in a
// test.
func ParseMigrationVersion(name string) (int, error) {
base := name
if i := strings.IndexAny(base, "_."); i > 0 {
base = base[:i]
}
version, err := strconv.Atoi(base)
if err != nil {
return 0, fmt.Errorf("%w: %q", errBadMigrationName, name)
}
return version, nil
}
// migrationSet is one embedded directory of numbered .sql migrations
// (000 bootstrap plus schema files).
type migrationSet struct {
fsys fs.FS
dir string
}
// collect returns the set's migration filenames sorted
// lexicographically, which is why they are zero-padded.
func (m migrationSet) collect() ([]string, error) {
entries, err := fs.ReadDir(m.fsys, m.dir)
if err != nil {
return nil, fmt.Errorf("failed to read schema directory: %w", err)
}
var migrations []string
for _, entry := range entries {
if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".sql") {
migrations = append(migrations, entry.Name())
}
}
sort.Strings(migrations)
return migrations, nil
}
// bootstrap ensures the schema_migrations table exists by applying
// 000.sql if the table is missing.
func (m migrationSet) bootstrap(
ctx context.Context, db *sql.DB, log *slog.Logger,
) error {
var tableExists int
err := db.QueryRowContext(ctx,
"SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='schema_migrations'",
).Scan(&tableExists)
if err != nil {
return fmt.Errorf("failed to check for migrations table: %w", err)
}
if tableExists > 0 {
return nil
}
content, err := fs.ReadFile(m.fsys, path.Join(m.dir, "000.sql"))
if err != nil {
return fmt.Errorf("failed to read bootstrap migration 000.sql: %w", err)
}
if log != nil {
log.Info("applying bootstrap migration", "version", bootstrapVersion)
}
_, err = db.ExecContext(ctx, string(content))
if err != nil {
return fmt.Errorf("failed to apply bootstrap migration: %w", err)
}
return nil
}
// applied reports whether the numbered migration has been recorded.
func (m migrationSet) applied(
ctx context.Context, db *sql.DB, version int,
) (bool, error) {
var count int
err := db.QueryRowContext(ctx,
"SELECT COUNT(*) FROM schema_migrations WHERE version = ?",
version,
).Scan(&count)
if err != nil {
return false, fmt.Errorf("failed to check migration status: %w", err)
}
return count > 0, nil
}
// applyOne reads, executes, and records one migration file.
func (m migrationSet) applyOne(
ctx context.Context, db *sql.DB, migration string, version int,
) error {
content, err := fs.ReadFile(m.fsys, path.Join(m.dir, migration))
if err != nil {
return fmt.Errorf("failed to read migration %s: %w", migration, err)
}
_, execErr := db.ExecContext(ctx, string(content))
if execErr != nil {
return fmt.Errorf("failed to apply migration %s: %w", migration, execErr)
}
_, recErr := db.ExecContext(ctx,
"INSERT INTO schema_migrations (version) VALUES (?)",
version,
)
if recErr != nil {
return fmt.Errorf("failed to record migration %s: %w", migration, recErr)
}
return nil
}
// apply runs all pending migrations of the set, in order. Idempotent:
// a second run over the same database applies nothing.
func (m migrationSet) apply(ctx context.Context, db *sql.DB, log *slog.Logger) error {
err := m.bootstrap(ctx, db, log)
if err != nil {
return err
}
migrations, err := m.collect()
if err != nil {
return err
}
for _, migration := range migrations {
version, parseErr := ParseMigrationVersion(migration)
if parseErr != nil {
return parseErr
}
done, checkErr := m.applied(ctx, db, version)
if checkErr != nil {
return checkErr
}
if done {
if log != nil {
log.Debug("migration already applied", "version", version)
}
continue
}
if log != nil {
log.Info("applying migration", "version", version)
}
applyErr := m.applyOne(ctx, db, migration, version)
if applyErr != nil {
return applyErr
}
if log != nil {
log.Info("migration applied successfully", "version", version)
}
}
return nil
}
+106
View File
@@ -0,0 +1,106 @@
package database
import (
"context"
"fmt"
"time"
// Go has no UUID in the standard library as of go1.25 — checked
// against this repo's toolchain, not assumed. Swap this import for
// the stdlib package the moment one lands; nothing else here
// depends on the implementation.
"github.com/google/uuid"
)
// timeFormat matches the strftime pattern the schema uses for its
// defaults, so rows written by Go and rows written by a DEFAULT sort
// against each other correctly.
const timeFormat = "2006-01-02T15:04:05.000Z"
// Widget is the example row type. It exists so that the migration
// runner, the query layer, the templates and the tests all exercise
// real data. Delete it when seeding a real project.
type Widget struct {
ID string
Name string
SizeBytes int64
CreatedAt time.Time
}
// CreateWidget inserts a widget and returns it as stored.
func (d *Database) CreateWidget(
ctx context.Context, name string, size int64,
) (*Widget, error) {
w := &Widget{
ID: uuid.NewString(),
Name: name,
SizeBytes: size,
CreatedAt: time.Now().UTC(),
}
_, err := d.db.ExecContext(ctx,
`INSERT INTO widgets (id, name, size_bytes, created_at) VALUES (?, ?, ?, ?)`,
w.ID, w.Name, w.SizeBytes, w.CreatedAt.Format(timeFormat),
)
if err != nil {
return nil, fmt.Errorf("inserting widget: %w", err)
}
return w, nil
}
// ListWidgets returns the most recently created widgets, newest first,
// up to limit.
func (d *Database) ListWidgets(ctx context.Context, limit int) ([]Widget, error) {
rows, err := d.db.QueryContext(ctx,
`SELECT id, name, size_bytes, created_at
FROM widgets
ORDER BY created_at DESC, id DESC
LIMIT ?`,
limit,
)
if err != nil {
return nil, fmt.Errorf("listing widgets: %w", err)
}
defer func() { _ = rows.Close() }()
widgets := []Widget{}
for rows.Next() {
var (
w Widget
createdAt string
)
scanErr := rows.Scan(&w.ID, &w.Name, &w.SizeBytes, &createdAt)
if scanErr != nil {
return nil, fmt.Errorf("scanning widget: %w", scanErr)
}
w.CreatedAt, scanErr = time.Parse(timeFormat, createdAt)
if scanErr != nil {
return nil, fmt.Errorf("parsing widget created_at %q: %w", createdAt, scanErr)
}
widgets = append(widgets, w)
}
err = rows.Err()
if err != nil {
return nil, fmt.Errorf("iterating widgets: %w", err)
}
return widgets, nil
}
// CountWidgets returns the number of widgets stored.
func (d *Database) CountWidgets(ctx context.Context) (int, error) {
var n int
err := d.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM widgets`).Scan(&n)
if err != nil {
return 0, fmt.Errorf("counting widgets: %w", err)
}
return n, nil
}
+15
View File
@@ -0,0 +1,15 @@
-- 000.sql: the bootstrap migration. It creates only the ledger that
-- records which migrations have run; every other migration is recorded
-- in it. Applied when the schema_migrations table is missing, and never
-- again.
--
-- Never edit an applied migration. Add a new numbered file instead: the
-- ledger records versions, not contents, so an edited file is applied
-- nowhere and diverges everywhere.
CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
INSERT OR IGNORE INTO schema_migrations (version) VALUES (0);
+15
View File
@@ -0,0 +1,15 @@
-- 001_widgets.sql: the example table. Delete it when seeding a real
-- project and start your own schema at 001 — nothing has been deployed
-- yet, so there is no ledger anywhere that would disagree.
--
-- It is here so that the template's migration runner, model layer and
-- tests all exercise a real table rather than an empty database.
CREATE TABLE IF NOT EXISTS widgets (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
size_bytes INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
CREATE INDEX IF NOT EXISTS widgets_created_at ON widgets (created_at);
+43
View File
@@ -0,0 +1,43 @@
// Package globals provides build-time variables injected via ldflags.
package globals
import (
"runtime"
"go.uber.org/fx"
)
// Build-time variables populated from main() and copied into the
// Globals object. main() sets them from its own ldflags-injected
// values; nothing else writes them.
//
//nolint:gochecknoglobals // Build-time variables set by main().
var (
Appname string
Version string
Buildarch string
)
// Globals holds build-time metadata about the application.
type Globals struct {
Appname string
Version string
Buildarch string
}
// New creates a Globals instance from the package-level build-time
// variables.
//
//nolint:revive // lc parameter is required by fx even if unused.
func New(lc fx.Lifecycle) (*Globals, error) {
arch := Buildarch
if arch == "" {
arch = runtime.GOARCH
}
return &Globals{
Appname: Appname,
Buildarch: arch,
Version: Version,
}, nil
}
+13
View File
@@ -0,0 +1,13 @@
package handlers
import "errors"
var (
// errBadWidgetName is a rejected form value, not a fault. It exists
// so that h.fail always has a non-nil error to log: a rejection
// with no error recorded is a rejection nobody can explain later.
errBadWidgetName = errors.New("widget name is empty or too long")
// errBadWidgetSize is a size field that is not a byte count.
errBadWidgetSize = errors.New("widget size is not a byte count")
)
+123
View File
@@ -0,0 +1,123 @@
// Package handlers holds the HTTP handlers. They are methods on one
// struct whose dependencies come from fx, so a handler never reaches
// for a package-level singleton and a test can build the struct with
// exactly the collaborators it wants.
package handlers
import (
"log/slog"
"net/http"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/database"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/middleware"
"sneak.berlin/go/simplexcalc/internal/render"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// Params defines dependencies for Handlers.
type Params struct {
fx.In
Config *config.Config
Globals *globals.Globals
Logger *logger.Logger
Database *database.Database
Renderer *render.Renderer
Sentry *telemetry.Sentry
}
// Handlers is the set of HTTP handlers.
type Handlers struct {
params Params
log *slog.Logger
}
// New creates the handler set.
//
//nolint:revive // lc parameter is required by fx even if unused.
func New(lc fx.Lifecycle, params Params) (*Handlers, error) {
return &Handlers{params: params, log: params.Logger.Get()}, nil
}
// NotFound answers unmatched routes with the error page rather than
// net/http's bare text, so a 404 still carries the site's own headers
// and chrome.
func (h *Handlers) NotFound() http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
data := render.ErrorPage{
Page: h.page(r),
Status: http.StatusNotFound,
Message: "No such page.",
}
err := h.params.Renderer.HTML(w, http.StatusNotFound, "error.html", data)
if err != nil {
h.log.Error("rendering 404 failed", "error", err)
http.Error(w, "not found", http.StatusNotFound)
}
}
}
// MethodNotAllowed answers a known path with the wrong method.
func (h *Handlers) MethodNotAllowed() http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
data := render.ErrorPage{
Page: h.page(r),
Status: http.StatusMethodNotAllowed,
Message: "That method is not allowed here.",
}
err := h.params.Renderer.HTML(w, http.StatusMethodNotAllowed, "error.html", data)
if err != nil {
h.log.Error("rendering 405 failed", "error", err)
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
}
}
}
// page builds the common template data for r.
func (h *Handlers) page(r *http.Request) render.Page {
return h.params.Renderer.NewPage(middleware.CSRFField(r))
}
// fail reports a handler error and answers with the error page.
//
// The message shown to the client is chosen by the caller and is never
// the error's text: an error from the database layer carries a query,
// possibly a value out of a row, and always more about the internals
// than a stranger should be given. The error itself goes to the log and
// to Sentry, tied to the request id that is also in the response
// header, so the two can be joined afterwards.
func (h *Handlers) fail(
w http.ResponseWriter, r *http.Request, status int, message string, err error,
) {
h.log.Error("handler error",
"id", middleware.RequestIDFrom(r.Context()),
"path", r.URL.Path,
"status", status,
"error", err,
)
if status >= http.StatusInternalServerError {
h.params.Sentry.CaptureError(err)
}
data := render.ErrorPage{
Page: h.page(r),
Status: status,
Message: message,
}
renderErr := h.params.Renderer.HTML(w, status, "error.html", data)
if renderErr != nil {
// The error page itself failed. Anything further would be
// another chance to fail, so this is the floor: a plain
// status, and the reason in the log.
h.log.Error("rendering error page failed", "error", renderErr)
http.Error(w, http.StatusText(status), status)
}
}
+77
View File
@@ -0,0 +1,77 @@
package handlers
import (
"encoding/json"
"net/http"
"slices"
)
// HealthResponse is the healthcheck body. It is a typed struct rather
// than a map so that the shape is part of the code and a change to it
// shows up in a diff: something is always parsing this.
type HealthResponse struct {
OK bool `json:"ok"`
App string `json:"app"`
Version string `json:"version"`
SchemaVersion int `json:"schema_version"`
DatabaseOK bool `json:"database_ok"`
SentryEnabled bool `json:"sentry_enabled"`
MetricsProtected bool `json:"metrics_protected"`
}
// Healthcheck answers with the process's own view of whether it is
// working. It touches the database on purpose: a health endpoint that
// only proves the HTTP server is up will report healthy through the
// entire outage that matters.
//
// A failure answers 503, not 200-with-ok-false. Everything that reads
// this — a load balancer, a container runtime, a monitoring probe —
// looks at the status code first, and several look at nothing else.
func (h *Handlers) Healthcheck() http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
resp := HealthResponse{
App: h.params.Globals.Appname,
Version: h.params.Globals.Version,
SentryEnabled: h.params.Sentry.Enabled(),
MetricsProtected: h.params.Config.MetricsUser != "",
}
versions, err := h.params.Database.AppliedVersions(r.Context())
if err == nil {
resp.DatabaseOK = true
resp.OK = true
if len(versions) > 0 {
resp.SchemaVersion = slices.Max(versions)
}
} else {
h.log.Error("healthcheck: database unreachable", "error", err)
}
status := http.StatusOK
if !resp.OK {
status = http.StatusServiceUnavailable
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
// A cached healthcheck is a healthcheck that reports the past.
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(status)
encodeErr := json.NewEncoder(w).Encode(resp)
if encodeErr != nil {
// The status and headers are already sent, so there is
// nothing to answer with; the log is the only record left.
h.log.Error("healthcheck: encoding response failed", "error", encodeErr)
}
}
}
// Panic is a route that panics, mounted only when DEBUG is on. It is
// how the panic recoverer is exercised by hand in a running process;
// the automated proof is in the middleware tests.
func (h *Handlers) Panic() http.HandlerFunc {
return func(_ http.ResponseWriter, _ *http.Request) {
panic("deliberate panic from the debug route")
}
}
+114
View File
@@ -0,0 +1,114 @@
package handlers
import (
"net/http"
"strings"
"github.com/dustin/go-humanize"
"sneak.berlin/go/simplexcalc/internal/render"
)
// widgetListLimit bounds the index query. An unbounded SELECT is fine
// on the day it is written and is the outage two years later.
const widgetListLimit = 50
// maxWidgetNameLen matches the maxlength on the form input. The form is
// a courtesy; this is the rule.
const maxWidgetNameLen = 200
// Index renders the front page from the embedded template, listing the
// most recent widgets.
func (h *Handlers) Index() http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
widgets, err := h.params.Database.ListWidgets(ctx, widgetListLimit)
if err != nil {
h.fail(w, r, http.StatusInternalServerError, "Could not load widgets.", err)
return
}
count, err := h.params.Database.CountWidgets(ctx)
if err != nil {
h.fail(w, r, http.StatusInternalServerError, "Could not count widgets.", err)
return
}
data := render.IndexPage{
Page: h.page(r),
WidgetCount: count,
Widgets: widgets,
}
err = h.params.Renderer.HTML(w, http.StatusOK, "index.html", data)
if err != nil {
h.fail(w, r, http.StatusInternalServerError, "Could not render the page.", err)
}
}
}
// CreateWidget handles the form POST. State-changing, so it is behind
// CSRF; see internal/server/routes.go for where that is applied.
func (h *Handlers) CreateWidget() http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
// ParseForm reads the body, which BodyLimit has already capped:
// an oversized submission fails here rather than being buffered
// in full first.
err := r.ParseForm()
if err != nil {
h.fail(w, r, http.StatusBadRequest, "Could not read the form.", err)
return
}
name := strings.TrimSpace(r.PostFormValue("name"))
if name == "" || len(name) > maxWidgetNameLen {
h.fail(w, r, http.StatusBadRequest,
"A widget needs a name of 1 to 200 characters.", errBadWidgetName)
return
}
size, err := parseSize(r.PostFormValue("size"))
if err != nil {
h.fail(w, r, http.StatusBadRequest,
"Size must be a byte count, like 4096 or 4KiB.", err)
return
}
_, err = h.params.Database.CreateWidget(r.Context(), name, size)
if err != nil {
h.fail(w, r, http.StatusInternalServerError, "Could not save the widget.", err)
return
}
// POST/redirect/GET: a reload must not repeat the write.
http.Redirect(w, r, "/", http.StatusSeeOther)
}
}
// parseSize accepts an empty value as zero and anything else as a
// human-readable byte size.
func parseSize(s string) (int64, error) {
s = strings.TrimSpace(s)
if s == "" {
return 0, nil
}
n, err := humanize.ParseBytes(s)
if err != nil {
return 0, errBadWidgetSize
}
// A size beyond this is not a widget, it is a typo with a suffix.
const maxWidgetSize = uint64(1) << 50
if n > maxWidgetSize {
return 0, errBadWidgetSize
}
return int64(n), nil
}
+92
View File
@@ -0,0 +1,92 @@
// Package logger provides structured logging using stdlib log/slog.
//
// JSON output always, one format for every environment, TTY or not. A
// log line is a record to be queried, not prose to be read.
package logger
import (
"fmt"
"io"
"log/slog"
"os"
"path/filepath"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/globals"
)
// Params defines dependencies for Logger.
type Params struct {
fx.In
Globals *globals.Globals
// Output is the log destination. Optional in the fx graph: when
// absent (production) it defaults to os.Stdout; tests inject a
// buffer here to assert on the output contract.
Output io.Writer `optional:"true"`
}
// Logger wraps slog with application-specific functionality.
type Logger struct {
log *slog.Logger
level *slog.LevelVar
globals *globals.Globals
}
// New creates a new Logger instance.
func New(_ fx.Lifecycle, params Params) (*Logger, error) {
l := &Logger{
level: new(slog.LevelVar),
globals: params.Globals,
}
l.level.Set(slog.LevelInfo)
out := params.Output
if out == nil {
out = os.Stdout
}
// replaceAttr simplifies the source attribute to "file.go:line".
replaceAttr := func(_ []string, a slog.Attr) slog.Attr {
if a.Key == slog.SourceKey {
if src, ok := a.Value.Any().(*slog.Source); ok {
a.Value = slog.StringValue(
fmt.Sprintf("%s:%d", filepath.Base(src.File), src.Line),
)
}
}
return a
}
handler := slog.NewJSONHandler(out, &slog.HandlerOptions{
Level: l.level,
AddSource: true,
ReplaceAttr: replaceAttr,
})
l.log = slog.New(handler)
return l, nil
}
// EnableDebugLogging sets the log level to debug.
func (l *Logger) EnableDebugLogging() {
l.level.Set(slog.LevelDebug)
l.log.Debug("debug logging enabled", "debug", true)
}
// Get returns the underlying slog.Logger.
func (l *Logger) Get() *slog.Logger {
return l.log
}
// Identify logs application startup information.
func (l *Logger) Identify() {
l.log.Info("starting",
"appname", l.globals.Appname,
"version", l.globals.Version,
"arch", l.globals.Buildarch,
)
}
+46
View File
@@ -0,0 +1,46 @@
package middleware
import (
"net/http"
"strconv"
)
// BodyLimit caps how much of a request body a handler can read.
//
// http.MaxBytesReader is the mechanism, and the reason to use it rather
// than checking Content-Length is that Content-Length is a claim: a
// chunked request does not send one, and a lying one is trivial to
// send. MaxBytesReader counts the bytes that actually arrive and makes
// the read fail past the cap, so the ceiling holds whatever the headers
// said.
//
// It also sets the response's error status itself (413) when the limit
// is hit during a read, so a handler that ignores the read error still
// cannot serve a success off a truncated body.
//
// Content-Length is still checked first, as an early refusal: it costs
// nothing and it lets an oversized upload be rejected before it is
// transferred.
func (m *Middleware) BodyLimit() func(http.Handler) http.Handler {
limit := m.cfg.MaxRequestBody
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.ContentLength > limit {
w.Header().Set("Content-Length", strconv.Itoa(len(tooLargeBody)))
http.Error(w, tooLargeBody, http.StatusRequestEntityTooLarge)
return
}
r.Body = http.MaxBytesReader(w, r.Body, limit)
next.ServeHTTP(w, r)
})
}
}
// tooLargeBody is the response to an oversized request. It names no
// limit: the number is an operational detail and telling a caller
// exactly where the ceiling is only helps them sit under it.
const tooLargeBody = "request body too large"
+116
View File
@@ -0,0 +1,116 @@
package middleware
import (
"crypto/rand"
"html/template"
"net/http"
"strings"
"github.com/gorilla/csrf"
)
// csrfCookieName is deliberately not the library default: a name that
// says which service issued it makes a cookie jar readable, and two
// services on sibling hosts do not fight over one name.
const csrfCookieName = "simplexcalc_csrf"
// csrfMaxAge bounds how long a token stays valid, in seconds.
const csrfMaxAge = 12 * 60 * 60
// csrfKeyBytes is the key length gorilla/csrf requires.
const csrfKeyBytes = 32
// CSRF protects state-changing routes (POST, PUT, PATCH, DELETE). Safe
// methods pass through and are issued a token.
//
// The key comes from config: CSRF_KEY when set, otherwise a random key
// generated here and logged as such. An ephemeral key is correct for
// development and wrong for anything with more than one replica or more
// than one process lifetime, because a token issued by one key is
// rejected by another — the user sees a failed form submission, not a
// security event. That is why it is a warning at startup and a
// documented configuration key rather than a silent default.
func (m *Middleware) CSRF() func(http.Handler) http.Handler {
key := m.cfg.CSRFKey
if m.cfg.CSRFKeyEphemeral {
key = make([]byte, csrfKeyBytes)
// crypto/rand.Read cannot fail on any supported platform; it
// panics internally rather than returning an error a caller
// might ignore. A key that is not random is not a key, so
// there is nothing to fall back to here anyway.
_, _ = rand.Read(key)
m.log.Warn("CSRF_KEY is not set; using a random key for this process",
"consequence", "tokens do not survive a restart and are not shared between replicas")
}
protect := csrf.Protect(
key,
// Secure cookies require TLS, which is absent in local
// development; tying the flag to the same switch that governs
// HSTS keeps "is this a production deployment" a single
// decision rather than two that can disagree.
csrf.Secure(m.cfg.HSTS),
csrf.HttpOnly(true),
csrf.SameSite(csrf.SameSiteLaxMode),
csrf.Path("/"),
csrf.CookieName(csrfCookieName),
csrf.MaxAge(csrfMaxAge),
csrf.ErrorHandler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
m.log.Warn("csrf rejection",
"id", RequestIDFrom(r.Context()),
"path", r.URL.Path,
"reason", csrf.FailureReason(r).Error(),
)
http.Error(w, "invalid CSRF token", http.StatusForbidden)
})),
)
// markScheme must be OUTSIDE protect: it sets a context value that
// protect reads, so it has to run first.
return func(next http.Handler) http.Handler {
return markScheme(protect(next))
}
}
// markScheme tells gorilla/csrf whether the browser's connection was
// plaintext, because the library cannot tell and assumes it was not.
//
// Its strict Referer check is for TLS only, and it treats every request
// as TLS unless a context value says otherwise. A service behind a
// TLS-terminating reverse proxy receives plaintext HTTP with an
// https:// Referer — the library then applies the TLS rules to a
// plaintext connection and rejects every form submission, which is a
// total outage of every state-changing route rather than a subtle bug.
// Left alone, the same misreading rejects plain HTTP in development for
// the mirror-image reason.
//
// The rule: HTTPS if the connection is TLS, or if a proxy said so with
// X-Forwarded-Proto. Trusting that header is safe in this one
// direction — the only thing an attacker gains by setting it is
// STRICTER checking of their own request. The reverse (inferring
// plaintext) is what would weaken the check, and nothing a client sends
// can cause it.
//
// A deployment behind a proxy that does not set X-Forwarded-Proto gets
// the plaintext ruleset: tokens still work, and the extra Referer check
// TLS would have added is not applied. Configure the proxy.
func markScheme(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.TLS == nil && !strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https") {
r = csrf.PlaintextHTTPRequest(r)
}
next.ServeHTTP(w, r)
})
}
// CSRFField returns the hidden input for r's token, for a template to
// place inside a form. Handlers call this rather than importing
// gorilla/csrf, so the library stays swappable behind this package.
func CSRFField(r *http.Request) template.HTML {
return csrf.TemplateField(r)
}
+163
View File
@@ -0,0 +1,163 @@
// Package middleware holds the HTTP middleware chain: request
// identity, logging, metrics, panic recovery, timeouts, body caps,
// security headers and CSRF.
//
// Order matters and is fixed in internal/server/routes.go, not here.
package middleware
import (
"context"
"log/slog"
"net/http"
"strconv"
"time"
"github.com/go-chi/chi/v5"
"github.com/google/uuid"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// contextKey is this package's private context key type, so no other
// package can collide with or read these values by accident.
type contextKey string
// requestIDKey carries the per-request id.
const requestIDKey contextKey = "request-id"
// RequestIDHeader is the response header the id is echoed in, so a
// user reporting a failure can quote something that finds the log line.
const RequestIDHeader = "X-Request-Id"
// Params defines dependencies for Middleware.
type Params struct {
fx.In
Config *config.Config
Logger *logger.Logger
Sentry *telemetry.Sentry
Metrics *telemetry.Metrics
}
// Middleware is the set of handlers, built once and reused.
type Middleware struct {
params Params
log *slog.Logger
cfg *config.Config
}
// New creates the middleware set.
//
//nolint:revive // lc parameter is required by fx even if unused.
func New(lc fx.Lifecycle, params Params) (*Middleware, error) {
return &Middleware{
params: params,
log: params.Logger.Get(),
cfg: params.Config,
}, nil
}
// RequestIDFrom returns the id assigned to r's context, or "" outside a
// request that went through RequestID.
func RequestIDFrom(ctx context.Context) string {
id, _ := ctx.Value(requestIDKey).(string)
return id
}
// RequestID assigns each request an id and echoes it. An id supplied by
// the client is ignored: it is attacker-controlled, it would let a
// caller collide two unrelated requests in the log, and there is no
// trusted proxy contract here that would make it meaningful.
func (m *Middleware) RequestID() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
id := uuid.NewString()
w.Header().Set(RequestIDHeader, id)
next.ServeHTTP(w, r.WithContext(
context.WithValue(r.Context(), requestIDKey, id),
))
})
}
}
// RequestLogger logs one line per completed request.
func (m *Middleware) RequestLogger() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
rec := newResponseRecorder(w)
next.ServeHTTP(rec, r)
m.log.Info("request",
"id", RequestIDFrom(r.Context()),
"method", r.Method,
"path", r.URL.Path,
"route", routePattern(r),
"status", rec.Status(),
"bytes", rec.written,
"duration_ms", time.Since(start).Milliseconds(),
)
})
}
}
// Metrics records the Prometheus series for each request.
func (m *Middleware) Metrics() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
rec := newResponseRecorder(w)
m.params.Metrics.InFlightAdd(1)
defer m.params.Metrics.InFlightAdd(-1)
next.ServeHTTP(rec, r)
m.params.Metrics.Observe(
r.Method,
routePattern(r),
strconv.Itoa(rec.Status()),
time.Since(start),
)
})
}
}
// Timeout bounds handler execution with the configured request timeout.
// The handler sees a context with a deadline; a handler that ignores it
// still runs to completion, so handlers must pass the context down to
// everything that can block.
func (m *Middleware) Timeout() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx, cancel := context.WithTimeout(r.Context(), m.cfg.RequestTimeout)
defer cancel()
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// routePattern returns the chi route pattern for r, or "unmatched" when
// no route matched (a 404). It is what the metrics and the log are
// labelled by; see the comment on the requests counter for why the path
// is not.
func routePattern(r *http.Request) string {
rctx := chi.RouteContext(r.Context())
if rctx == nil {
return "unmatched"
}
pattern := rctx.RoutePattern()
if pattern == "" {
return "unmatched"
}
return pattern
}
+314
View File
@@ -0,0 +1,314 @@
package middleware_test
import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/middleware"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// newMiddleware builds the set against a given config, with logging
// discarded and telemetry disabled.
func newMiddleware(t *testing.T, cfg *config.Config) *middleware.Middleware {
t.Helper()
g := &globals.Globals{Appname: "simplexcalc", Version: "test"}
log, err := logger.New(nil, logger.Params{Globals: g, Output: io.Discard})
if err != nil {
t.Fatalf("building logger: %v", err)
}
sentry, err := telemetry.NewSentry(nil, telemetry.SentryParams{
Config: cfg, Globals: g, Logger: log,
})
if err != nil {
t.Fatalf("building sentry: %v", err)
}
metrics, err := telemetry.NewMetrics(telemetry.MetricsParams{Config: cfg})
if err != nil {
t.Fatalf("building metrics: %v", err)
}
mw, err := middleware.New(nil, middleware.Params{
Config: cfg, Logger: log, Sentry: sentry, Metrics: metrics,
})
if err != nil {
t.Fatalf("building middleware: %v", err)
}
return mw
}
// getReq and postReq build requests carrying the test's context, so a
// handler that respects cancellation is exercised the way the server
// exercises it.
func getReq(t *testing.T) *http.Request {
t.Helper()
return httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", nil)
}
func postReq(t *testing.T, body string) *http.Request {
t.Helper()
return httptest.NewRequestWithContext(
t.Context(), http.MethodPost, "/", strings.NewReader(body),
)
}
func testConfig() *config.Config {
return &config.Config{
Port: 8080,
HSTS: true,
MaxRequestBody: 1024,
RequestTimeout: time.Second,
ShutdownGrace: time.Second,
CSRFKeyEphemeral: true,
}
}
// TestRecovererAnswers500 is the point of the panic middleware: net/http
// on its own drops the connection, which tells the client nothing about
// whose fault it was.
func TestRecovererAnswers500(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
h := mw.Recoverer()(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {
panic("boom")
}))
w := httptest.NewRecorder()
h.ServeHTTP(w, getReq(t))
if w.Code != http.StatusInternalServerError {
t.Fatalf("status = %d, want 500", w.Code)
}
if w.Body.Len() == 0 {
t.Error("a 500 with no body tells the client nothing")
}
// The panic value must not reach the client.
if strings.Contains(w.Body.String(), "boom") {
t.Error("the panic value was leaked in the response body")
}
}
// TestRecovererPassesThroughSuccess: the recovery wrapper must be
// invisible when nothing goes wrong, including for the response body.
func TestRecovererPassesThroughSuccess(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
h := mw.Recoverer()(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusTeapot)
_, _ = w.Write([]byte("fine"))
}))
w := httptest.NewRecorder()
h.ServeHTTP(w, getReq(t))
if w.Code != http.StatusTeapot || w.Body.String() != "fine" {
t.Errorf("status = %d body = %q", w.Code, w.Body.String())
}
}
// TestSecurityHeadersOnEveryResponse, including responses the handler
// never got to write.
func TestSecurityHeadersOnEveryResponse(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
notFound := func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "not found", http.StatusNotFound)
}
h := mw.SecurityHeaders()(http.HandlerFunc(notFound))
w := httptest.NewRecorder()
h.ServeHTTP(w, getReq(t))
want := map[string]string{
"X-Frame-Options": "DENY",
"X-Content-Type-Options": "nosniff",
"Referrer-Policy": "strict-origin-when-cross-origin",
"Strict-Transport-Security": "max-age=31536000; includeSubDomains",
}
for header, value := range want {
if got := w.Header().Get(header); got != value {
t.Errorf("%s = %q, want %q", header, got, value)
}
}
csp := w.Header().Get("Content-Security-Policy")
if !strings.Contains(csp, "default-src 'self'") {
t.Errorf("CSP = %q", csp)
}
if strings.Contains(csp, "unsafe-inline") {
t.Error("the CSP permits inline script or style")
}
}
// TestHSTSOffWhenDisabled: the header must be absent, not empty, so a
// developer's browser is never pinned to HTTPS on localhost.
func TestHSTSOffWhenDisabled(t *testing.T) {
t.Parallel()
cfg := testConfig()
cfg.HSTS = false
mw := newMiddleware(t, cfg)
noop := func(_ http.ResponseWriter, _ *http.Request) {}
h := mw.SecurityHeaders()(http.HandlerFunc(noop))
w := httptest.NewRecorder()
h.ServeHTTP(w, getReq(t))
if _, ok := w.Header()["Strict-Transport-Security"]; ok {
t.Error("HSTS was sent with HSTS disabled")
}
}
// TestBodyLimitRefusesDeclaredOversize: a Content-Length over the cap is
// refused before the body transfers.
func TestBodyLimitRefusesDeclaredOversize(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
reached := false
h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {
reached = true
}))
req := postReq(t, strings.Repeat("x", 2048))
w := httptest.NewRecorder()
h.ServeHTTP(w, req)
if w.Code != http.StatusRequestEntityTooLarge {
t.Errorf("status = %d, want 413", w.Code)
}
if reached {
t.Error("the handler ran for an oversized request")
}
}
// TestBodyLimitCapsUndeclaredBody is the case Content-Length cannot
// catch: a body that arrives without one, or with a lying one, must
// still fail at the cap rather than being read in full.
func TestBodyLimitCapsUndeclaredBody(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
var readErr error
h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
_, readErr = io.ReadAll(r.Body)
}))
req := postReq(t, strings.Repeat("x", 4096))
// Undeclared length: what a chunked upload looks like here.
req.ContentLength = -1
h.ServeHTTP(httptest.NewRecorder(), req)
if readErr == nil {
t.Error("reading past the cap succeeded; the limit is not enforced on the read")
}
}
// TestBodyLimitAllowsNormalRequests, so the cap is not just "refuse
// everything".
func TestBodyLimitAllowsNormalRequests(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
got := ""
h := mw.BodyLimit()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
b, _ := io.ReadAll(r.Body)
got = string(b)
}))
req := postReq(t, "small")
h.ServeHTTP(httptest.NewRecorder(), req)
if got != "small" {
t.Errorf("body = %q, want %q", got, "small")
}
}
// TestRequestIDIsAssignedAndNotBorrowed: the id must be this process's,
// so a client cannot collide two unrelated requests in the log.
func TestRequestIDIsAssignedAndNotBorrowed(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
var inHandler string
h := mw.RequestID()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
inHandler = middleware.RequestIDFrom(r.Context())
}))
req := getReq(t)
req.Header.Set(middleware.RequestIDHeader, "client-supplied")
w := httptest.NewRecorder()
h.ServeHTTP(w, req)
if inHandler == "" {
t.Fatal("no request id reached the handler")
}
if inHandler == "client-supplied" {
t.Error("the client's request id was trusted")
}
if w.Header().Get(middleware.RequestIDHeader) != inHandler {
t.Error("the response header does not carry the id the handler saw")
}
}
// TestTimeoutGivesHandlerADeadline. The handler is what has to respect
// it, so what is asserted here is that the deadline is there at all.
func TestTimeoutGivesHandlerADeadline(t *testing.T) {
t.Parallel()
mw := newMiddleware(t, testConfig())
var hasDeadline bool
h := mw.Timeout()(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
_, hasDeadline = r.Context().Deadline()
}))
h.ServeHTTP(httptest.NewRecorder(), getReq(t))
if !hasDeadline {
t.Error("the handler's context carries no deadline")
}
}
+59
View File
@@ -0,0 +1,59 @@
package middleware
import (
"net/http"
"runtime/debug"
)
// panicBody is the entire response a recovered panic produces. No
// template, no detail: the client learns that the request failed, and
// everything about why goes to the log and to Sentry, where it is not
// attacker-readable.
const panicBody = "internal server error"
// Recoverer turns a panicking handler into a 500 rather than a dropped
// connection.
//
// net/http already recovers panics, but what it does is close the
// connection without a response, so the client sees a transport error
// and no status. Answering 500 is the difference between "the service
// is broken" and "the network is broken" for everyone downstream.
//
// A panic after the response has started cannot be turned into a 500 —
// the status is already on the wire — so in that case the connection is
// deliberately dropped by re-panicking to net/http, which is the only
// honest signal left that the body is truncated. A truncated 200 that
// looks complete is worse than a broken connection.
func (m *Middleware) Recoverer() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
rec := newResponseRecorder(w)
defer func() {
v := recover()
if v == nil {
return
}
// http.ErrAbortHandler is net/http's documented way for
// a handler to abandon a response on purpose. It is not
// a bug, so it is not reported; it is re-raised for
// net/http to handle as it always does.
//nolint:errorlint,err113 // a sentinel value, compared as net/http documents.
if v == http.ErrAbortHandler {
panic(v)
}
m.params.Sentry.CapturePanic(v, debug.Stack())
if rec.Written() {
panic(v)
}
http.Error(rec, panicBody, http.StatusInternalServerError)
}()
next.ServeHTTP(rec, r)
})
}
}
+63
View File
@@ -0,0 +1,63 @@
package middleware
import (
"net/http"
)
// responseRecorder remembers the status code and byte count for the
// logger and the metrics middleware. net/http gives no way to read
// them back off an http.ResponseWriter, so the only way to know what
// was answered is to be the thing that answered it.
type responseRecorder struct {
http.ResponseWriter
status int
written int64
wrote bool
}
func newResponseRecorder(w http.ResponseWriter) *responseRecorder {
// A handler that writes a body without calling WriteHeader has
// sent 200; recording that up front means Status() is right for
// the common case without waiting for a call that never comes.
return &responseRecorder{ResponseWriter: w, status: http.StatusOK}
}
func (r *responseRecorder) WriteHeader(status int) {
if r.wrote {
return
}
r.status = status
r.wrote = true
r.ResponseWriter.WriteHeader(status)
}
func (r *responseRecorder) Write(b []byte) (int, error) {
r.wrote = true
n, err := r.ResponseWriter.Write(b)
r.written += int64(n)
//nolint:wrapcheck // pass-through writer: wrapping would obscure the underlying error.
return n, err
}
// Status returns the status code that was sent.
func (r *responseRecorder) Status() int {
return r.status
}
// Written reports whether anything has been sent yet. The panic
// recoverer needs this: it can only substitute a 500 for a response
// that has not started.
func (r *responseRecorder) Written() bool {
return r.wrote
}
// Unwrap lets http.ResponseController reach the underlying writer, so
// wrapping does not cost the handler flushing or deadline control.
func (r *responseRecorder) Unwrap() http.ResponseWriter {
return r.ResponseWriter
}
+67
View File
@@ -0,0 +1,67 @@
package middleware
import "net/http"
// Security response headers.
const (
// hstsValue is served even where TLS terminates at a reverse
// proxy, so the browser enforces HTTPS end to end. Off when
// config.HSTS is false (development), because pinning a
// developer's browser to HTTPS on localhost is a self-inflicted
// outage that outlives the process.
hstsValue = "max-age=31536000; includeSubDomains"
// cspValue is the baseline. Every template ships with external CSS
// and no inline script, style or event handler, so nothing needs
// 'unsafe-inline' and nothing should be given it: the moment a
// project seeded from this template adds 'unsafe-inline', the
// policy stops being a defence against injected script and becomes
// decoration.
cspValue = "default-src 'self'; " +
"base-uri 'self'; " +
"form-action 'self'; " +
"frame-ancestors 'none'; " +
"object-src 'none'"
// permissionsPolicyValue denies the browser features this
// application does not use.
permissionsPolicyValue = "accelerometer=(), autoplay=(), camera=(), " +
"display-capture=(), encrypted-media=(), geolocation=(), " +
"gyroscope=(), magnetometer=(), microphone=(), midi=(), " +
"payment=(), picture-in-picture=(), " +
"publickey-credentials-get=(), screen-wake-lock=(), usb=(), " +
"xr-spatial-tracking=()"
referrerPolicyValue = "strict-origin-when-cross-origin"
frameOptionsValue = "DENY"
contentTypeOptsVal = "nosniff"
)
// SecurityHeaders sets the response security headers before the handler
// runs, so they are on every response the router produces — 404s,
// handler error bodies, static assets, and the bare 500 the panic
// recoverer writes.
//
// X-Frame-Options duplicates the CSP frame-ancestors directive on
// purpose, for browsers that do not implement the latter.
func (m *Middleware) SecurityHeaders() func(http.Handler) http.Handler {
hsts := m.cfg.HSTS
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
h := w.Header()
if hsts {
h.Set("Strict-Transport-Security", hstsValue)
}
h.Set("Content-Security-Policy", cspValue)
h.Set("X-Frame-Options", frameOptionsValue)
h.Set("X-Content-Type-Options", contentTypeOptsVal)
h.Set("Referrer-Policy", referrerPolicyValue)
h.Set("Permissions-Policy", permissionsPolicyValue)
next.ServeHTTP(w, r)
})
}
}
+13
View File
@@ -0,0 +1,13 @@
package render
import "errors"
var (
// errNoTemplates means the embed matched nothing — a build that
// produced a binary with no pages in it.
errNoTemplates = errors.New("no page templates were embedded")
// errUnknownTemplate means a handler asked for a page that is not
// in the embedded set: a typo, caught by that handler's test.
errUnknownTemplate = errors.New("unknown template")
)
+240
View File
@@ -0,0 +1,240 @@
// Package render parses the embedded templates once, at startup, and
// executes them into a buffer before writing anything to the client.
//
// Two decisions worth keeping when this is seeded into a real project:
//
// - Every template is parsed in New. A template that does not compile
// is a process that does not start, rather than a 500 the first time
// someone visits the page it broke.
// - Execution goes to a buffer first. A template that fails halfway
// through would otherwise have already written a 200 and half a
// page, and the error could no longer be reported as one.
package render
import (
"bytes"
"fmt"
"html/template"
"io"
"io/fs"
"net/http"
"path"
"sort"
"strings"
"time"
"github.com/dustin/go-humanize"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/database"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/templates"
)
// baseTemplate is the outer document every page is rendered through.
const baseTemplate = "base"
// Params defines dependencies for Renderer.
type Params struct {
fx.In
Globals *globals.Globals
}
// Renderer holds one compiled template set per page.
type Renderer struct {
pages map[string]*template.Template
globals *globals.Globals
started time.Time
}
// Page is the data every template can rely on, embedded by the
// page-specific types below so that promoted fields keep the templates
// free of a data-envelope prefix.
type Page struct {
AppName string
Version string
Buildarch string
Uptime string
// CSRFField is the hidden input gorilla/csrf validates. It is
// template.HTML because it is markup this process generated, not
// input; nothing user-supplied is ever assigned to it.
CSRFField template.HTML
}
// IndexPage is the data for index.html.
type IndexPage struct {
Page
WidgetCount int
Widgets []database.Widget
}
// ErrorPage is the data for error.html.
type ErrorPage struct {
Page
Status int
Message string
}
// funcs are the template helpers. Deliberately few: logic belongs in
// the handler, where it can be tested without parsing HTML.
func funcs() template.FuncMap {
return template.FuncMap{
// bytes renders a byte count the way an operator reads one.
"bytes": func(n int64) string {
if n < 0 {
return "-"
}
return humanize.IBytes(uint64(n))
},
// since renders a timestamp as "3 minutes ago".
"since": humanize.Time,
}
}
// New compiles every page template against the base document and the
// partials.
//
//nolint:revive // lc parameter is required by fx even if unused.
func New(lc fx.Lifecycle, params Params) (*Renderer, error) {
r := &Renderer{
pages: map[string]*template.Template{},
globals: params.Globals,
started: time.Now(),
}
shared, pages, err := split(templates.FS)
if err != nil {
return nil, err
}
for _, page := range pages {
// Each page gets its own set: pages define blocks of the same
// names ("title", "content"), so parsing them all into one
// template would leave whichever was parsed last defining both
// for everybody.
set := template.New(baseTemplate).Funcs(funcs())
set, err = set.ParseFS(templates.FS, append(append([]string{}, shared...), page)...)
if err != nil {
return nil, fmt.Errorf("parsing template %s: %w", page, err)
}
r.pages[path.Base(page)] = set
}
if len(r.pages) == 0 {
return nil, errNoTemplates
}
return r, nil
}
// split separates the embedded set into the files every page needs
// (the base document and the partials) and the page templates
// themselves.
func split(fsys fs.FS) ([]string, []string, error) {
partials, err := fs.Glob(fsys, "partials/*.html")
if err != nil {
return nil, nil, fmt.Errorf("globbing partials: %w", err)
}
top, err := fs.Glob(fsys, "*.html")
if err != nil {
return nil, nil, fmt.Errorf("globbing templates: %w", err)
}
var pages []string
shared := append([]string{}, partials...)
for _, f := range top {
if strings.TrimSuffix(path.Base(f), ".html") == baseTemplate {
shared = append(shared, f)
continue
}
pages = append(pages, f)
}
sort.Strings(shared)
sort.Strings(pages)
return shared, pages, nil
}
// NewPage returns the common data, filled in from build-time globals
// and the request's CSRF field.
func (r *Renderer) NewPage(csrfField template.HTML) Page {
return Page{
AppName: r.globals.Appname,
Version: r.globals.Version,
Buildarch: r.globals.Buildarch,
Uptime: time.Since(r.started).Round(time.Second).String(),
CSRFField: csrfField,
}
}
// Execute renders a page into w. It buffers first: see the package
// comment.
func (r *Renderer) Execute(w io.Writer, name string, data any) error {
set, ok := r.pages[name]
if !ok {
return fmt.Errorf("%w: %s", errUnknownTemplate, name)
}
var buf bytes.Buffer
err := set.ExecuteTemplate(&buf, baseTemplate, data)
if err != nil {
return fmt.Errorf("executing template %s: %w", name, err)
}
_, err = buf.WriteTo(w)
if err != nil {
return fmt.Errorf("writing rendered template %s: %w", name, err)
}
return nil
}
// HTML renders a page to an http.ResponseWriter with the given status.
// A render failure after the buffer succeeded cannot happen, so the
// status written here is always the status the client sees.
func (r *Renderer) HTML(
w http.ResponseWriter, status int, name string, data any,
) error {
var buf bytes.Buffer
err := r.Execute(&buf, name, data)
if err != nil {
return err
}
w.Header().Set("Content-Type", "text/html; charset=utf-8")
w.WriteHeader(status)
_, err = buf.WriteTo(w)
if err != nil {
return fmt.Errorf("writing response: %w", err)
}
return nil
}
// Names returns the compiled page names, sorted. Tests use it to assert
// that every embedded page really compiled.
func (r *Renderer) Names() []string {
names := make([]string, 0, len(r.pages))
for name := range r.pages {
names = append(names, name)
}
sort.Strings(names)
return names
}
+197
View File
@@ -0,0 +1,197 @@
package render_test
import (
"bytes"
"html/template"
"net/http/httptest"
"strings"
"testing"
"time"
"sneak.berlin/go/simplexcalc/internal/database"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/render"
"sneak.berlin/go/simplexcalc/templates"
)
func newRenderer(t *testing.T) *render.Renderer {
t.Helper()
r, err := render.New(nil, render.Params{
Globals: &globals.Globals{
Appname: "simplexcalc", Version: "test", Buildarch: "amd64",
},
})
if err != nil {
t.Fatalf("compiling templates: %v", err)
}
return r
}
// TestEveryEmbeddedPageCompiles is the reason New parses everything at
// startup: a template that does not compile must be a process that does
// not start, and this is what proves the set is complete rather than
// just non-empty.
func TestEveryEmbeddedPageCompiles(t *testing.T) {
t.Parallel()
r := newRenderer(t)
embedded, err := templates.FS.ReadDir(".")
if err != nil {
t.Fatalf("reading embedded templates: %v", err)
}
want := 0
for _, e := range embedded {
if !e.IsDir() && strings.HasSuffix(e.Name(), ".html") && e.Name() != "base.html" {
want++
}
}
if want == 0 {
t.Fatal("no page templates were embedded, so this test proves nothing")
}
if got := len(r.Names()); got != want {
t.Errorf("compiled %d pages (%v), embedded %d", got, r.Names(), want)
}
}
// TestIndexRendersEmbeddedContent renders the page the service serves
// at /, with real data, and checks that the base document, both
// partials and the page body all made it into one response.
func TestIndexRendersEmbeddedContent(t *testing.T) {
t.Parallel()
r := newRenderer(t)
data := render.IndexPage{
Page: r.NewPage(template.HTML(`<input type="hidden" name="csrf" />`)),
WidgetCount: 1,
Widgets: []database.Widget{
{
ID: "abc", Name: "a widget", SizeBytes: 4096,
CreatedAt: time.Now().Add(-time.Hour),
},
},
}
var buf bytes.Buffer
err := r.Execute(&buf, "index.html", data)
if err != nil {
t.Fatalf("rendering index: %v", err)
}
out := buf.String()
for _, want := range []string{
"<!doctype html>", // base document
`href="/static/css/style.css"`, // base document links the embedded asset
"<nav>", // navbar partial
"<footer>", // footer partial
"a widget", // page data
"4.0 KiB", // the bytes template func
"ago", // the since template func
`name="csrf"`, // the CSRF field was placed
} {
if !strings.Contains(out, want) {
t.Errorf("rendered page is missing %q\n---\n%s", want, out)
}
}
}
// TestHTMLEscapesUserData: the name comes from a form, and html/template
// is only a defence if the data actually goes through it.
func TestHTMLEscapesUserData(t *testing.T) {
t.Parallel()
r := newRenderer(t)
data := render.IndexPage{
Page: r.NewPage(""),
WidgetCount: 1,
Widgets: []database.Widget{
{ID: "x", Name: `<script>alert(1)</script>`, CreatedAt: time.Now()},
},
}
var buf bytes.Buffer
err := r.Execute(&buf, "index.html", data)
if err != nil {
t.Fatalf("rendering: %v", err)
}
if strings.Contains(buf.String(), "<script>alert(1)</script>") {
t.Error("a widget name was rendered as live markup")
}
}
// TestPagesDoNotShareBlocks: index.html and error.html both define
// "title" and "content". Parsed into one set, the last one parsed would
// define both for everybody, and every page would render as whichever
// that was.
func TestPagesDoNotShareBlocks(t *testing.T) {
t.Parallel()
r := newRenderer(t)
var buf bytes.Buffer
err := r.Execute(&buf, "error.html", render.ErrorPage{
Page: r.NewPage(""), Status: 404, Message: "gone fishing",
})
if err != nil {
t.Fatalf("rendering error page: %v", err)
}
out := buf.String()
if !strings.Contains(out, "gone fishing") {
t.Errorf("error page did not render its own content:\n%s", out)
}
if strings.Contains(out, "widgets (") {
t.Error("the error page rendered the index's content block")
}
}
// TestUnknownTemplateIsAnError, rather than an empty 200.
func TestUnknownTemplateIsAnError(t *testing.T) {
t.Parallel()
r := newRenderer(t)
err := r.Execute(&bytes.Buffer{}, "nope.html", nil)
if err == nil {
t.Fatal("rendering a template that does not exist must fail")
}
}
// TestHTMLWritesStatusAndType checks the response contract, since the
// handlers rely on it for every page they serve.
func TestHTMLWritesStatusAndType(t *testing.T) {
t.Parallel()
r := newRenderer(t)
w := httptest.NewRecorder()
err := r.HTML(w, 404, "error.html", render.ErrorPage{
Page: r.NewPage(""), Status: 404, Message: "no",
})
if err != nil {
t.Fatalf("rendering: %v", err)
}
if w.Code != 404 {
t.Errorf("status = %d, want 404", w.Code)
}
if ct := w.Header().Get("Content-Type"); !strings.HasPrefix(ct, "text/html") {
t.Errorf("Content-Type = %q", ct)
}
}
+55
View File
@@ -0,0 +1,55 @@
package server
import (
"log/slog"
"net/http"
"time"
)
// HTTP server hardening limits. These are the connection-level bounds;
// the per-request deadline is middleware.Timeout, driven by
// REQUEST_TIMEOUT.
const (
// readHeaderTimeout is the slowloris bound: request headers must
// arrive within it, and it starts on connection accept.
readHeaderTimeout = 5 * time.Second
// readTimeout covers headers plus body. Bodies are capped by
// MAX_REQUEST_BODY, which transfers well inside this even on a
// slow mobile link.
readTimeout = 30 * time.Second
// writeTimeout must exceed the per-request timeout: it starts when
// the headers are read, so it spans handler execution, and a
// smaller value would cut the connection instead of letting the
// request context deadline end the request with a status. The
// margin is added to whatever REQUEST_TIMEOUT is configured to.
writeTimeoutMargin = 15 * time.Second
// idleTimeout bounds how long an idle keep-alive connection is
// held; browsers reconnect transparently.
idleTimeout = 120 * time.Second
maxHeaderBytes = 1 << 20
)
// newHTTPServer builds the http.Server.
//
// ErrorLog is set on purpose: net/http internals (TLS handshake
// errors, request parse errors, panics net/http itself recovers) write
// through it, and unset they would emit plain text on stderr via the
// default log package — a few lines of unstructured output in the
// middle of a JSON log stream, which is exactly the kind of thing a log
// pipeline drops on the floor.
func (s *Server) newHTTPServer(listenAddr string) *http.Server {
return &http.Server{
Addr: listenAddr,
ReadHeaderTimeout: readHeaderTimeout,
ReadTimeout: readTimeout,
WriteTimeout: s.params.Config.RequestTimeout + writeTimeoutMargin,
IdleTimeout: idleTimeout,
MaxHeaderBytes: maxHeaderBytes,
Handler: s,
ErrorLog: slog.NewLogLogger(s.log.Handler(), slog.LevelError),
}
}
+87
View File
@@ -0,0 +1,87 @@
package server
import (
"net/http"
"github.com/go-chi/chi/v5"
"sneak.berlin/go/simplexcalc/static"
)
// staticPrefix is where the embedded assets are mounted.
const staticPrefix = "/static/"
// staticCacheControl is a year, because the asset set is baked into the
// binary: a new build is a new deployment, and a deployment is the only
// thing that can change these bytes. Bump the path if that stops being
// true.
const staticCacheControl = "public, max-age=31536000, immutable"
// SetupRoutes builds the router. The middleware order is the contract
// of this file, and it is this way for reasons:
//
// 1. RequestID first, so every later line of log and every captured
// error can name the request it came from.
// 2. Recoverer next-to-outermost, so it covers every handler and every
// middleware below it. Above the logger, so a panic still produces
// a logged request line.
// 3. RequestLogger and Metrics before the work, so both see the final
// status of every request including the 500 the recoverer wrote.
// 4. SecurityHeaders before anything can write a body — the headers
// have to be set before the first Write, and a 404 or a panic
// response needs them as much as a page does.
// 5. Timeout and BodyLimit before any handler reads a body.
// 6. CSRF innermost of the global chain, wrapping only the routes that
// can change state.
//
// /metrics and the healthcheck sit outside CSRF (neither is
// state-changing, and a scraper has no token) and outside nothing else.
func (s *Server) SetupRoutes() {
r := chi.NewRouter()
r.Use(s.mw.RequestID())
r.Use(s.mw.Recoverer())
r.Use(s.mw.RequestLogger())
r.Use(s.mw.Metrics())
r.Use(s.mw.SecurityHeaders())
r.Use(s.mw.Timeout())
r.Use(s.mw.BodyLimit())
r.NotFound(s.h.NotFound())
r.MethodNotAllowed(s.h.MethodNotAllowed())
// Operational endpoints: no CSRF, no session, no HTML.
r.Get("/.well-known/healthcheck.json", s.h.Healthcheck())
r.Method(http.MethodGet, "/metrics", s.metrics.Handler())
r.Handle(staticPrefix+"*", s.staticHandler())
// Everything a browser drives, behind CSRF. Safe methods are
// unaffected by it except that they are issued a token.
r.Group(func(r chi.Router) {
r.Use(s.mw.CSRF())
r.Get("/", s.h.Index())
r.Post("/widgets", s.h.CreateWidget())
if s.params.Config.Debug {
// Only with DEBUG=true: a route that panics on demand is a
// denial-of-service primitive in production.
r.Get("/debug/panic", s.h.Panic())
}
})
s.router = r
}
// staticHandler serves the embedded assets, without directory listings
// and with a long cache lifetime.
func (s *Server) staticHandler() http.Handler {
fileServer := http.FileServer(filesOnly{inner: http.FS(static.FS)})
return http.StripPrefix(staticPrefix, http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", staticCacheControl)
fileServer.ServeHTTP(w, r)
},
))
}
+180
View File
@@ -0,0 +1,180 @@
// Package server owns the HTTP server lifecycle and the route table.
package server
import (
"context"
"errors"
"fmt"
"log/slog"
"net"
"net/http"
"strconv"
"time"
"github.com/go-chi/chi/v5"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/handlers"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/middleware"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// Params defines dependencies for Server.
type Params struct {
fx.In
Logger *logger.Logger
Globals *globals.Globals
Config *config.Config
Middleware *middleware.Middleware
Handlers *handlers.Handlers
Metrics *telemetry.Metrics
}
// Server is the HTTP server and its lifecycle state.
type Server struct {
params Params
log *slog.Logger
mw *middleware.Middleware
h *handlers.Handlers
metrics *telemetry.Metrics
httpServer *http.Server
listener net.Listener
router *chi.Mux
serveErr chan error
}
// New creates the Server and hooks it into the fx lifecycle.
//
// The listener is opened during OnStart, not in a goroutine after it:
// binding is the part that fails (port in use, permission denied), and
// a failure there has to be a failure to start. A server that logs
// "listen failed" from a goroutine while fx reports a successful
// startup is a process that is up and serving nothing.
func New(lc fx.Lifecycle, params Params) (*Server, error) {
s := &Server{
params: params,
log: params.Logger.Get(),
mw: params.Middleware,
h: params.Handlers,
metrics: params.Metrics,
serveErr: make(chan error, 1),
}
lc.Append(fx.Hook{
OnStart: s.start,
OnStop: s.stop,
})
return s, nil
}
// Addr returns the address actually bound, which is what a test that
// asked for port 0 needs.
func (s *Server) Addr() string {
if s.listener == nil {
return ""
}
return s.listener.Addr().String()
}
// Handler exposes the router so tests can drive it with
// httptest.NewRecorder without binding a port.
func (s *Server) Handler() http.Handler {
return s.router
}
// ServeHTTP dispatches requests through the router.
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.router.ServeHTTP(w, r)
}
// start builds the routes, binds, and begins serving.
func (s *Server) start(ctx context.Context) error {
s.SetupRoutes()
addr := ":" + strconv.Itoa(s.params.Config.Port)
s.httpServer = s.newHTTPServer(addr)
// ListenConfig rather than net.Listen: the start context bounds
// name resolution and the bind, so a start that cannot bind fails
// within fx's start timeout instead of hanging inside it.
var lc net.ListenConfig
ln, err := lc.Listen(ctx, "tcp", addr)
if err != nil {
return fmt.Errorf("listening on %s: %w", addr, err)
}
s.listener = ln
s.log.Info("http listening",
"addr", ln.Addr().String(),
"debug", s.params.Config.Debug,
"metrics_protected", s.metrics.AuthRequired(),
)
go func() {
serveErr := s.httpServer.Serve(ln)
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
s.log.Error("http serve failed", "error", serveErr)
s.serveErr <- serveErr
return
}
s.serveErr <- nil
}()
return nil
}
// stop drains in-flight requests, bounded by the stop context.
//
// The bound is whichever of the two comes first: the deadline fx gives
// this hook, or the configured grace period. Taking the minimum is the
// point — a shutdown that outlives the context it was given is a
// shutdown the supervisor kills with SIGKILL, which is precisely the
// abrupt termination the draining was supposed to avoid.
//
// Past the deadline, Shutdown returns and every remaining connection is
// closed. That is a deliberate choice of a bounded, reported failure
// over an unbounded wait.
func (s *Server) stop(ctx context.Context) error {
if s.httpServer == nil {
return nil
}
grace := s.params.Config.ShutdownGrace
stopCtx, cancel := context.WithTimeout(ctx, grace)
defer cancel()
s.log.Info("http shutting down", "grace", grace.String())
start := time.Now()
err := s.httpServer.Shutdown(stopCtx)
if err != nil {
// Connections were still open at the deadline. Report it —
// silence here would hide requests that were cut off — but do
// not fail the stop sequence: there is nothing left to retry
// and the rest of the graph still has to close cleanly.
s.log.Error("http shutdown did not drain in time",
"error", err,
"waited", time.Since(start).Round(time.Millisecond).String(),
)
return nil
}
s.log.Info("http shutdown complete",
"waited", time.Since(start).Round(time.Millisecond).String(),
)
return nil
}
+560
View File
@@ -0,0 +1,560 @@
package server_test
import (
"encoding/json"
"io"
"net"
"net/http"
"net/http/cookiejar"
"net/url"
"path/filepath"
"regexp"
"strings"
"testing"
"time"
"go.uber.org/fx/fxtest"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/database"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/handlers"
"sneak.berlin/go/simplexcalc/internal/logger"
"sneak.berlin/go/simplexcalc/internal/middleware"
"sneak.berlin/go/simplexcalc/internal/render"
"sneak.berlin/go/simplexcalc/internal/server"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
// formNameField is the widget form's name input; refererHeader is the
// header gorilla/csrf checks on every state-changing request.
const (
formNameField = "name"
refererHeader = "Referer"
)
// response is a finished exchange: status, headers and the body, read
// and closed before the helper returns. Handing the tests a value
// rather than an *http.Response means no test can leak a connection by
// forgetting to close one.
type response struct {
status int
header http.Header
body string
}
// instance is a running service: a real listener on an ephemeral port,
// a real sqlite database in a temp dir, and the whole middleware chain.
// Nothing is stubbed, because the things most likely to be wrong —
// middleware order, route registration, embedded assets — are exactly
// what a stub would paper over.
type instance struct {
base string
client *http.Client
}
// testConfig is the configuration every instance starts from.
func testConfig(t *testing.T) *config.Config {
t.Helper()
dir := t.TempDir()
return &config.Config{
// Port 0: the OS picks a free one, so parallel tests do not
// fight over a fixed port.
Port: 0,
DataDir: dir,
DBPath: filepath.Join(dir, "test.db"),
Debug: true,
HSTS: false,
MaxRequestBody: 1 << 20,
RequestTimeout: 5 * time.Second,
ShutdownGrace: 2 * time.Second,
CSRFKeyEphemeral: true,
}
}
// build constructs the graph by hand. fx would do the same wiring; done
// explicitly, a constructor that starts needing a new dependency shows
// up here as a compile error rather than as a runtime resolution
// failure inside a container.
func build(t *testing.T, lc *fxtest.Lifecycle, cfg *config.Config) *server.Server {
t.Helper()
g := &globals.Globals{Appname: "simplexcalc", Version: "test", Buildarch: "amd64"}
log, err := logger.New(lc, logger.Params{Globals: g, Output: io.Discard})
if err != nil {
t.Fatalf("logger: %v", err)
}
db, err := database.New(lc, database.Params{Config: cfg, Logger: log})
if err != nil {
t.Fatalf("database: %v", err)
}
sentry, err := telemetry.NewSentry(lc, telemetry.SentryParams{
Config: cfg, Globals: g, Logger: log,
})
if err != nil {
t.Fatalf("sentry: %v", err)
}
metrics, err := telemetry.NewMetrics(telemetry.MetricsParams{Config: cfg})
if err != nil {
t.Fatalf("metrics: %v", err)
}
renderer, err := render.New(lc, render.Params{Globals: g})
if err != nil {
t.Fatalf("render: %v", err)
}
mw, err := middleware.New(lc, middleware.Params{
Config: cfg, Logger: log, Sentry: sentry, Metrics: metrics,
})
if err != nil {
t.Fatalf("middleware: %v", err)
}
h, err := handlers.New(lc, handlers.Params{
Config: cfg, Globals: g, Logger: log,
Database: db, Renderer: renderer, Sentry: sentry,
})
if err != nil {
t.Fatalf("handlers: %v", err)
}
srv, err := server.New(lc, server.Params{
Logger: log, Globals: g, Config: cfg,
Middleware: mw, Handlers: h, Metrics: metrics,
})
if err != nil {
t.Fatalf("server: %v", err)
}
return srv
}
// newInstance starts the service and registers its shutdown.
func newInstance(t *testing.T, mutate func(*config.Config)) *instance {
t.Helper()
cfg := testConfig(t)
if mutate != nil {
mutate(cfg)
}
lc := fxtest.NewLifecycle(t)
srv := build(t, lc, cfg)
lc.RequireStart()
t.Cleanup(lc.RequireStop)
jar, err := cookiejar.New(nil)
if err != nil {
t.Fatalf("cookie jar: %v", err)
}
// The listener binds the unspecified address, so Addr() reads
// "[::]:PORT". Dialling that works, but a cookie jar keyed on "::"
// does not, and the CSRF flow depends on cookies coming back — so
// the tests talk to the loopback address explicitly.
_, port, err := net.SplitHostPort(srv.Addr())
if err != nil {
t.Fatalf("parsing listen address %q: %v", srv.Addr(), err)
}
return &instance{
base: "http://127.0.0.1:" + port,
client: &http.Client{
Jar: jar,
Timeout: 10 * time.Second,
// The redirect after a successful POST is under test, so
// it must not be followed silently.
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
},
},
}
}
// host is the Host header value the service is reached by.
func (i *instance) host() string {
return strings.TrimPrefix(i.base, "http://")
}
func (i *instance) get(t *testing.T, path string) response {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, i.base+path, nil)
if err != nil {
t.Fatalf("building GET %s: %v", path, err)
}
return i.do(t, req)
}
// post submits a form the way a browser does — including the Referer,
// which gorilla/csrf checks against the request's Host on every
// state-changing request. Omitting it is itself a CSRF rejection, so a
// test that leaves it out proves nothing about the token.
func (i *instance) post(t *testing.T, path string, form url.Values) response {
t.Helper()
return i.postWithHeaders(t, path, form, map[string]string{
refererHeader: i.base + "/",
})
}
// postWithHeaders is post with the headers stated explicitly, for the
// reverse-proxy cases.
func (i *instance) postWithHeaders(
t *testing.T, path string, form url.Values, headers map[string]string,
) response {
t.Helper()
req, err := http.NewRequestWithContext(
t.Context(), http.MethodPost, i.base+path, strings.NewReader(form.Encode()),
)
if err != nil {
t.Fatalf("building POST %s: %v", path, err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
for k, v := range headers {
req.Header.Set(k, v)
}
return i.do(t, req)
}
func (i *instance) do(t *testing.T, req *http.Request) response {
t.Helper()
resp, err := i.client.Do(req)
if err != nil {
t.Fatalf("%s %s: %v", req.Method, req.URL.Path, err)
}
body, readErr := io.ReadAll(resp.Body)
closeErr := resp.Body.Close()
if closeErr != nil {
t.Errorf("closing response body: %v", closeErr)
}
if readErr != nil {
t.Fatalf("reading %s %s: %v", req.Method, req.URL.Path, readErr)
}
return response{status: resp.StatusCode, header: resp.Header, body: string(body)}
}
// TestServesPageFromEmbeddedTemplate is one of the definition-of-done
// claims: a running process answers / from the template compiled into
// it.
func TestServesPageFromEmbeddedTemplate(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
resp := i.get(t, "/")
if resp.status != http.StatusOK {
t.Fatalf("status = %d, want 200", resp.status)
}
for _, want := range []string{"<!doctype html>", "<nav>", "<footer>", "widgets ("} {
if !strings.Contains(resp.body, want) {
t.Errorf("page is missing %q", want)
}
}
if resp.header.Get("X-Content-Type-Options") != "nosniff" {
t.Error("security headers are missing from a normal page response")
}
if resp.header.Get(middleware.RequestIDHeader) == "" {
t.Error("no request id on the response")
}
}
// TestServesEmbeddedStaticAsset, with the cache policy the route sets.
func TestServesEmbeddedStaticAsset(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
resp := i.get(t, "/static/css/style.css")
if resp.status != http.StatusOK {
t.Fatalf("status = %d, want 200", resp.status)
}
if !strings.Contains(resp.body, "--fg:") {
t.Error("the served asset is not the embedded stylesheet")
}
if !strings.Contains(resp.header.Get("Cache-Control"), "immutable") {
t.Errorf("Cache-Control = %q", resp.header.Get("Cache-Control"))
}
}
// TestStaticDirectoriesAreNotListed: enumerating what the binary ships
// is a capability no caller needs.
func TestStaticDirectoriesAreNotListed(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
for _, path := range []string{"/static/", "/static/css/"} {
resp := i.get(t, path)
if resp.status != http.StatusNotFound {
t.Errorf("GET %s: status = %d, want 404", path, resp.status)
}
if strings.Contains(resp.body, "style.css") {
t.Errorf("GET %s listed the directory contents", path)
}
}
}
// TestHealthcheckReportsMigratedSchema proves the embedded migrations
// ran during startup, from a directory that was empty a moment ago.
func TestHealthcheckReportsMigratedSchema(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
resp := i.get(t, "/.well-known/healthcheck.json")
if resp.status != http.StatusOK {
t.Fatalf("status = %d, want 200", resp.status)
}
var health handlers.HealthResponse
err := json.Unmarshal([]byte(resp.body), &health)
if err != nil {
t.Fatalf("decoding healthcheck: %v (body %q)", err, resp.body)
}
if !health.OK || !health.DatabaseOK {
t.Errorf("healthcheck reports unhealthy: %+v", health)
}
if health.SchemaVersion != 1 {
t.Errorf("schema_version = %d, want 1", health.SchemaVersion)
}
}
// TestMetricsEndpointIsServedAndGated covers both halves of the metrics
// requirement in one running service.
func TestMetricsEndpointIsServedAndGated(t *testing.T) {
t.Parallel()
t.Run("open when unconfigured", func(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
// One request first: the middleware records a request after
// the response is written, so a scrape that is the very first
// request to the process legitimately has no HTTP series in it
// yet.
i.get(t, "/")
resp := i.get(t, "/metrics")
if resp.status != http.StatusOK {
t.Fatalf("status = %d, want 200", resp.status)
}
want := `http_requests_total{code="200",method="GET",route="/"}`
if !strings.Contains(resp.body, want) {
t.Errorf("the scrape does not report the request that preceded it:\n%s", resp.body)
}
})
t.Run("gated when configured", func(t *testing.T) {
t.Parallel()
i := newInstance(t, func(c *config.Config) {
c.MetricsUser = "scraper"
c.MetricsPassword = "hunter2"
})
resp := i.get(t, "/metrics")
if resp.status != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401", resp.status)
}
})
}
// TestUnknownRouteRendersErrorPage: a 404 still carries the site's
// headers and chrome rather than net/http's bare text.
func TestUnknownRouteRendersErrorPage(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
resp := i.get(t, "/no-such-page")
if resp.status != http.StatusNotFound {
t.Fatalf("status = %d, want 404", resp.status)
}
if !strings.Contains(resp.body, "No such page") ||
!strings.Contains(resp.body, "<nav>") {
t.Errorf("404 is not the rendered error page:\n%s", resp.body)
}
}
// TestPanicIsAnsweredWith500 exercises the recovery middleware through
// the whole stack, on the debug-only route.
func TestPanicIsAnsweredWith500(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
resp := i.get(t, "/debug/panic")
if resp.status != http.StatusInternalServerError {
t.Fatalf("status = %d, want 500", resp.status)
}
if strings.Contains(resp.body, "deliberate panic") {
t.Error("the panic value reached the client")
}
}
// TestDebugRouteAbsentWithoutDebug: the panic route is a denial-of-
// service primitive, so it must not exist in a normal deployment.
func TestDebugRouteAbsentWithoutDebug(t *testing.T) {
t.Parallel()
i := newInstance(t, func(c *config.Config) { c.Debug = false })
resp := i.get(t, "/debug/panic")
if resp.status != http.StatusNotFound {
t.Errorf("status = %d, want 404 with DEBUG off", resp.status)
}
}
// csrfField finds the hidden token in a rendered form.
var csrfField = regexp.MustCompile(
`<input type="hidden" name="gorilla\.csrf\.Token" value="([^"]+)"`,
)
// token fetches the index page and returns the CSRF token its form
// carries.
func (i *instance) token(t *testing.T) string {
t.Helper()
page := i.get(t, "/")
match := csrfField.FindStringSubmatch(page.body)
if match == nil {
t.Fatalf("no CSRF field in the rendered form:\n%s", page.body)
}
return match[1]
}
// TestCSRFProtectsStateChange is the whole point of the CSRF
// middleware: the same POST must be refused without a token and
// accepted with one.
func TestCSRFProtectsStateChange(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
// A well-formed cross-site request: everything a browser would
// send except the token.
resp := i.post(t, "/widgets", url.Values{formNameField: {"unauthorised"}})
if resp.status != http.StatusForbidden {
t.Fatalf("POST without a CSRF token: status = %d, want 403", resp.status)
}
resp = i.post(t, "/widgets", url.Values{
formNameField: {"authorised"},
"size": {"4KiB"},
"gorilla.csrf.Token": {i.token(t)},
})
if resp.status != http.StatusSeeOther {
t.Fatalf("POST with a CSRF token: status = %d, want 303", resp.status)
}
// And the write actually happened.
after := i.get(t, "/")
if !strings.Contains(after.body, "authorised") {
t.Error("the created widget is not on the page")
}
if !strings.Contains(after.body, "4.0 KiB") {
t.Error("the size was not parsed and rendered")
}
if strings.Contains(after.body, "unauthorised") {
t.Error("the rejected POST created a widget anyway")
}
}
// TestCSRFBehindTLSTerminatingProxy is the deployment shape this
// service actually runs in: the browser speaks HTTPS to a proxy and the
// proxy speaks plaintext HTTP here, saying so with X-Forwarded-Proto.
// The token must still be accepted, with the stricter TLS-side Referer
// rules applied to the browser's real scheme rather than to the
// connection's.
func TestCSRFBehindTLSTerminatingProxy(t *testing.T) {
t.Parallel()
i := newInstance(t, nil)
form := url.Values{
formNameField: {"through the proxy"},
"gorilla.csrf.Token": {i.token(t)},
}
// A cleartext Referer on a connection the proxy says was HTTPS is
// exactly the machine-in-the-middle case the strict check exists
// for, and must be refused.
resp := i.postWithHeaders(t, "/widgets", form, map[string]string{
"X-Forwarded-Proto": "https",
refererHeader: "http://" + i.host() + "/",
})
if resp.status != http.StatusForbidden {
t.Errorf("cleartext referer under X-Forwarded-Proto=https: status = %d, want 403",
resp.status)
}
// The genuine article: a same-origin HTTPS referer.
resp = i.postWithHeaders(t, "/widgets", form, map[string]string{
"X-Forwarded-Proto": "https",
refererHeader: "https://" + i.host() + "/",
})
if resp.status != http.StatusSeeOther {
t.Errorf("valid submission through a TLS-terminating proxy: status = %d, want 303",
resp.status)
}
}
// TestOversizedRequestIsRefused: the body cap is wired into the running
// chain, not just unit-tested in isolation.
func TestOversizedRequestIsRefused(t *testing.T) {
t.Parallel()
i := newInstance(t, func(c *config.Config) { c.MaxRequestBody = 1024 })
// The cap is enforced ahead of CSRF in the chain, so an oversized
// body is refused as oversized rather than as untokened.
resp := i.post(t, "/widgets",
url.Values{formNameField: {strings.Repeat("x", 4096)}})
if resp.status != http.StatusRequestEntityTooLarge {
t.Errorf("status = %d, want 413", resp.status)
}
}
+48
View File
@@ -0,0 +1,48 @@
package server
import (
"io/fs"
"net/http"
)
// filesOnly wraps an http.FileSystem so that directories do not exist as
// far as http.FileServer is concerned. FileServer asks the filesystem for
// the directory first and generates its index from what it gets back;
// refusing the Open is therefore the whole control, and it leaves the
// path handling, content sniffing and range support of FileServer intact.
//
// Enumerating what a binary ships is a capability no caller needs, and
// the embedded set grows as a project seeded from this template adds to
// it.
type filesOnly struct {
inner http.FileSystem
}
// Open serves a file and refuses a directory with fs.ErrNotExist, which
// http.FileServer maps to 404 — the same answer a path that is not
// embedded at all gets.
func (f filesOnly) Open(name string) (http.File, error) {
file, err := f.inner.Open(name)
if err != nil {
// Unwrapped on purpose: http.FileServer inspects this error,
// and wrapping hides fs.ErrNotExist from it.
//nolint:wrapcheck // see above.
return nil, err
}
info, err := file.Stat()
if err != nil {
_ = file.Close()
//nolint:wrapcheck // as above.
return nil, err
}
if info.IsDir() {
_ = file.Close()
return nil, fs.ErrNotExist
}
return file, nil
}
+162
View File
@@ -0,0 +1,162 @@
package telemetry
import (
"crypto/sha256"
"crypto/subtle"
"net/http"
"time"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/collectors"
"github.com/prometheus/client_golang/prometheus/promhttp"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
)
// durationBuckets span a fast in-process handler through a slow
// upstream call. Prometheus's defaults top out at 10s, which hides the
// tail this service's request timeout permits.
//
//nolint:gochecknoglobals // a bucket list is a declaration, not mutable state.
var durationBuckets = []float64{
0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30,
}
// MetricsParams defines dependencies for Metrics.
type MetricsParams struct {
fx.In
Config *config.Config
}
// Metrics owns the registry and the HTTP series. A private registry,
// not the global default: what this process exports is then exactly
// what this code registered, and a linked library cannot quietly add to
// it.
type Metrics struct {
registry *prometheus.Registry
requests *prometheus.CounterVec
duration *prometheus.HistogramVec
inflight prometheus.Gauge
user string
password string
}
// NewMetrics builds the registry and registers the collectors.
func NewMetrics(params MetricsParams) (*Metrics, error) {
m := &Metrics{
registry: prometheus.NewRegistry(),
user: params.Config.MetricsUser,
password: params.Config.MetricsPassword,
}
m.requests = prometheus.NewCounterVec(
prometheus.CounterOpts{
Name: "http_requests_total",
Help: "Total HTTP requests by method, route pattern and status code.",
},
// The route PATTERN, never the path: labelling by path turns
// every distinct URL into a new time series, and a crawler
// then owns the memory of the process.
[]string{"method", "route", "code"},
)
m.duration = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Name: "http_request_duration_seconds",
Help: "HTTP request duration by method and route pattern.",
Buckets: durationBuckets,
},
[]string{"method", "route"},
)
m.inflight = prometheus.NewGauge(prometheus.GaugeOpts{
Name: "http_requests_in_flight",
Help: "HTTP requests currently being served.",
})
m.registry.MustRegister(
m.requests,
m.duration,
m.inflight,
collectors.NewGoCollector(),
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
)
return m, nil
}
// Observe records one finished request.
func (m *Metrics) Observe(method, route, code string, d time.Duration) {
m.requests.WithLabelValues(method, route, code).Inc()
m.duration.WithLabelValues(method, route).Observe(d.Seconds())
}
// InFlightAdd adjusts the in-flight gauge.
func (m *Metrics) InFlightAdd(delta float64) {
m.inflight.Add(delta)
}
// Registry exposes the registry so tests can gather what was recorded.
func (m *Metrics) Registry() *prometheus.Registry {
return m.registry
}
// AuthRequired reports whether /metrics is credential-gated.
func (m *Metrics) AuthRequired() bool {
return m.user != "" && m.password != ""
}
// Handler serves the exposition format, behind HTTP basic auth when
// credentials are configured.
//
// Metrics are not public: they leak route names, traffic volume,
// version and process memory layout. When no credentials are set the
// endpoint is served open, which is correct for a private network and
// documented as such in the README; config refuses the half-configured
// case, so "open" is always something the operator chose rather than
// something a typo produced.
func (m *Metrics) Handler() http.Handler {
h := promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{
// A collector that errors should not take the scrape down
// with a 500 the operator has to go and interpret.
ErrorHandling: promhttp.ContinueOnError,
})
if !m.AuthRequired() {
return h
}
return m.basicAuth(h)
}
// basicAuth gates h. Comparison is over SHA-256 digests through
// subtle.ConstantTimeCompare: comparing the raw strings would leak the
// credential length and the position of the first wrong byte through
// timing, and hashing first makes the comparison fixed-width.
func (m *Metrics) basicAuth(h http.Handler) http.Handler {
wantUser := sha256.Sum256([]byte(m.user))
wantPass := sha256.Sum256([]byte(m.password))
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, pass, ok := r.BasicAuth()
if ok {
gotUser := sha256.Sum256([]byte(user))
gotPass := sha256.Sum256([]byte(pass))
userOK := subtle.ConstantTimeCompare(gotUser[:], wantUser[:]) == 1
passOK := subtle.ConstantTimeCompare(gotPass[:], wantPass[:]) == 1
if userOK && passOK {
h.ServeHTTP(w, r)
return
}
}
w.Header().Set("WWW-Authenticate", `Basic realm="metrics", charset="UTF-8"`)
http.Error(w, "unauthorized", http.StatusUnauthorized)
})
}
+152
View File
@@ -0,0 +1,152 @@
package telemetry_test
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/telemetry"
)
func newMetrics(t *testing.T, user, password string) *telemetry.Metrics {
t.Helper()
m, err := telemetry.NewMetrics(telemetry.MetricsParams{
Config: &config.Config{MetricsUser: user, MetricsPassword: password},
})
if err != nil {
t.Fatalf("building metrics: %v", err)
}
return m
}
// TestMetricsRequireCredentialsWhenConfigured: an exposition endpoint
// leaks route names, traffic volume and process layout, so credentials
// have to actually be enforced.
func TestMetricsRequireCredentialsWhenConfigured(t *testing.T) {
t.Parallel()
m := newMetrics(t, user, pass)
if !m.AuthRequired() {
t.Fatal("credentials are configured but AuthRequired is false")
}
const (
unauthorized = http.StatusUnauthorized
ok = http.StatusOK
)
cases := []struct {
name string
user, pass string
useAuth bool
want int
}{
{name: "no credentials", want: unauthorized},
{name: "wrong password", user: user, pass: "no", useAuth: true, want: unauthorized},
{name: "wrong user", user: "nobody", pass: pass, useAuth: true, want: unauthorized},
{name: "correct", user: user, pass: pass, useAuth: true, want: ok},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
req := scrapeReq(t)
if tc.useAuth {
req.SetBasicAuth(tc.user, tc.pass)
}
w := httptest.NewRecorder()
m.Handler().ServeHTTP(w, req)
if w.Code != tc.want {
t.Errorf("status = %d, want %d", w.Code, tc.want)
}
if tc.want == http.StatusUnauthorized {
if w.Header().Get("WWW-Authenticate") == "" {
t.Error("a 401 with no WWW-Authenticate gives the client nothing to do")
}
if strings.Contains(w.Body.String(), "http_requests_total") {
t.Error("metrics were served in the body of a 401")
}
}
})
}
}
// TestMetricsOpenWhenNoCredentials documents the other half: with
// nothing configured the endpoint is open, which config permits only
// when BOTH values are absent.
func TestMetricsOpenWhenNoCredentials(t *testing.T) {
t.Parallel()
m := newMetrics(t, "", "")
if m.AuthRequired() {
t.Fatal("no credentials configured but AuthRequired is true")
}
w := httptest.NewRecorder()
m.Handler().ServeHTTP(w, scrapeReq(t))
if w.Code != http.StatusOK {
t.Errorf("status = %d, want 200", w.Code)
}
}
// TestObservedRequestsAreExported: the middleware records through
// Observe, and what it records has to come back out of the endpoint.
func TestObservedRequestsAreExported(t *testing.T) {
t.Parallel()
m := newMetrics(t, "", "")
m.Observe(http.MethodGet, "/widgets/{id}", "200", 25*time.Millisecond)
w := httptest.NewRecorder()
m.Handler().ServeHTTP(w, scrapeReq(t))
body := w.Body.String()
for _, want := range []string{
`http_requests_total{code="200",method="GET",route="/widgets/{id}"} 1`,
"http_request_duration_seconds_bucket",
"go_goroutines", // the Go collector is registered
} {
if !strings.Contains(body, want) {
t.Errorf("exposition output is missing %q", want)
}
}
}
// TestSentryDisabledWithoutDSN: every method must be safe with no DSN,
// because that is how the service runs in development and in tests.
func TestSentryDisabledWithoutDSN(t *testing.T) {
t.Parallel()
s, err := telemetry.NewSentry(nil, telemetry.SentryParams{
Config: &config.Config{},
Globals: testGlobals(),
Logger: testLogger(t),
})
if err != nil {
t.Fatalf("building sentry: %v", err)
}
if s.Enabled() {
t.Error("sentry reports enabled with no DSN")
}
// Must not panic.
s.CaptureError(nil)
s.CaptureError(errTest)
s.CapturePanic("boom", []byte("stack"))
}
+112
View File
@@ -0,0 +1,112 @@
// Package telemetry owns error reporting (Sentry) and metrics
// (Prometheus). Both are optional at runtime and neither is allowed to
// take the process down: a monitoring backend that is unreachable must
// not stop the service it monitors.
package telemetry
import (
"context"
"fmt"
"log/slog"
"time"
"github.com/getsentry/sentry-go"
"go.uber.org/fx"
"sneak.berlin/go/simplexcalc/internal/config"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/logger"
)
// flushTimeout bounds how long shutdown waits for queued events to
// reach Sentry. Exceeding it drops the tail rather than hanging the
// stop sequence.
const flushTimeout = 2 * time.Second
// SentryParams defines dependencies for Sentry.
type SentryParams struct {
fx.In
Config *config.Config
Globals *globals.Globals
Logger *logger.Logger
}
// Sentry wraps the client. When SENTRY_DSN is unset the wrapper still
// exists and every method is a no-op, so no caller needs a nil check
// and no caller behaves differently in development.
type Sentry struct {
enabled bool
log *slog.Logger
}
// NewSentry initialises the client if a DSN is configured. A DSN that
// is present but malformed has already failed config parsing; a DSN the
// client itself rejects is logged and reporting stays off, because a
// telemetry backend is not a reason to refuse to serve.
func NewSentry(lc fx.Lifecycle, params SentryParams) (*Sentry, error) {
s := &Sentry{log: params.Logger.Get()}
if params.Config.SentryDSN == "" {
s.log.Info("sentry disabled", "reason", "no DSN configured")
return s, nil
}
err := sentry.Init(sentry.ClientOptions{
Dsn: params.Config.SentryDSN,
Environment: params.Config.SentryEnvironment,
Release: params.Globals.Appname + "@" + params.Globals.Version,
// Panics are reported explicitly by the recovery middleware,
// which also has to answer the request; letting the SDK
// re-raise them would take the process down.
Debug: params.Config.Debug,
})
if err != nil {
s.log.Error("sentry init failed; error reporting is off", "error", err)
return s, nil
}
s.enabled = true
s.log.Info("sentry enabled", "environment", params.Config.SentryEnvironment)
lc.Append(fx.Hook{
OnStop: func(_ context.Context) error {
sentry.Flush(flushTimeout)
return nil
},
})
return s, nil
}
// Enabled reports whether events are actually being sent.
func (s *Sentry) Enabled() bool {
return s.enabled
}
// CaptureError reports an error, and always logs it. Logging is not
// conditional on Sentry being on: the log is the record of record, and
// Sentry is a convenience on top of it.
func (s *Sentry) CaptureError(err error) {
if err == nil {
return
}
s.log.Error("captured error", "error", err)
if s.enabled {
sentry.CaptureException(err)
}
}
// CapturePanic reports a recovered panic value with its stack.
func (s *Sentry) CapturePanic(v any, stack []byte) {
s.log.Error("recovered panic", "panic", fmt.Sprint(v), "stack", string(stack))
if s.enabled {
sentry.CurrentHub().Recover(v)
}
}
+43
View File
@@ -0,0 +1,43 @@
package telemetry_test
import (
"errors"
"io"
"net/http"
"net/http/httptest"
"testing"
"sneak.berlin/go/simplexcalc/internal/globals"
"sneak.berlin/go/simplexcalc/internal/logger"
)
// scrapeReq is a request to /metrics carrying the test's context.
func scrapeReq(t *testing.T) *http.Request {
t.Helper()
return httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/metrics", nil)
}
// errTest is a stand-in error for the capture paths.
var errTest = errors.New("test error")
// The metrics credentials used across these tests.
const (
user = "scraper"
pass = "hunter2"
)
func testGlobals() *globals.Globals {
return &globals.Globals{Appname: "simplexcalc", Version: "test", Buildarch: "amd64"}
}
func testLogger(t *testing.T) *logger.Logger {
t.Helper()
log, err := logger.New(nil, logger.Params{Globals: testGlobals(), Output: io.Discard})
if err != nil {
t.Fatalf("building logger: %v", err)
}
return log
}