Files
terdut-server/internal/api/helpers.go
T
Niklas Ye 01922291f6
CI / chart (pull_request) Successful in 3s
CI / security (pull_request) Successful in 25s
CI / test (pull_request) Successful in 5m42s
Satisfy gosec on the proxy count and the request log
G115: the trusted-proxy count is stored as an int64 instead of narrowing
it to int32. G706: the request logger and serverError quote the request
method and route; the logger's remaining taint comes from the wrapped
response writer, so it carries a justified nosec like the other quoted log
lines.

Claude-Session: https://claude.ai/code/session_016mBLURvJoMuUEr9cB2RpUN
2026-10-09 15:27:45 +02:00

145 lines
5.5 KiB
Go

package api
import (
"context"
"database/sql"
"encoding/json"
"errors"
"github.com/go-chi/chi/v5"
"log"
"net/http"
"strconv"
"strings"
"github.com/jackc/pgerrcode"
"github.com/jackc/pgx/v5/pgconn"
)
// sqlArgs accumulates query arguments and hands back the placeholder for each.
//
// Postgres numbers its placeholders, so a dynamically assembled WHERE clause has
// to keep its $1, $2, … in step with the order of the values. Handing out the placeholder and storing the value
// in one call is what keeps them in step: a filter can be added, removed or
// reordered without renumbering anything by hand.
type sqlArgs struct{ vals []any }
// add stores v and returns the placeholder that refers to it.
func (a *sqlArgs) add(v any) string {
a.vals = append(a.vals, v)
return "$" + strconv.Itoa(len(a.vals))
}
// addList stores every value and returns their placeholders as "$1, $2, …",
// ready to drop into an IN (…) clause. Returns an empty string for no values,
// which no caller should reach: `IN ()` is a syntax error in Postgres, so callers check for an empty set before building the query.
func (a *sqlArgs) addList(vs []any) string {
parts := make([]string, len(vs))
for i, v := range vs {
parts[i] = a.add(v)
}
return strings.Join(parts, ", ")
}
// all returns the accumulated values, to be passed straight to Query or Exec.
func (a *sqlArgs) all() []any { return a.vals }
// nowEpoch is the SQL expression for "now, as unix seconds", matching how every
// timestamp in this schema is stored.
//
// FLOOR, not a bare cast: EXTRACT returns fractional seconds and casting to
// bigint rounds half up, so a row written at .6 of a second would claim a
// timestamp one second in the future — off by one against the time.Now().Unix()
// the Go side stamps, which is what the expiry tests measure.
const nowEpoch = "FLOOR(EXTRACT(EPOCH FROM now()))::bigint"
// isUniqueViolation reports whether err is a broken unique constraint, which
// callers turn into 409 Conflict rather than 500.
//
// Postgres reports it as SQLSTATE 23505 on a typed error. Matching the code
// means a renamed constraint or a translated message cannot quietly turn a
// conflict back into a 500.
func isUniqueViolation(err error) bool {
var pgErr *pgconn.PgError
return errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation
}
func respond(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
json.NewEncoder(w).Encode(v)
}
// serverError answers 500 and logs why. The response stays opaque, so the log
// line is the only record of what failed.
func serverError(w http.ResponseWriter, r *http.Request, err error) {
// The route pattern, not the path: two routes carry a credential in it.
route := r.URL.Path
if rc := chi.RouteContext(r.Context()); rc != nil && rc.RoutePattern() != "" {
route = rc.RoutePattern()
}
log.Printf("%s %s: %v", strconv.Quote(r.Method), strconv.Quote(route), err)
respond(w, http.StatusInternalServerError, errResp("internal error"))
}
// maxBodyBytes caps an ordinary JSON request body. 1 MiB is far more than any
// endpoint below needs — it exists so an unauthenticated caller (signup,
// login, bootstrap) can't make the server buffer an arbitrarily large body
// before the request is even validated.
const maxBodyBytes = 1 << 20
func decodeJSON(r *http.Request, v any) error {
return decodeJSONLimit(r, v, maxBodyBytes)
}
// decodeJSONLimit is decodeJSON with an explicit cap, for the one endpoint
// (the Alertmanager webhook, see maxWebhookBodyBytes) whose real payloads can
// legitimately be larger than maxBodyBytes.
func decodeJSONLimit(r *http.Request, v any, limit int64) error {
defer r.Body.Close()
// w is nil: there is no ResponseWriter here to disable keep-alive with,
// which net/http documents as fine — the limit is still enforced, the
// connection just isn't closed early on a request that blows past it.
r.Body = http.MaxBytesReader(nil, r.Body, limit)
return json.NewDecoder(r.Body).Decode(v)
}
func errResp(msg string) map[string]string {
return map[string]string{"error": msg}
}
// withAdvisoryLock runs fn only if it can take the named Postgres advisory lock on a
// dedicated connection, and skips fn otherwise. This is what keeps the archiver and
// notifier safe to run on more than one replica: whichever instance's tick gets there
// first does the work; the rest see the lock held and simply wait for their next tick
// instead of running the same pass concurrently.
//
// pg_try_advisory_lock is session-scoped, so taking and releasing it must happen on the
// same connection, reserved via db.Conn rather than borrowed from the pool's shared
// connections fn itself may use — and released (unlocked, then closed) before returning,
// since a session lock otherwise outlives this call and leaks onto whatever reuses the
// pooled connection next.
func withAdvisoryLock(ctx context.Context, db *sql.DB, key int64, name string, fn func()) {
conn, err := db.Conn(ctx)
if err != nil {
log.Printf("%s: advisory lock: acquire connection: %v", name, err)
return
}
defer conn.Close()
var locked bool
if err := conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", key).Scan(&locked); err != nil {
log.Printf("%s: advisory lock: %v", name, err)
return
}
if !locked {
return // another replica is already running this pass
}
defer func() {
if _, err := conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", key); err != nil {
log.Printf("%s: advisory unlock: %v", name, err)
}
}()
fn()
}