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:
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
)
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
},
|
||||
))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user