A SimpleX Chat bot that answers arithmetic (closes #1)
check / check (push) Successful in 54s
check / check (push) Successful in 54s
Remove the template's HTTP service, database and fx wiring. Add exact arithmetic on go/parser and go/constant, a client that runs simplex-chat as a child process and drives its WebSocket API, and the bot, which keeps an auto-accepting address and replies to each message. The image adds the checksum-pinned simplex-chat v7.0.2 on Ubuntu 22.04. Model: opus-5-5
This commit is contained in:
@@ -1,72 +0,0 @@
|
||||
// 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,230 @@
|
||||
// Package bot is the calculator: it runs the SimpleX Chat client, makes
|
||||
// sure the bot has a contact address that accepts everyone, and answers
|
||||
// every text message with the value of the arithmetic in it.
|
||||
package bot
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/simplexcalc/internal/calc"
|
||||
"sneak.berlin/go/simplexcalc/internal/simplex"
|
||||
)
|
||||
|
||||
// DisplayName is the name of the bot's SimpleX profile, given to it
|
||||
// when the profile is created on the first start.
|
||||
const DisplayName = "calc"
|
||||
|
||||
// Welcome is sent to everyone whose contact request the bot accepts.
|
||||
const Welcome = "Send me arithmetic, such as 2 + 2 or 5 * 5/2, " +
|
||||
"and I will reply with the result."
|
||||
|
||||
const (
|
||||
// chatPort is where the chat client serves its API, on localhost
|
||||
// inside the bot's own container.
|
||||
chatPort = 5225
|
||||
|
||||
// connectTimeout bounds the wait for a freshly started chat client
|
||||
// to open its API, which includes creating or migrating the
|
||||
// database.
|
||||
connectTimeout = 60 * time.Second
|
||||
|
||||
// setupTimeout bounds each setup command. Creating an address
|
||||
// talks to SimpleX relays over the network.
|
||||
setupTimeout = 2 * time.Minute
|
||||
|
||||
// retryInterval paces the connection attempts.
|
||||
retryInterval = 250 * time.Millisecond
|
||||
|
||||
dataDirMode = 0o700
|
||||
)
|
||||
|
||||
var errExited = errors.New("simplex-chat exited")
|
||||
|
||||
// Run starts the chat client with its database in dataDir, connects to
|
||||
// it, sets up the bot's address, and answers messages until ctx is
|
||||
// cancelled — which is a clean stop and returns nil — or until the chat
|
||||
// client or the connection to it fails, which returns the error.
|
||||
func Run(ctx context.Context, log *slog.Logger, dataDir string) error {
|
||||
err := os.MkdirAll(dataDir, dataDirMode)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating data directory: %w", err)
|
||||
}
|
||||
|
||||
// Cancelling this stops the chat client; the deferred wait makes
|
||||
// Run return only once it has exited, whatever path Run takes.
|
||||
cliCtx, stopCLI := context.WithCancel(ctx)
|
||||
|
||||
cli, err := simplex.StartCLI(cliCtx, log, filepath.Join(dataDir, "simplex"),
|
||||
DisplayName, chatPort)
|
||||
if err != nil {
|
||||
stopCLI()
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
defer func() {
|
||||
stopCLI()
|
||||
<-cli.Done()
|
||||
}()
|
||||
|
||||
client, err := connect(ctx, log, cli)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
err = setUp(ctx, log, client)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil
|
||||
case <-client.Done():
|
||||
return fmt.Errorf("%w: %w", simplex.ErrClosed, client.Err())
|
||||
case <-cli.Done():
|
||||
return fmt.Errorf("%w: %w", errExited, cli.Err())
|
||||
}
|
||||
}
|
||||
|
||||
// connect waits for the chat client to open its API and connects to it.
|
||||
func connect(
|
||||
ctx context.Context, log *slog.Logger, cli *simplex.CLI,
|
||||
) (*simplex.Client, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, connectTimeout)
|
||||
defer cancel()
|
||||
|
||||
url := "ws://127.0.0.1:" + strconv.Itoa(chatPort)
|
||||
|
||||
for {
|
||||
client, err := simplex.Dial(ctx, url, log, handle(log))
|
||||
if err == nil {
|
||||
return client, nil
|
||||
}
|
||||
|
||||
select {
|
||||
case <-cli.Done():
|
||||
return nil, fmt.Errorf("%w before opening its API: %w", errExited, cli.Err())
|
||||
case <-ctx.Done():
|
||||
return nil, fmt.Errorf("waiting for simplex-chat: %w (last error: %w)",
|
||||
ctx.Err(), err)
|
||||
case <-time.After(retryInterval):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// setUp gives the bot a long-term contact address, creating it on the
|
||||
// first start, and sets it to accept every contact request and to greet
|
||||
// each new contact. The settings are written on every start, so an
|
||||
// address whose settings were changed by hand is put right.
|
||||
func setUp(ctx context.Context, log *slog.Logger, client *simplex.Client) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, setupTimeout)
|
||||
defer cancel()
|
||||
|
||||
user, err := client.ActiveUser(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading the bot's profile: %w", err)
|
||||
}
|
||||
|
||||
link, ok, err := client.Address(ctx, user.UserID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading the bot's address: %w", err)
|
||||
}
|
||||
|
||||
if !ok {
|
||||
log.Info("creating the bot's address")
|
||||
|
||||
link, err = client.CreateAddress(ctx, user.UserID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating the bot's address: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = client.SetAddressSettings(ctx, user.UserID, simplex.AddressSettings{
|
||||
AutoAccept: &simplex.AutoAccept{AcceptIncognito: false},
|
||||
AutoReply: &simplex.MsgContent{Type: "text", Text: Welcome},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("setting the bot's address to accept everyone: %w", err)
|
||||
}
|
||||
|
||||
log.Info("ready",
|
||||
"display_name", user.Profile.DisplayName,
|
||||
"address", link.ShortLink,
|
||||
"full_address", link.FullLink,
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// handle answers each text message a contact sends.
|
||||
func handle(log *slog.Logger) simplex.EventHandler {
|
||||
return func(c *simplex.Client, ev simplex.Event) {
|
||||
switch ev.Type {
|
||||
case simplex.TypeNewChatItems:
|
||||
var r simplex.NewChatItems
|
||||
|
||||
err := ev.Decode(&r)
|
||||
if err != nil {
|
||||
log.Warn("ignoring an event", "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
for _, item := range r.ChatItems {
|
||||
msg, ok := item.Message()
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
err = c.SendText(msg.ContactID, msg.ItemID, Reply(msg.Text))
|
||||
if err != nil {
|
||||
log.Error("replying", "contact_id", msg.ContactID, "error", err)
|
||||
}
|
||||
}
|
||||
case simplex.TypeContactConnected:
|
||||
var r simplex.ContactConnected
|
||||
|
||||
err := ev.Decode(&r)
|
||||
if err != nil {
|
||||
log.Warn("ignoring an event", "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
log.Info("contact connected", "contact_id", r.Contact.ContactID)
|
||||
default:
|
||||
log.Debug("event", "type", ev.Type)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Reply is the bot's answer to a message: the value of the arithmetic
|
||||
// in it, or a short explanation of why there is none.
|
||||
func Reply(text string) string {
|
||||
result, err := calc.Evaluate(text)
|
||||
|
||||
switch {
|
||||
case err == nil:
|
||||
return result
|
||||
case errors.Is(err, calc.ErrTooLong):
|
||||
return fmt.Sprintf("That is too long for me: at most %d characters, please.",
|
||||
calc.MaxInputLength)
|
||||
case errors.Is(err, calc.ErrDivisionByZero):
|
||||
return "I cannot divide by zero."
|
||||
case errors.Is(err, calc.ErrTooLarge):
|
||||
return "The result is too large for me."
|
||||
default:
|
||||
return "I only understand arithmetic: numbers, + - * / and " +
|
||||
"parentheses, such as 5 * 5/2."
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package bot_test
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sneak.berlin/go/simplexcalc/internal/bot"
|
||||
"sneak.berlin/go/simplexcalc/internal/calc"
|
||||
)
|
||||
|
||||
// TestReply: a result is sent bare, and every way of failing gets its
|
||||
// own short explanation rather than silence.
|
||||
func TestReply(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for in, want := range map[string]string{
|
||||
"2 + 2": "4",
|
||||
"5 * 5/2": "12.5",
|
||||
} {
|
||||
if got := bot.Reply(in); got != want {
|
||||
t.Errorf("Reply(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
for in, want := range map[string]string{
|
||||
"hello": "I only understand arithmetic",
|
||||
"1 / 0": "I cannot divide by zero.",
|
||||
"1e400": "The result is too large for me.",
|
||||
strings.Repeat("1+", calc.MaxInputLength) + "1": "That is too long for me",
|
||||
} {
|
||||
if got := bot.Reply(in); !strings.HasPrefix(got, want) {
|
||||
t.Errorf("Reply(%q) = %q, want it to start %q", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
// Package calc evaluates the arithmetic people send the bot: decimal
|
||||
// numbers, + - * /, unary minus and parentheses.
|
||||
//
|
||||
// The expression is parsed by go/parser and computed by go/constant,
|
||||
// which does exact rational arithmetic: 5 * 5/2 is exactly 12.5, and
|
||||
// 0.1 + 0.2 is exactly 0.3, so a result carries no binary floating
|
||||
// point noise until the moment it is formatted.
|
||||
package calc
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"go/ast"
|
||||
"go/constant"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"math"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// MaxInputLength caps an expression, in bytes, so a message cannot make
|
||||
// the bot do unbounded work. Every operation's cost grows with the size
|
||||
// of its operands, and the operands can only grow with the input.
|
||||
const MaxInputLength = 256
|
||||
|
||||
// Results of magnitude plainUpper or more are written in exponent form
|
||||
// (1e+21 rather than twenty-two digits), and so are fractions smaller
|
||||
// than plainLower (1e-07 rather than 0.0000001).
|
||||
const (
|
||||
plainUpper = 1e21
|
||||
plainLower = 1e-6
|
||||
)
|
||||
|
||||
// Errors returned by Evaluate. The bot turns each into a reply.
|
||||
var (
|
||||
ErrTooLong = errors.New("expression too long")
|
||||
ErrNotArithmetic = errors.New("not an arithmetic expression")
|
||||
ErrDivisionByZero = errors.New("division by zero")
|
||||
ErrTooLarge = errors.New("result too large")
|
||||
)
|
||||
|
||||
// decimalLiteral is the only number syntax accepted. Go's own literal
|
||||
// syntax is wider, and parts of it are traps for someone typing
|
||||
// arithmetic: 010 is octal 8, and 0x10, 1_000 and 1i are not what a
|
||||
// calculator user means by a number.
|
||||
var decimalLiteral = regexp.MustCompile(
|
||||
`^([0-9]+\.?[0-9]*|\.[0-9]+)([eE][+-]?[0-9]+)?$`,
|
||||
)
|
||||
|
||||
// Evaluate computes an arithmetic expression and returns its result as
|
||||
// text: whole numbers without a decimal point, fractions in the
|
||||
// shortest form that reads back as the same float64.
|
||||
func Evaluate(input string) (string, error) {
|
||||
s := strings.TrimSpace(input)
|
||||
if len(s) > MaxInputLength {
|
||||
return "", ErrTooLong
|
||||
}
|
||||
|
||||
if s == "" {
|
||||
return "", ErrNotArithmetic
|
||||
}
|
||||
|
||||
expr, err := parser.ParseExpr(s)
|
||||
if err != nil {
|
||||
return "", ErrNotArithmetic
|
||||
}
|
||||
|
||||
v, err := eval(expr)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return format(v)
|
||||
}
|
||||
|
||||
// eval walks the syntax tree, allowing only the node types and
|
||||
// operators of arithmetic. Anything else — identifiers, calls, strings,
|
||||
// shifts, comparisons — is refused, not evaluated.
|
||||
func eval(e ast.Expr) (constant.Value, error) {
|
||||
switch n := e.(type) {
|
||||
case *ast.BasicLit:
|
||||
return literal(n)
|
||||
case *ast.ParenExpr:
|
||||
return eval(n.X)
|
||||
case *ast.UnaryExpr:
|
||||
if n.Op != token.ADD && n.Op != token.SUB {
|
||||
return nil, ErrNotArithmetic
|
||||
}
|
||||
|
||||
x, err := eval(n.X)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return constant.UnaryOp(n.Op, x, 0), nil
|
||||
case *ast.BinaryExpr:
|
||||
return binary(n)
|
||||
default:
|
||||
return nil, ErrNotArithmetic
|
||||
}
|
||||
}
|
||||
|
||||
func binary(n *ast.BinaryExpr) (constant.Value, error) {
|
||||
switch n.Op { //nolint:exhaustive // every other operator is refused.
|
||||
case token.ADD, token.SUB, token.MUL, token.QUO:
|
||||
default:
|
||||
return nil, ErrNotArithmetic
|
||||
}
|
||||
|
||||
x, err := eval(n.X)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
y, err := eval(n.Y)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// constant.BinaryOp panics on a zero divisor.
|
||||
if n.Op == token.QUO && constant.Sign(y) == 0 {
|
||||
return nil, ErrDivisionByZero
|
||||
}
|
||||
|
||||
// token.QUO divides exactly, integers included: 25/2 is 12.5.
|
||||
v := constant.BinaryOp(x, n.Op, y)
|
||||
|
||||
// go/constant represents an overflow to infinity as Unknown.
|
||||
if v.Kind() == constant.Unknown {
|
||||
return nil, ErrTooLarge
|
||||
}
|
||||
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func literal(n *ast.BasicLit) (constant.Value, error) {
|
||||
if n.Kind != token.INT && n.Kind != token.FLOAT {
|
||||
return nil, ErrNotArithmetic
|
||||
}
|
||||
|
||||
if !decimalLiteral.MatchString(n.Value) {
|
||||
return nil, ErrNotArithmetic
|
||||
}
|
||||
|
||||
// Read as FLOAT whatever the token says, which makes every literal
|
||||
// decimal: as INT, a leading zero would make it octal.
|
||||
v := constant.MakeFromLiteral(n.Value, token.FLOAT, 0)
|
||||
|
||||
// The syntax was checked above, so Unknown here means the exponent
|
||||
// overflowed.
|
||||
if v.Kind() == constant.Unknown {
|
||||
return nil, ErrTooLarge
|
||||
}
|
||||
|
||||
return v, nil
|
||||
}
|
||||
|
||||
// format writes a result for a person to read. A whole number of
|
||||
// ordinary size is written exactly, digit for digit; anything else goes
|
||||
// through float64, whose shortest round-trip form is free of the noise
|
||||
// (0.30000000000000004) that printing a binary fraction to a fixed
|
||||
// precision produces.
|
||||
func format(v constant.Value) (string, error) {
|
||||
f, _ := constant.Float64Val(v)
|
||||
if math.IsInf(f, 0) || math.IsNaN(f) {
|
||||
return "", ErrTooLarge
|
||||
}
|
||||
|
||||
abs := math.Abs(f)
|
||||
|
||||
if i := constant.ToInt(v); i.Kind() == constant.Int && abs < plainUpper {
|
||||
return i.ExactString(), nil
|
||||
}
|
||||
|
||||
if abs >= plainUpper || abs < plainLower {
|
||||
return strconv.FormatFloat(f, 'g', -1, 64), nil
|
||||
}
|
||||
|
||||
return strconv.FormatFloat(f, 'f', -1, 64), nil
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package calc_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"sneak.berlin/go/simplexcalc/internal/calc"
|
||||
)
|
||||
|
||||
// TestEvaluate covers the two examples the bot was specified with, and
|
||||
// the arithmetic around them.
|
||||
func TestEvaluate(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := map[string]string{
|
||||
// The specification's own examples.
|
||||
"2 + 2": "4",
|
||||
"5 * 5/2": "12.5",
|
||||
|
||||
"2+2": "4",
|
||||
" 7 - 10 \n": "-3",
|
||||
"-3 * 2": "-6",
|
||||
"+4": "4",
|
||||
"-(-4)": "4",
|
||||
"(1 + 2) * 3": "9",
|
||||
"1 + 2 * 3": "7",
|
||||
"((2))": "2",
|
||||
"8 / 2 / 2": "2",
|
||||
"10 - 2 - 3": "5",
|
||||
"7 / 2": "3.5",
|
||||
"25/2": "12.5",
|
||||
"1 / 3": "0.3333333333333333",
|
||||
"2 / 3": "0.6666666666666666",
|
||||
"0.1 + 0.2": "0.3",
|
||||
"1.5 * 2": "3",
|
||||
"2.50 * 2": "5",
|
||||
".5 + .5": "1",
|
||||
"3. * 2": "6",
|
||||
"1e3 + 1": "1001",
|
||||
"2.5e-1": "0.25",
|
||||
"010 + 1": "11",
|
||||
"-0": "0",
|
||||
"0 / 5": "0",
|
||||
// Exact: a float64 would print 99999999980000000000.
|
||||
"9999999999 * 9999999999": "99999999980000000001",
|
||||
"1e21": "1e+21",
|
||||
"1e20": "100000000000000000000",
|
||||
"1 / 1e7": "1e-07",
|
||||
"1 / 1e6": "0.000001",
|
||||
"1234567.5": "1234567.5",
|
||||
"-1 / 4": "-0.25",
|
||||
"1e300 * 1e8": "1e+308",
|
||||
}
|
||||
|
||||
for in, want := range cases {
|
||||
t.Run(in, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, err := calc.Evaluate(in)
|
||||
if err != nil {
|
||||
t.Fatalf("Evaluate(%q) failed: %v", in, err)
|
||||
}
|
||||
|
||||
if got != want {
|
||||
t.Errorf("Evaluate(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEvaluateRefuses covers what must be answered with an error rather
|
||||
// than a number, and never with a panic.
|
||||
func TestEvaluateRefuses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := map[string]error{
|
||||
"": calc.ErrNotArithmetic,
|
||||
" ": calc.ErrNotArithmetic,
|
||||
"hello": calc.ErrNotArithmetic,
|
||||
"/help": calc.ErrNotArithmetic,
|
||||
"2 +": calc.ErrNotArithmetic,
|
||||
"2 2": calc.ErrNotArithmetic,
|
||||
"2 + 2 =": calc.ErrNotArithmetic,
|
||||
"x + 1": calc.ErrNotArithmetic,
|
||||
"len(\"abc\")": calc.ErrNotArithmetic,
|
||||
"\"a\" + \"b\"": calc.ErrNotArithmetic,
|
||||
"'a' + 1": calc.ErrNotArithmetic,
|
||||
"2i * 2i": calc.ErrNotArithmetic,
|
||||
"0x10 + 1": calc.ErrNotArithmetic,
|
||||
"1_000 + 1": calc.ErrNotArithmetic,
|
||||
"7 % 2": calc.ErrNotArithmetic,
|
||||
"2 ^ 3": calc.ErrNotArithmetic,
|
||||
"1 << 10": calc.ErrNotArithmetic,
|
||||
"1 == 1": calc.ErrNotArithmetic,
|
||||
"!1": calc.ErrNotArithmetic,
|
||||
"func() int { return 1 }()": calc.ErrNotArithmetic,
|
||||
"1 / 0": calc.ErrDivisionByZero,
|
||||
"1 / (2 - 2)": calc.ErrDivisionByZero,
|
||||
"5 / 0.0": calc.ErrDivisionByZero,
|
||||
"1e400": calc.ErrTooLarge,
|
||||
"1e300 * 1e300": calc.ErrTooLarge,
|
||||
"1e999999999 * 1e999999999": calc.ErrTooLarge,
|
||||
"1 / 1e-400": calc.ErrTooLarge,
|
||||
}
|
||||
|
||||
for in, want := range cases {
|
||||
t.Run(in, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, err := calc.Evaluate(in)
|
||||
if !errors.Is(err, want) {
|
||||
t.Errorf("Evaluate(%q) = %q, %v; want error %v", in, got, err, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEvaluateCapsInput: the length cap is what bounds the work a
|
||||
// message can cause, so it must hold exactly at the boundary.
|
||||
func TestEvaluateCapsInput(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// "1+1+...+1" with the last term padded to land exactly on the cap.
|
||||
longest := strings.Repeat("1+", calc.MaxInputLength/2-1) + "10"
|
||||
if len(longest) != calc.MaxInputLength {
|
||||
t.Fatalf("test setup: expression is %d bytes, want %d",
|
||||
len(longest), calc.MaxInputLength)
|
||||
}
|
||||
|
||||
got, err := calc.Evaluate(longest)
|
||||
if err != nil {
|
||||
t.Fatalf("an expression of exactly MaxInputLength bytes was refused: %v", err)
|
||||
}
|
||||
|
||||
if want := "137"; got != want {
|
||||
t.Errorf("Evaluate(longest) = %q, want %q", got, want)
|
||||
}
|
||||
|
||||
_, err = calc.Evaluate(longest + "0")
|
||||
if !errors.Is(err, calc.ErrTooLong) {
|
||||
t.Errorf("an expression over MaxInputLength gave %v, want ErrTooLong", err)
|
||||
}
|
||||
|
||||
// Surrounding whitespace is not part of the expression.
|
||||
_, err = calc.Evaluate(" " + longest + "\n")
|
||||
if err != nil {
|
||||
t.Errorf("whitespace around a maximal expression counted against the cap: %v", err)
|
||||
}
|
||||
}
|
||||
+26
-279
@@ -3,34 +3,30 @@
|
||||
//
|
||||
// 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.
|
||||
// operator who writes DEBUG=yes has said something specific that this
|
||||
// program does not understand, and starting anyway with the default
|
||||
// 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.
|
||||
// broken deployment takes one restart to diagnose rather than several.
|
||||
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`
|
||||
// `DEBUG=true`
|
||||
// (without the backticks, of course)
|
||||
_ "github.com/joho/godotenv/autoload"
|
||||
)
|
||||
@@ -38,54 +34,12 @@ import (
|
||||
// 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"
|
||||
EnvDataDir = "DATA_DIR"
|
||||
EnvDebug = "DEBUG"
|
||||
)
|
||||
|
||||
// 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
|
||||
)
|
||||
// DefaultDataDir applies when DATA_DIR is absent.
|
||||
const DefaultDataDir = "./data"
|
||||
|
||||
// ErrInvalidConfig is the sentinel every configuration failure wraps,
|
||||
// so callers can distinguish "the operator got it wrong" from "the
|
||||
@@ -95,43 +49,12 @@ 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.
|
||||
// wrong, and it is before anything starts.
|
||||
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
|
||||
// DataDir holds the SimpleX Chat database: the bot's profile, its
|
||||
// address and its contacts. Losing it loses the address.
|
||||
DataDir string
|
||||
Debug bool
|
||||
}
|
||||
|
||||
// loader parses one environment into a Config, accumulating every
|
||||
@@ -166,32 +89,10 @@ func (l *loader) str(key, def string) string {
|
||||
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.
|
||||
// parse failure on purpose: guessing at it is how a setting ends up the
|
||||
// opposite of what the operator meant.
|
||||
func (l *loader) boolean(key string, def bool) bool {
|
||||
s, ok := l.raw(key)
|
||||
if !ok {
|
||||
@@ -208,76 +109,10 @@ func (l *loader) boolean(key string, def bool) bool {
|
||||
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) {
|
||||
// New parses and validates the environment. An error here aborts
|
||||
// startup before the chat client is launched, so there is no partially
|
||||
// configured running state to reason about.
|
||||
func New() (*Config, error) {
|
||||
v := viper.New()
|
||||
v.AutomaticEnv()
|
||||
|
||||
@@ -289,34 +124,10 @@ func New(lc fx.Lifecycle, _ Params) (*Config, error) {
|
||||
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)
|
||||
c := &Config{
|
||||
DataDir: l.str(EnvDataDir, DefaultDataDir),
|
||||
Debug: l.boolean(EnvDebug, false),
|
||||
}
|
||||
|
||||
if len(l.errs) > 0 {
|
||||
return nil, errors.Join(l.errs...)
|
||||
@@ -324,67 +135,3 @@ func load(v *viper.Viper) (*Config, error) {
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
+37
-198
@@ -2,20 +2,12 @@ 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 {
|
||||
@@ -33,63 +25,40 @@ func env(kv map[string]string) *viper.Viper {
|
||||
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 {
|
||||
for name, kv := range map[string]map[string]string{
|
||||
"unset": nil,
|
||||
"whitespace only": {config.EnvDataDir: " ", config.EnvDebug: " "},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
c, err := config.Load(env(kv))
|
||||
if err != nil {
|
||||
t.Fatalf("absent values must be valid, got: %v", err)
|
||||
}
|
||||
|
||||
if c.DataDir != config.DefaultDataDir {
|
||||
t.Errorf("DataDir = %q, want %q", c.DataDir, config.DefaultDataDir)
|
||||
}
|
||||
|
||||
if c.Debug {
|
||||
t.Error("Debug must default off")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestSetButUnparseableAborts is the central contract of this package:
|
||||
// a value an operator plausibly types, and that does not parse, fails
|
||||
// startup rather than being replaced by the default.
|
||||
func TestSetButUnparseableAborts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, raw := range []string{"yes", "on", "enabled", "2"} {
|
||||
t.Run(raw, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
c, err := config.Load(env(map[string]string{config.EnvDebug: raw}))
|
||||
if err == nil {
|
||||
t.Fatalf("wanted a startup failure, got a Config: %+v", c)
|
||||
}
|
||||
@@ -105,153 +74,23 @@ func TestSetButUnparseableAborts(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 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",
|
||||
config.EnvDebug: "true",
|
||||
}))
|
||||
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)
|
||||
if c.DataDir != "/var/lib/example" {
|
||||
t.Errorf("DataDir = %q, want /var/lib/example", c.DataDir)
|
||||
}
|
||||
|
||||
if !c.Debug {
|
||||
t.Error("Debug = false, want true")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,7 +11,3 @@ package config
|
||||
//
|
||||
//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
|
||||
|
||||
@@ -1,29 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,211 +0,0 @@
|
||||
// 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
|
||||
}
|
||||
@@ -1,219 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,199 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,106 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,15 +0,0 @@
|
||||
-- 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);
|
||||
@@ -1,15 +0,0 @@
|
||||
-- 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);
|
||||
@@ -1,43 +0,0 @@
|
||||
// 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
|
||||
}
|
||||
@@ -1,13 +0,0 @@
|
||||
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")
|
||||
)
|
||||
@@ -1,123 +0,0 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -1,77 +0,0 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -1,114 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -8,43 +8,15 @@ 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
|
||||
// New returns a JSON logger writing to out, at debug level when debug is
|
||||
// set and at info level otherwise.
|
||||
func New(out io.Writer, debug bool) *slog.Logger {
|
||||
level := slog.LevelInfo
|
||||
if debug {
|
||||
level = slog.LevelDebug
|
||||
}
|
||||
|
||||
// replaceAttr simplifies the source attribute to "file.go:line".
|
||||
@@ -60,33 +32,9 @@ func New(_ fx.Lifecycle, params Params) (*Logger, error) {
|
||||
return a
|
||||
}
|
||||
|
||||
handler := slog.NewJSONHandler(out, &slog.HandlerOptions{
|
||||
Level: l.level,
|
||||
return slog.New(slog.NewJSONHandler(out, &slog.HandlerOptions{
|
||||
Level: 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,
|
||||
)
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -1,46 +0,0 @@
|
||||
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"
|
||||
@@ -1,116 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,163 +0,0 @@
|
||||
// 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
|
||||
}
|
||||
@@ -1,314 +0,0 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -1,59 +0,0 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,63 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,67 +0,0 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,13 +0,0 @@
|
||||
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")
|
||||
)
|
||||
@@ -1,240 +0,0 @@
|
||||
// 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
|
||||
}
|
||||
@@ -1,197 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,55 +0,0 @@
|
||||
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),
|
||||
}
|
||||
}
|
||||
@@ -1,87 +0,0 @@
|
||||
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)
|
||||
},
|
||||
))
|
||||
}
|
||||
@@ -1,180 +0,0 @@
|
||||
// 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
|
||||
}
|
||||
@@ -1,560 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,48 +0,0 @@
|
||||
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,149 @@
|
||||
package simplex
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os/exec"
|
||||
"strconv"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Binary is the chat client's executable, looked up on PATH.
|
||||
const Binary = "simplex-chat"
|
||||
|
||||
// stopGrace is how long the chat client gets to exit after SIGTERM
|
||||
// before it is killed.
|
||||
const stopGrace = 10 * time.Second
|
||||
|
||||
// maxLine bounds how much of one unterminated output line is held
|
||||
// before it is logged anyway.
|
||||
const maxLine = 64 << 10
|
||||
|
||||
// CLI is a running chat client process.
|
||||
type CLI struct {
|
||||
done chan struct{}
|
||||
err error // how the process ended; valid once done is closed
|
||||
}
|
||||
|
||||
// StartCLI launches the chat client with its database at dbPrefix,
|
||||
// serving its API on localhost at port. On the first start, with no
|
||||
// database yet, the client creates a bot profile named displayName;
|
||||
// every later start uses the profile already in the database.
|
||||
//
|
||||
// Cancelling ctx stops the client: SIGTERM, then SIGKILL if it has not
|
||||
// exited after stopGrace. Its output is logged line by line, so the
|
||||
// process emits one log format.
|
||||
func StartCLI(
|
||||
ctx context.Context, log *slog.Logger, dbPrefix, displayName string, port int,
|
||||
) (*CLI, error) {
|
||||
// No shell is involved: each argument reaches the client as one
|
||||
// argv entry, whatever it contains.
|
||||
//nolint:gosec // G204: the arguments are this program's own settings.
|
||||
cmd := exec.CommandContext(ctx, Binary,
|
||||
"--database", dbPrefix,
|
||||
"--chat-server-port", strconv.Itoa(port),
|
||||
"--create-bot-display-name", displayName,
|
||||
// Confirms the database migrations a newer client brings,
|
||||
// which it would otherwise wait to have confirmed on a
|
||||
// terminal that nobody is at.
|
||||
"--yes-migrate",
|
||||
)
|
||||
|
||||
cmd.Cancel = func() error {
|
||||
return cmd.Process.Signal(syscall.SIGTERM)
|
||||
}
|
||||
cmd.WaitDelay = stopGrace
|
||||
|
||||
out := &lineLogger{log: log}
|
||||
cmd.Stdout = out
|
||||
cmd.Stderr = out
|
||||
|
||||
err := cmd.Start()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("starting %s: %w", Binary, err)
|
||||
}
|
||||
|
||||
cli := &CLI{done: make(chan struct{})}
|
||||
|
||||
go func() {
|
||||
defer close(cli.done)
|
||||
|
||||
cli.err = cmd.Wait()
|
||||
|
||||
out.flush()
|
||||
}()
|
||||
|
||||
return cli, nil
|
||||
}
|
||||
|
||||
// Done is closed when the process has exited; Err then says how.
|
||||
func (c *CLI) Done() <-chan struct{} {
|
||||
return c.done
|
||||
}
|
||||
|
||||
// Err returns how the process ended. Call it only after Done is closed.
|
||||
func (c *CLI) Err() error {
|
||||
return c.err
|
||||
}
|
||||
|
||||
// lineLogger is the chat client's stdout and stderr.
|
||||
type lineLogger struct {
|
||||
log *slog.Logger
|
||||
|
||||
// mu covers the final flush after Wait, which can overlap a last
|
||||
// Write when Wait gave up on the output after stopGrace.
|
||||
mu sync.Mutex
|
||||
buf []byte
|
||||
}
|
||||
|
||||
func (l *lineLogger) Write(p []byte) (int, error) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
l.buf = append(l.buf, p...)
|
||||
|
||||
for {
|
||||
line, rest, found := bytes.Cut(l.buf, []byte{'\n'})
|
||||
if !found {
|
||||
break
|
||||
}
|
||||
|
||||
l.emit(line)
|
||||
l.buf = rest
|
||||
}
|
||||
|
||||
if len(l.buf) > maxLine {
|
||||
l.flushLocked()
|
||||
}
|
||||
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// flush logs what is left of an unterminated last line.
|
||||
func (l *lineLogger) flush() {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
l.flushLocked()
|
||||
}
|
||||
|
||||
func (l *lineLogger) flushLocked() {
|
||||
if len(l.buf) > 0 {
|
||||
l.emit(l.buf)
|
||||
}
|
||||
|
||||
l.buf = nil
|
||||
}
|
||||
|
||||
func (l *lineLogger) emit(line []byte) {
|
||||
line = bytes.TrimSpace(line)
|
||||
if len(line) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
l.log.Info("simplex-chat", "output", string(line))
|
||||
}
|
||||
@@ -0,0 +1,359 @@
|
||||
// Package simplex runs the SimpleX Chat command-line client and talks
|
||||
// to it over its WebSocket API.
|
||||
//
|
||||
// The protocol, documented in the simplex-chat repository under bots/:
|
||||
// a command goes out as {"corrId": "...", "cmd": "..."}, and the client
|
||||
// answers it with {"corrId": "...", "resp": {...}} carrying the same id.
|
||||
// Everything it sends without a corrId is an event. The API has no
|
||||
// authentication; the client binds it to localhost only.
|
||||
package simplex
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// maxMessageSize bounds one message from the chat client. The largest
|
||||
// thing it sends is a record carrying a contact's profile picture, well
|
||||
// under this.
|
||||
const maxMessageSize = 16 << 20
|
||||
|
||||
var (
|
||||
// ErrClosed is returned by Command once the connection has ended.
|
||||
ErrClosed = errors.New("connection to the chat client closed")
|
||||
|
||||
errUnexpected = errors.New("unexpected response")
|
||||
errCommand = errors.New("command failed")
|
||||
)
|
||||
|
||||
// EventHandler receives each event the chat client sends. Events are
|
||||
// delivered one at a time, on the goroutine that also delivers command
|
||||
// responses: a handler may Send, but must never wait on Command, whose
|
||||
// response could then never arrive.
|
||||
type EventHandler func(c *Client, ev Event)
|
||||
|
||||
// Client is a connection to the chat client's WebSocket API.
|
||||
type Client struct {
|
||||
conn *websocket.Conn
|
||||
log *slog.Logger
|
||||
onEvent EventHandler
|
||||
|
||||
// writeMu serialises writes: the connection allows one writer at
|
||||
// a time, and Command and Send are called from different
|
||||
// goroutines.
|
||||
writeMu sync.Mutex
|
||||
|
||||
mu sync.Mutex
|
||||
lastID uint64
|
||||
waiting map[string]chan Event
|
||||
|
||||
done chan struct{}
|
||||
err error // why the read loop ended; valid once done is closed
|
||||
}
|
||||
|
||||
// Dial connects to the chat client's API at url and starts reading from
|
||||
// it, passing every event to onEvent.
|
||||
func Dial(
|
||||
ctx context.Context, url string, log *slog.Logger, onEvent EventHandler,
|
||||
) (*Client, error) {
|
||||
conn, resp, err := websocket.DefaultDialer.DialContext(ctx, url, nil)
|
||||
if resp != nil {
|
||||
_ = resp.Body.Close()
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connecting to %s: %w", url, err)
|
||||
}
|
||||
|
||||
conn.SetReadLimit(maxMessageSize)
|
||||
|
||||
c := &Client{
|
||||
conn: conn,
|
||||
log: log,
|
||||
onEvent: onEvent,
|
||||
waiting: make(map[string]chan Event),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
go c.read()
|
||||
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// Done is closed when the connection ends; Err then says why.
|
||||
func (c *Client) Done() <-chan struct{} {
|
||||
return c.done
|
||||
}
|
||||
|
||||
// Err returns why the connection ended. Call it only after Done is
|
||||
// closed.
|
||||
func (c *Client) Err() error {
|
||||
return c.err
|
||||
}
|
||||
|
||||
// Close ends the connection.
|
||||
func (c *Client) Close() error {
|
||||
err := c.conn.Close()
|
||||
if err != nil {
|
||||
return fmt.Errorf("closing connection: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ActiveUser returns the chat client's active user profile.
|
||||
func (c *Client) ActiveUser(ctx context.Context) (User, error) {
|
||||
var r struct {
|
||||
User User `json:"user"`
|
||||
}
|
||||
|
||||
err := c.command(ctx, cmdShowActiveUser, TypeActiveUser, &r)
|
||||
|
||||
return r.User, err
|
||||
}
|
||||
|
||||
// Address returns the user's long-term contact address, and false if
|
||||
// the user has none.
|
||||
func (c *Client) Address(ctx context.Context, userID int64) (ConnLink, bool, error) {
|
||||
//nolint:tagliatelle // the chat client's wire format.
|
||||
var r struct {
|
||||
ContactLink struct {
|
||||
ConnLinkContact ConnLink `json:"connLinkContact"`
|
||||
} `json:"contactLink"`
|
||||
}
|
||||
|
||||
err := c.command(ctx, cmdShowAddress(userID), TypeUserContactLink, &r)
|
||||
|
||||
var cerr *CommandError
|
||||
if errors.As(err, &cerr) && cerr.Detail == "userContactLinkNotFound" {
|
||||
return ConnLink{}, false, nil
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return ConnLink{}, false, err
|
||||
}
|
||||
|
||||
return r.ContactLink.ConnLinkContact, true, nil
|
||||
}
|
||||
|
||||
// CreateAddress creates the user's long-term contact address.
|
||||
func (c *Client) CreateAddress(ctx context.Context, userID int64) (ConnLink, error) {
|
||||
//nolint:tagliatelle // the chat client's wire format.
|
||||
var r struct {
|
||||
ConnLinkContact ConnLink `json:"connLinkContact"`
|
||||
}
|
||||
|
||||
err := c.command(ctx, cmdCreateAddress(userID), TypeUserContactLinkCreated, &r)
|
||||
|
||||
return r.ConnLinkContact, err
|
||||
}
|
||||
|
||||
// SetAddressSettings replaces the settings of the user's address.
|
||||
func (c *Client) SetAddressSettings(
|
||||
ctx context.Context, userID int64, s AddressSettings,
|
||||
) error {
|
||||
cmd, err := cmdSetAddressSettings(userID, s)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return c.command(ctx, cmd, TypeUserContactLinkUpdated, nil)
|
||||
}
|
||||
|
||||
// SendText sends a text message to a contact, as a reply to the message
|
||||
// quotedItemID (0 for none). It does not wait for the chat client to
|
||||
// accept it; a failure is logged when the client's answer arrives.
|
||||
func (c *Client) SendText(contactID, quotedItemID int64, text string) error {
|
||||
cmd, err := cmdSendText(contactID, quotedItemID, text)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
id := c.nextID()
|
||||
c.mu.Unlock()
|
||||
|
||||
return c.write(id, cmd)
|
||||
}
|
||||
|
||||
// CommandError is a command the chat client refused. Type and Detail
|
||||
// are the discriminators of its chatError record, such as "errorStore"
|
||||
// and "userContactLinkNotFound".
|
||||
type CommandError struct {
|
||||
Type string
|
||||
Detail string
|
||||
}
|
||||
|
||||
func (e *CommandError) Error() string {
|
||||
return fmt.Sprintf("%s: %s/%s", errCommand, e.Type, e.Detail)
|
||||
}
|
||||
|
||||
func (e *CommandError) Unwrap() error {
|
||||
return errCommand
|
||||
}
|
||||
|
||||
// command sends cmd, waits for its response, and, if the response has
|
||||
// type want, decodes it into out (unless out is nil).
|
||||
func (c *Client) command(ctx context.Context, cmd, want string, out any) error {
|
||||
ch := make(chan Event, 1)
|
||||
|
||||
c.mu.Lock()
|
||||
id := c.nextID()
|
||||
c.waiting[id] = ch
|
||||
c.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
c.mu.Lock()
|
||||
delete(c.waiting, id)
|
||||
c.mu.Unlock()
|
||||
}()
|
||||
|
||||
err := c.write(id, cmd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var ev Event
|
||||
|
||||
select {
|
||||
case ev = <-ch:
|
||||
case <-c.done:
|
||||
return fmt.Errorf("%w: %w", ErrClosed, c.err)
|
||||
case <-ctx.Done():
|
||||
return fmt.Errorf("waiting for a response to %q: %w", cmdName(cmd), ctx.Err())
|
||||
}
|
||||
|
||||
switch ev.Type {
|
||||
case want:
|
||||
if out == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return ev.Decode(out)
|
||||
case TypeChatCmdError:
|
||||
return commandError(ev)
|
||||
default:
|
||||
return fmt.Errorf("%w to %q: %s", errUnexpected, cmdName(cmd), ev.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// nextID returns a fresh correlation id. The caller holds c.mu.
|
||||
func (c *Client) nextID() string {
|
||||
c.lastID++
|
||||
|
||||
return strconv.FormatUint(c.lastID, 10)
|
||||
}
|
||||
|
||||
func (c *Client) write(id, cmd string) error {
|
||||
c.writeMu.Lock()
|
||||
defer c.writeMu.Unlock()
|
||||
|
||||
err := c.conn.WriteJSON(command{CorrID: id, Cmd: cmd})
|
||||
if err != nil {
|
||||
return fmt.Errorf("sending %q: %w", cmdName(cmd), err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// read is the only reader of the connection. It runs until the
|
||||
// connection fails or is closed.
|
||||
func (c *Client) read() {
|
||||
defer close(c.done)
|
||||
|
||||
for {
|
||||
_, data, err := c.conn.ReadMessage()
|
||||
if err != nil {
|
||||
c.err = err
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
c.dispatch(data)
|
||||
}
|
||||
}
|
||||
|
||||
// dispatch routes one message: a response to whoever waits for it, an
|
||||
// event to the handler. A message that does not parse is logged and
|
||||
// skipped rather than ending the connection; the API documentation
|
||||
// warns that records change between releases.
|
||||
func (c *Client) dispatch(data []byte) {
|
||||
var env envelope
|
||||
|
||||
err := json.Unmarshal(data, &env)
|
||||
if err != nil {
|
||||
c.log.Warn("undecodable message from the chat client", "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
var head tagged
|
||||
|
||||
err = json.Unmarshal(env.Resp, &head)
|
||||
if err != nil {
|
||||
c.log.Warn("undecodable record from the chat client", "error", err)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
ev := Event{Type: head.Type, raw: env.Resp}
|
||||
|
||||
if env.CorrID == "" {
|
||||
c.onEvent(c, ev)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
ch, ok := c.waiting[env.CorrID]
|
||||
c.mu.Unlock()
|
||||
|
||||
if ok {
|
||||
ch <- ev
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
// The response to a SendText: nobody waits for it, so a failure
|
||||
// is reported here or nowhere.
|
||||
if ev.Type == TypeChatCmdError {
|
||||
c.log.Warn("sending a message failed", "error", commandError(ev))
|
||||
}
|
||||
}
|
||||
|
||||
func commandError(ev Event) error {
|
||||
var r cmdError
|
||||
|
||||
err := ev.Decode(&r)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
e := &CommandError{Type: r.ChatError.Type}
|
||||
|
||||
for _, detail := range []*tagged{
|
||||
r.ChatError.ErrorType, r.ChatError.StoreError, r.ChatError.AgentError,
|
||||
} {
|
||||
if detail != nil {
|
||||
e.Detail = detail.Type
|
||||
}
|
||||
}
|
||||
|
||||
return e
|
||||
}
|
||||
|
||||
// cmdName is a command without its arguments, for error messages: the
|
||||
// arguments of /_send are a message someone wrote.
|
||||
func cmdName(cmd string) string {
|
||||
name, _, _ := strings.Cut(cmd, " ")
|
||||
|
||||
return name
|
||||
}
|
||||
@@ -0,0 +1,320 @@
|
||||
package simplex_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"sneak.berlin/go/simplexcalc/internal/simplex"
|
||||
)
|
||||
|
||||
const testTimeout = 5 * time.Second
|
||||
|
||||
// cmdUser is the command that reads the active user profile.
|
||||
const cmdUser = "/user"
|
||||
|
||||
// Records as the chat client sends them, with fields this package does
|
||||
// not read left in, since ignoring those is part of the contract.
|
||||
const (
|
||||
activeUser = `{"type":"activeUser","user":{"userId":1,"agentUserId":1,
|
||||
"profile":{"profileId":1,"displayName":"calc","fullName":"",
|
||||
"peerType":"bot","localAlias":""},"activeUser":true}}`
|
||||
|
||||
addressNotFound = `{"type":"chatCmdError","chatError":{"type":"errorStore",
|
||||
"storeError":{"type":"userContactLinkNotFound"}}}`
|
||||
|
||||
addressCreated = `{"type":"userContactLinkCreated","user":{"userId":1},
|
||||
"connLinkContact":{"connFullLink":"simplex:/contact#/?v=2-7&smp=x",
|
||||
"connShortLink":"https://smp.example/a#key"}}`
|
||||
|
||||
addressUpdated = `{"type":"userContactLinkUpdated","user":{"userId":1},
|
||||
"contactLink":{"userContactLinkId":1}}`
|
||||
|
||||
noActiveUser = `{"type":"chatCmdError","chatError":{"type":"error",
|
||||
"errorType":{"type":"noActiveUser"}}}`
|
||||
|
||||
contactConnected = `{"type":"contactConnected","user":{"userId":1},
|
||||
"contact":{"contactId":3,"localDisplayName":"alice"}}`
|
||||
)
|
||||
|
||||
// fakeChat stands in for the chat client's API. It answers each command
|
||||
// with the record in replies under the command's first word, stays
|
||||
// silent for a command it has no record for, and reports every command
|
||||
// it receives on got.
|
||||
type fakeChat struct {
|
||||
replies map[string]string
|
||||
got chan string
|
||||
|
||||
mu sync.Mutex
|
||||
conn *websocket.Conn
|
||||
up chan struct{}
|
||||
}
|
||||
|
||||
func newFakeChat(t *testing.T, replies map[string]string) (*fakeChat, string) {
|
||||
t.Helper()
|
||||
|
||||
f := &fakeChat{
|
||||
replies: replies,
|
||||
got: make(chan string, 16),
|
||||
up: make(chan struct{}),
|
||||
}
|
||||
|
||||
srv := httptest.NewServer(f)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
return f, "ws" + strings.TrimPrefix(srv.URL, "http")
|
||||
}
|
||||
|
||||
func (f *fakeChat) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := (&websocket.Upgrader{}).Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
f.mu.Lock()
|
||||
f.conn = conn
|
||||
f.mu.Unlock()
|
||||
close(f.up)
|
||||
|
||||
for {
|
||||
//nolint:tagliatelle // the chat client's wire format.
|
||||
var cmd struct {
|
||||
CorrID string `json:"corrId"`
|
||||
Cmd string `json:"cmd"`
|
||||
}
|
||||
|
||||
err = conn.ReadJSON(&cmd)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
f.got <- cmd.Cmd
|
||||
|
||||
name, _, _ := strings.Cut(cmd.Cmd, " ")
|
||||
if resp, ok := f.replies[name]; ok {
|
||||
f.send(cmd.CorrID, resp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// send writes one record, as a response to corrID, or as an event when
|
||||
// corrID is empty.
|
||||
func (f *fakeChat) send(corrID, resp string) {
|
||||
msg := map[string]any{"resp": json.RawMessage(resp)}
|
||||
if corrID != "" {
|
||||
msg["corrId"] = corrID
|
||||
}
|
||||
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
||||
_ = f.conn.WriteJSON(msg)
|
||||
}
|
||||
|
||||
func (f *fakeChat) hangUp() {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
||||
_ = f.conn.Close()
|
||||
}
|
||||
|
||||
func (f *fakeChat) next(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
select {
|
||||
case cmd := <-f.got:
|
||||
return cmd
|
||||
case <-time.After(testTimeout):
|
||||
t.Fatal("the client sent no command")
|
||||
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func dial(
|
||||
t *testing.T, url string, onEvent simplex.EventHandler,
|
||||
) (*simplex.Client, context.Context) {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(t.Context(), testTimeout)
|
||||
t.Cleanup(cancel)
|
||||
|
||||
if onEvent == nil {
|
||||
onEvent = func(*simplex.Client, simplex.Event) {}
|
||||
}
|
||||
|
||||
c, err := simplex.Dial(ctx, url, slog.New(slog.DiscardHandler), onEvent)
|
||||
if err != nil {
|
||||
t.Fatalf("Dial: %v", err)
|
||||
}
|
||||
|
||||
t.Cleanup(func() { _ = c.Close() })
|
||||
|
||||
return c, ctx
|
||||
}
|
||||
|
||||
// TestAddressSetup walks the calls the bot makes on its first start,
|
||||
// and checks the exact commands that reach the chat client.
|
||||
func TestAddressSetup(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
f, url := newFakeChat(t, map[string]string{
|
||||
cmdUser: activeUser,
|
||||
"/_show_address": addressNotFound,
|
||||
"/_address": addressCreated,
|
||||
"/_address_settings": addressUpdated,
|
||||
})
|
||||
c, ctx := dial(t, url, nil)
|
||||
|
||||
user, err := c.ActiveUser(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ActiveUser: %v", err)
|
||||
}
|
||||
|
||||
if user.UserID != 1 || user.Profile.DisplayName != "calc" {
|
||||
t.Errorf("ActiveUser = %+v, want user 1 named calc", user)
|
||||
}
|
||||
|
||||
_, ok, err := c.Address(ctx, 1)
|
||||
if err != nil || ok {
|
||||
t.Fatalf("Address = %v, %v; want no address and no error", ok, err)
|
||||
}
|
||||
|
||||
link, err := c.CreateAddress(ctx, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateAddress: %v", err)
|
||||
}
|
||||
|
||||
if link.ShortLink != "https://smp.example/a#key" ||
|
||||
link.FullLink != "simplex:/contact#/?v=2-7&smp=x" {
|
||||
t.Errorf("CreateAddress = %+v", link)
|
||||
}
|
||||
|
||||
err = c.SetAddressSettings(ctx, 1, simplex.AddressSettings{
|
||||
AutoAccept: &simplex.AutoAccept{},
|
||||
AutoReply: &simplex.MsgContent{Type: "text", Text: "hi"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SetAddressSettings: %v", err)
|
||||
}
|
||||
|
||||
for _, want := range []string{
|
||||
cmdUser,
|
||||
"/_show_address 1",
|
||||
"/_address 1",
|
||||
`/_address_settings 1 {"businessAddress":false,` +
|
||||
`"autoAccept":{"acceptIncognito":false},` +
|
||||
`"autoReply":{"type":"text","text":"hi"}}`,
|
||||
} {
|
||||
if got := f.next(t); got != want {
|
||||
t.Errorf("command = %s\nwant %s", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRefusedCommand: a command the chat client refuses is an error
|
||||
// that names the reason.
|
||||
func TestRefusedCommand(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, url := newFakeChat(t, map[string]string{cmdUser: noActiveUser})
|
||||
c, ctx := dial(t, url, nil)
|
||||
|
||||
_, err := c.ActiveUser(ctx)
|
||||
|
||||
var cerr *simplex.CommandError
|
||||
if !errors.As(err, &cerr) || cerr.Detail != "noActiveUser" {
|
||||
t.Errorf("ActiveUser error = %v, want a CommandError for noActiveUser", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUnexpectedResponse: a response of the wrong type is an error, not
|
||||
// a zero value.
|
||||
func TestUnexpectedResponse(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, url := newFakeChat(t, map[string]string{cmdUser: addressUpdated})
|
||||
c, ctx := dial(t, url, nil)
|
||||
|
||||
_, err := c.ActiveUser(ctx)
|
||||
if err == nil {
|
||||
t.Error("ActiveUser accepted a userContactLinkUpdated response")
|
||||
}
|
||||
}
|
||||
|
||||
// TestEventsAndReplies: an event reaches the handler, and a reply sent
|
||||
// from inside the handler reaches the chat client.
|
||||
func TestEventsAndReplies(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
f, url := newFakeChat(t, nil)
|
||||
|
||||
seen := make(chan simplex.Event, 1)
|
||||
|
||||
dial(t, url, func(c *simplex.Client, ev simplex.Event) {
|
||||
seen <- ev
|
||||
|
||||
err := c.SendText(3, 7, `4 "exactly"`)
|
||||
if err != nil {
|
||||
t.Errorf("SendText: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
<-f.up
|
||||
f.send("", contactConnected)
|
||||
|
||||
select {
|
||||
case ev := <-seen:
|
||||
var r simplex.ContactConnected
|
||||
|
||||
err := ev.Decode(&r)
|
||||
if err != nil || ev.Type != simplex.TypeContactConnected ||
|
||||
r.Contact.ContactID != 3 {
|
||||
t.Errorf("event = %s %+v (%v), want contactConnected for contact 3",
|
||||
ev.Type, r, err)
|
||||
}
|
||||
case <-time.After(testTimeout):
|
||||
t.Fatal("the event never reached the handler")
|
||||
}
|
||||
|
||||
want := `/_send @3 json [{"quotedItemId":7,` +
|
||||
`"msgContent":{"type":"text","text":"4 \"exactly\""},"mentions":{}}]`
|
||||
if got := f.next(t); got != want {
|
||||
t.Errorf("command = %s\nwant %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestConnectionLoss: when the chat client goes away, Done closes and a
|
||||
// command fails instead of waiting for an answer that cannot come.
|
||||
func TestConnectionLoss(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
f, url := newFakeChat(t, nil)
|
||||
c, ctx := dial(t, url, nil)
|
||||
|
||||
<-f.up
|
||||
f.hangUp()
|
||||
|
||||
select {
|
||||
case <-c.Done():
|
||||
case <-time.After(testTimeout):
|
||||
t.Fatal("Done did not close after the connection ended")
|
||||
}
|
||||
|
||||
if c.Err() == nil {
|
||||
t.Error("Err is nil after the connection ended")
|
||||
}
|
||||
|
||||
_, err := c.ActiveUser(ctx)
|
||||
if err == nil {
|
||||
t.Error("a command on a closed connection succeeded")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
package simplex
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Response and event types this package and the bot act on. The chat
|
||||
// client sends many more; every other type is ignored, as its API
|
||||
// documentation requires of clients.
|
||||
const (
|
||||
TypeActiveUser = "activeUser"
|
||||
TypeUserContactLink = "userContactLink"
|
||||
TypeUserContactLinkCreated = "userContactLinkCreated"
|
||||
TypeUserContactLinkUpdated = "userContactLinkUpdated"
|
||||
TypeNewChatItems = "newChatItems"
|
||||
TypeContactConnected = "contactConnected"
|
||||
TypeChatCmdError = "chatCmdError"
|
||||
)
|
||||
|
||||
// Event is one message from the chat client: a response to a command,
|
||||
// or an event it sends unprompted. The protocol is a discriminated
|
||||
// union on "type"; the rest of the record is decoded on demand, into a
|
||||
// struct declaring only the fields the caller reads, so a record whose
|
||||
// other fields changed shape between releases still decodes.
|
||||
type Event struct {
|
||||
Type string
|
||||
raw json.RawMessage
|
||||
}
|
||||
|
||||
// Decode unmarshals the whole record into v.
|
||||
func (e Event) Decode(v any) error {
|
||||
err := json.Unmarshal(e.raw, v)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decoding %s: %w", e.Type, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Wire types, reduced to the fields this program uses. The field names
|
||||
// are the chat client's, hence camelCase in the tags.
|
||||
//
|
||||
//nolint:tagliatelle // the chat client's wire format, not ours to name.
|
||||
type (
|
||||
// User is the chat client's local user profile: the bot itself.
|
||||
User struct {
|
||||
UserID int64 `json:"userId"`
|
||||
Profile struct {
|
||||
DisplayName string `json:"displayName"`
|
||||
} `json:"profile"`
|
||||
}
|
||||
|
||||
// ConnLink is a SimpleX link. The short form is what people share;
|
||||
// the full form is what older clients understand.
|
||||
ConnLink struct {
|
||||
FullLink string `json:"connFullLink"`
|
||||
ShortLink string `json:"connShortLink,omitempty"`
|
||||
}
|
||||
|
||||
// Contact is a person connected to the bot.
|
||||
Contact struct {
|
||||
ContactID int64 `json:"contactId"`
|
||||
}
|
||||
|
||||
// NewChatItems is the record of a newChatItems event: messages
|
||||
// received, or sent from this profile elsewhere.
|
||||
NewChatItems struct {
|
||||
ChatItems []AChatItem `json:"chatItems"`
|
||||
}
|
||||
|
||||
// ContactConnected is the record of a contactConnected event.
|
||||
ContactConnected struct {
|
||||
Contact Contact `json:"contact"`
|
||||
}
|
||||
|
||||
// AChatItem is one message together with the chat it belongs to.
|
||||
AChatItem struct {
|
||||
ChatInfo struct {
|
||||
Type string `json:"type"`
|
||||
Contact *Contact `json:"contact,omitempty"`
|
||||
} `json:"chatInfo"`
|
||||
ChatItem struct {
|
||||
ChatDir tagged `json:"chatDir"`
|
||||
Meta struct {
|
||||
ItemID int64 `json:"itemId"`
|
||||
} `json:"meta"`
|
||||
Content struct {
|
||||
Type string `json:"type"`
|
||||
MsgContent *MsgContent `json:"msgContent,omitempty"`
|
||||
} `json:"content"`
|
||||
} `json:"chatItem"`
|
||||
}
|
||||
|
||||
// MsgContent is a message body. Only "text" is sent or read here.
|
||||
MsgContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
// AddressSettings configures the bot's long-term address.
|
||||
AddressSettings struct {
|
||||
BusinessAddress bool `json:"businessAddress"`
|
||||
AutoAccept *AutoAccept `json:"autoAccept,omitempty"`
|
||||
AutoReply *MsgContent `json:"autoReply,omitempty"`
|
||||
}
|
||||
|
||||
// AutoAccept makes the chat client accept every contact request
|
||||
// to the address itself, when present in AddressSettings.
|
||||
AutoAccept struct {
|
||||
AcceptIncognito bool `json:"acceptIncognito"`
|
||||
}
|
||||
|
||||
composedMessage struct {
|
||||
QuotedItemID int64 `json:"quotedItemId,omitempty"`
|
||||
MsgContent MsgContent `json:"msgContent"`
|
||||
Mentions map[string]int64 `json:"mentions"`
|
||||
}
|
||||
|
||||
envelope struct {
|
||||
CorrID string `json:"corrId,omitempty"`
|
||||
Resp json.RawMessage `json:"resp"`
|
||||
}
|
||||
|
||||
command struct {
|
||||
CorrID string `json:"corrId"`
|
||||
Cmd string `json:"cmd"`
|
||||
}
|
||||
|
||||
cmdError struct {
|
||||
ChatError struct {
|
||||
Type string `json:"type"`
|
||||
ErrorType *tagged `json:"errorType,omitempty"`
|
||||
StoreError *tagged `json:"storeError,omitempty"`
|
||||
AgentError *tagged `json:"agentError,omitempty"`
|
||||
} `json:"chatError"`
|
||||
}
|
||||
|
||||
tagged struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
)
|
||||
|
||||
// Message is a text message a contact sent to the bot.
|
||||
type Message struct {
|
||||
ContactID int64
|
||||
ItemID int64
|
||||
Text string
|
||||
}
|
||||
|
||||
// Message returns the text message a contact sent in a direct chat, and
|
||||
// false for anything else: group messages, files, the bot's own
|
||||
// messages, and the event items the client records in a chat.
|
||||
func (a AChatItem) Message() (Message, bool) {
|
||||
item := a.ChatItem
|
||||
|
||||
if a.ChatInfo.Type != "direct" || a.ChatInfo.Contact == nil ||
|
||||
item.ChatDir.Type != "directRcv" || item.Content.Type != "rcvMsgContent" ||
|
||||
item.Content.MsgContent == nil || item.Content.MsgContent.Type != "text" {
|
||||
return Message{}, false
|
||||
}
|
||||
|
||||
return Message{
|
||||
ContactID: a.ChatInfo.Contact.ContactID,
|
||||
ItemID: item.Meta.ItemID,
|
||||
Text: item.Content.MsgContent.Text,
|
||||
}, true
|
||||
}
|
||||
|
||||
// Command strings. Their syntax is documented per command in the
|
||||
// simplex-chat repository, bots/api/COMMANDS.md.
|
||||
|
||||
const cmdShowActiveUser = "/user"
|
||||
|
||||
func cmdShowAddress(userID int64) string {
|
||||
return "/_show_address " + strconv.FormatInt(userID, 10)
|
||||
}
|
||||
|
||||
func cmdCreateAddress(userID int64) string {
|
||||
return "/_address " + strconv.FormatInt(userID, 10)
|
||||
}
|
||||
|
||||
func cmdSetAddressSettings(userID int64, s AddressSettings) (string, error) {
|
||||
b, err := json.Marshal(s)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("encoding address settings: %w", err)
|
||||
}
|
||||
|
||||
return "/_address_settings " + strconv.FormatInt(userID, 10) + " " + string(b), nil
|
||||
}
|
||||
|
||||
func cmdSendText(contactID, quotedItemID int64, text string) (string, error) {
|
||||
b, err := json.Marshal([]composedMessage{{
|
||||
QuotedItemID: quotedItemID,
|
||||
MsgContent: MsgContent{Type: "text", Text: text},
|
||||
Mentions: map[string]int64{},
|
||||
}})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("encoding message: %w", err)
|
||||
}
|
||||
|
||||
return "/_send @" + strconv.FormatInt(contactID, 10) + " json " + string(b), nil
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package simplex_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"sneak.berlin/go/simplexcalc/internal/simplex"
|
||||
)
|
||||
|
||||
// newChatItems holds one message the bot must answer and four it must
|
||||
// not, each shaped as the chat client sends it.
|
||||
const newChatItems = `{"type":"newChatItems","user":{"userId":1},"chatItems":[
|
||||
{"chatInfo":{"type":"direct","contact":{"contactId":3,"localDisplayName":"alice"}},
|
||||
"chatItem":{"chatDir":{"type":"directRcv"},
|
||||
"meta":{"itemId":41,"itemText":"2 + 2","itemEdited":false},
|
||||
"content":{"type":"rcvMsgContent","msgContent":{"type":"text","text":"2 + 2"}},
|
||||
"mentions":{},"reactions":[]}},
|
||||
{"chatInfo":{"type":"direct","contact":{"contactId":3}},
|
||||
"chatItem":{"chatDir":{"type":"directSnd"},"meta":{"itemId":42},
|
||||
"content":{"type":"sndMsgContent","msgContent":{"type":"text","text":"4"}}}},
|
||||
{"chatInfo":{"type":"group","groupInfo":{"groupId":9}},
|
||||
"chatItem":{"chatDir":{"type":"groupRcv","groupMember":{"groupMemberId":5}},
|
||||
"meta":{"itemId":43},
|
||||
"content":{"type":"rcvMsgContent","msgContent":{"type":"text","text":"1 + 1"}}}},
|
||||
{"chatInfo":{"type":"direct","contact":{"contactId":3}},
|
||||
"chatItem":{"chatDir":{"type":"directRcv"},"meta":{"itemId":44},
|
||||
"content":{"type":"rcvMsgContent","msgContent":{"type":"file","text":"3 * 3"}}}},
|
||||
{"chatInfo":{"type":"direct","contact":{"contactId":3}},
|
||||
"chatItem":{"chatDir":{"type":"directRcv"},"meta":{"itemId":45},
|
||||
"content":{"type":"rcvDirectEvent","rcvDirectEvent":{"type":"contactDeleted"}}}}
|
||||
]}`
|
||||
|
||||
// TestMessage: only a text message a contact sent in a direct chat is a
|
||||
// message to answer. Answering the bot's own messages would loop.
|
||||
func TestMessage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var r simplex.NewChatItems
|
||||
|
||||
err := json.Unmarshal([]byte(newChatItems), &r)
|
||||
if err != nil {
|
||||
t.Fatalf("decoding: %v", err)
|
||||
}
|
||||
|
||||
var got []simplex.Message
|
||||
|
||||
for _, item := range r.ChatItems {
|
||||
if msg, ok := item.Message(); ok {
|
||||
got = append(got, msg)
|
||||
}
|
||||
}
|
||||
|
||||
want := simplex.Message{ContactID: 3, ItemID: 41, Text: "2 + 2"}
|
||||
if len(got) != 1 || got[0] != want {
|
||||
t.Errorf("messages = %+v, want only %+v", got, want)
|
||||
}
|
||||
}
|
||||
@@ -1,162 +0,0 @@
|
||||
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)
|
||||
})
|
||||
}
|
||||
@@ -1,152 +0,0 @@
|
||||
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"))
|
||||
}
|
||||
@@ -1,112 +0,0 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -1,43 +0,0 @@
|
||||
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