first commit
This commit is contained in:
135
go-api/cmd/api/main.go
Normal file
135
go-api/cmd/api/main.go
Normal file
@@ -0,0 +1,135 @@
|
||||
// Command api is the Krow HTTP API.
|
||||
//
|
||||
// Configuration, a verified PostgreSQL pool, /health, the sign-in endpoints,
|
||||
// and the entity and current-user endpoints described in docs/api-contract.md.
|
||||
//
|
||||
// Every request outside the public allowlist carries a session cookie that this
|
||||
// process resolves to a real user. The development organization that used to be
|
||||
// injected into every request is gone: identity now comes from the sessions
|
||||
// table, and a database with no users is a database nobody can sign in to,
|
||||
// which is the correct behaviour rather than a gap.
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
)
|
||||
|
||||
func main() {
|
||||
if err := run(); err != nil {
|
||||
slog.Error("fatal", "error", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func run() error {
|
||||
cfg, err := config.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
log := newLogger(cfg.Log.Level)
|
||||
log.Info("starting krow-api", "env", cfg.AppEnv, "database", cfg.DB.Redacted())
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||
defer stop()
|
||||
|
||||
database, err := db.Open(ctx, cfg.DB)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer database.Close()
|
||||
log.Info("database connected", "schema", cfg.DB.Schema)
|
||||
|
||||
server, err := httpserver.New(cfg, database, log)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// The sweeper's context is cancelled by the same signal that stops the
|
||||
// server, so the ticker goes away with the process rather than outliving
|
||||
// the pool it queries.
|
||||
go sweepSessions(ctx, server.Sessions(), log)
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- server.Start() }()
|
||||
log.Info("listening", "addr", server.Addr(), "endpoints", server.Endpoints(),
|
||||
"health", "http://"+server.Addr()+"/health",
|
||||
"cors_origins", cfg.HTTP.CORSOrigins)
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
log.Info("shutdown signal received, draining")
|
||||
// context.Background: ctx is already cancelled, and Shutdown needs a
|
||||
// live deadline of its own to drain in-flight requests.
|
||||
return server.Shutdown(context.Background())
|
||||
}
|
||||
}
|
||||
|
||||
// sweepInterval is how often dead sessions are collected.
|
||||
//
|
||||
// Sweeping is housekeeping, not correctness: Manager.Authenticate already
|
||||
// refuses an expired session and deletes the row as it finds it, so a session
|
||||
// is never usable between its expiry and the next sweep. This only collects the
|
||||
// rows nobody comes back for. Fifteen minutes keeps the table from growing
|
||||
// without putting a DELETE on any hot path.
|
||||
const sweepInterval = 15 * time.Minute
|
||||
|
||||
// sweepSessions deletes expired sessions until the context is cancelled.
|
||||
//
|
||||
// It runs once immediately so a process that has been down for a while does not
|
||||
// carry a backlog for a further fifteen minutes, then on the ticker. A failed
|
||||
// sweep is logged and retried at the next tick: the table being briefly larger
|
||||
// than it should be is not worth stopping the API for.
|
||||
func sweepSessions(ctx context.Context, sessions *auth.Manager, log *slog.Logger) {
|
||||
ticker := time.NewTicker(sweepInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
sweep := func() {
|
||||
// A deadline of its own, so a slow or wedged DELETE cannot leave this
|
||||
// goroutine blocked past shutdown.
|
||||
sweepCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer cancel()
|
||||
n, err := sessions.Sweep(sweepCtx)
|
||||
switch {
|
||||
case err != nil && ctx.Err() != nil:
|
||||
// Shutting down; the cancellation is expected, not a failure.
|
||||
case err != nil:
|
||||
log.Warn("session sweep failed", "error", err)
|
||||
case n > 0:
|
||||
log.Info("swept expired sessions", "deleted", n)
|
||||
default:
|
||||
log.Debug("session sweep found nothing to delete")
|
||||
}
|
||||
}
|
||||
|
||||
sweep()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
log.Debug("session sweeper stopped")
|
||||
return
|
||||
case <-ticker.C:
|
||||
sweep()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func newLogger(level string) *slog.Logger {
|
||||
var lvl slog.Level
|
||||
if err := lvl.UnmarshalText([]byte(level)); err != nil {
|
||||
lvl = slog.LevelInfo
|
||||
}
|
||||
return slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: lvl}))
|
||||
}
|
||||
75
go-api/cmd/seed/main.go
Normal file
75
go-api/cmd/seed/main.go
Normal file
@@ -0,0 +1,75 @@
|
||||
// Command seed loads the frontend's demo dataset into PostgreSQL.
|
||||
//
|
||||
// Safe to run repeatedly: every record's key is derived deterministically from
|
||||
// its source id and written with ON CONFLICT DO UPDATE inside one transaction,
|
||||
// so re-running restores the seeded values without duplicating a row or
|
||||
// deleting anything. See internal/seeder.
|
||||
//
|
||||
// make seed
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/seeder"
|
||||
)
|
||||
|
||||
func main() {
|
||||
if err := run(); err != nil {
|
||||
fmt.Fprintln(os.Stderr, "seed failed:", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func run() error {
|
||||
cfg, err := config.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fixture, err := seeder.Load(cfg.Seed.FixturePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
database, err := db.Open(ctx, cfg.DB)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer database.Close()
|
||||
|
||||
started := time.Now()
|
||||
result, err := seeder.New(database.Pool, fixture, time.Now()).Run(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
names := make([]string, 0, len(result.Counts))
|
||||
total := 0
|
||||
for name, n := range result.Counts {
|
||||
names = append(names, name)
|
||||
total += n
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
fmt.Printf("seeded into %s (organization %s)\n", cfg.DB.Name, result.OrgID)
|
||||
for _, name := range names {
|
||||
fmt.Printf(" %-16s %4d\n", name, result.Counts[name])
|
||||
}
|
||||
fmt.Printf(" %-16s %4d records in %s\n", "TOTAL", total, time.Since(started).Round(time.Millisecond))
|
||||
if result.Pruned > 0 {
|
||||
// Shift records are a rolling window; ones that fell out of it are
|
||||
// removed. Said out loud, because a seed that deletes should say so.
|
||||
fmt.Printf(" %-16s %4d stale shift record(s) outside the current window\n", "PRUNED", result.Pruned)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
259
go-api/cmd/setpassword/main.go
Normal file
259
go-api/cmd/setpassword/main.go
Normal file
@@ -0,0 +1,259 @@
|
||||
// Command setpassword sets a user's password.
|
||||
//
|
||||
// It exists because migration 000001 left users.password_hash nullable and
|
||||
// NULL, and the seeded demo user still has no password. Nothing in the seed
|
||||
// fixture, the migrations or this repository contains, generates or defaults a
|
||||
// password: a password enters the system here, typed by a person, and nowhere
|
||||
// else.
|
||||
//
|
||||
// # prompt for the password, twice, with the input hidden
|
||||
// cd go-api && go run ./cmd/setpassword -email demo@krow.app
|
||||
// cd go-api && go run ./cmd/setpassword -id 9a1f...-uuid
|
||||
//
|
||||
// # non-interactive, for a provisioning script — the password arrives on
|
||||
// # stdin, never in argv, so it does not reach `ps` or the shell history
|
||||
// printf '%s' "$NEW_PASSWORD" | go run ./cmd/setpassword -email demo@krow.app -stdin
|
||||
//
|
||||
// There is deliberately no -password flag. A password in argv is visible to
|
||||
// every process on the machine through `ps`, and lands in the shell history
|
||||
// besides. stdin is the only non-interactive route.
|
||||
//
|
||||
// The password, the confirmation and the resulting hash are never printed,
|
||||
// never logged and never written anywhere but the users.password_hash column,
|
||||
// through a bind parameter.
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"golang.org/x/term"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
)
|
||||
|
||||
func main() {
|
||||
if err := run(); err != nil {
|
||||
// The error strings in this file name rules and identifiers only. No
|
||||
// path here can carry the password into this line.
|
||||
fmt.Fprintf(os.Stderr, "setpassword: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
type target struct {
|
||||
id string
|
||||
email string
|
||||
role string
|
||||
hadHash bool
|
||||
}
|
||||
|
||||
func run() error {
|
||||
var (
|
||||
email = flag.String("email", "", "the user's email address")
|
||||
id = flag.String("id", "", "the user's UUID")
|
||||
fromStdin = flag.Bool("stdin", false, "read the password from stdin instead of prompting")
|
||||
)
|
||||
flag.Usage = func() {
|
||||
fmt.Fprintf(flag.CommandLine.Output(),
|
||||
"Usage: setpassword (-email <address> | -id <uuid>) [-stdin]\n\n"+
|
||||
"Sets one user's password, hashed with argon2id. The password is never\n"+
|
||||
"echoed, printed or logged, and there is no -password flag by design.\n\n")
|
||||
flag.PrintDefaults()
|
||||
}
|
||||
flag.Parse()
|
||||
|
||||
if flag.NArg() > 0 {
|
||||
// A bare argument is most likely someone typing the password after the
|
||||
// command. Refuse loudly rather than ignoring it — and say nothing
|
||||
// about what the argument was.
|
||||
return errors.New("unexpected positional argument; pass -email or -id, and supply the password when prompted")
|
||||
}
|
||||
if (*email == "") == (*id == "") {
|
||||
return errors.New("pass exactly one of -email or -id")
|
||||
}
|
||||
|
||||
cfg, err := config.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
database, err := db.Open(ctx, cfg.DB)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer database.Close()
|
||||
|
||||
// Resolve and show the target BEFORE asking for a password, so nobody
|
||||
// types a secret at a prompt that turns out to be pointed at the wrong
|
||||
// user, or at no user at all.
|
||||
t, err := resolve(ctx, database, *email, *id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, "database: %s\nuser: %s <%s>\nrole: %s\npassword: %s\n\n",
|
||||
cfg.DB.Name, t.id, t.email, t.role, existingState(t.hadHash))
|
||||
|
||||
password, err := readPassword(*fromStdin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// The plaintext lives in this slice and nowhere else. Wipe it as soon as
|
||||
// the hash exists. Go's garbage collector may still have copied it, so
|
||||
// this is a reduction in exposure rather than a guarantee — worth doing,
|
||||
// not worth trusting.
|
||||
defer wipe(password)
|
||||
|
||||
if err := auth.ValidatePassword(string(password)); err != nil {
|
||||
return describePolicy(err)
|
||||
}
|
||||
|
||||
hash, err := auth.HashPassword(string(password))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Parameterized, and keyed by the UUID resolved above rather than by the
|
||||
// string the operator typed. Neither the hash nor the password is ever
|
||||
// interpolated into SQL.
|
||||
const q = `UPDATE users SET password_hash = $2::text, updated_date = now() WHERE id = $1::uuid`
|
||||
tag, err := database.Pool.Exec(ctx, q, t.id, hash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update password: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() != 1 {
|
||||
return fmt.Errorf("expected to update exactly one user, updated %d", tag.RowsAffected())
|
||||
}
|
||||
|
||||
// Confirms the identity and nothing about the secret: no hash, no length,
|
||||
// no prefix.
|
||||
fmt.Fprintf(os.Stderr, "password set for %s (%s)\n", t.email, t.id)
|
||||
return nil
|
||||
}
|
||||
|
||||
func existingState(had bool) string {
|
||||
if had {
|
||||
return "already set (it will be replaced)"
|
||||
}
|
||||
return "not set yet"
|
||||
}
|
||||
|
||||
// resolve finds exactly one user by email or by id.
|
||||
//
|
||||
// Email lookup relies on the citext column, so it is case-insensitive, and on
|
||||
// the global unique index added by migration 000004, so it cannot match two
|
||||
// users in two organizations.
|
||||
func resolve(ctx context.Context, database *db.DB, email, id string) (target, error) {
|
||||
var (
|
||||
t target
|
||||
err error
|
||||
)
|
||||
if email != "" {
|
||||
const q = `SELECT id::text, email::text, role, password_hash IS NOT NULL
|
||||
FROM users WHERE email = $1::citext`
|
||||
err = database.Pool.QueryRow(ctx, q, strings.TrimSpace(email)).
|
||||
Scan(&t.id, &t.email, &t.role, &t.hadHash)
|
||||
} else {
|
||||
const q = `SELECT id::text, email::text, role, password_hash IS NOT NULL
|
||||
FROM users WHERE id = $1::uuid`
|
||||
err = database.Pool.QueryRow(ctx, q, strings.TrimSpace(id)).
|
||||
Scan(&t.id, &t.email, &t.role, &t.hadHash)
|
||||
}
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return t, errors.New("no such user")
|
||||
}
|
||||
if err != nil {
|
||||
return t, fmt.Errorf("look up user: %w", err)
|
||||
}
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// readPassword collects the password without echoing it.
|
||||
//
|
||||
// Interactively it asks twice and compares, because a mistyped password that
|
||||
// nobody can see is otherwise only discovered at the next login. With -stdin
|
||||
// it reads the stream verbatim, minus one trailing newline, so
|
||||
// `printf '%s' "$P" | setpassword -stdin` and a here-string both work.
|
||||
func readPassword(fromStdin bool) ([]byte, error) {
|
||||
if fromStdin {
|
||||
raw, err := io.ReadAll(os.Stdin)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read password from stdin: %w", err)
|
||||
}
|
||||
return trimOneNewline(raw), nil
|
||||
}
|
||||
|
||||
fd := int(os.Stdin.Fd())
|
||||
if !term.IsTerminal(fd) {
|
||||
// Falling back to an echoing read here would print the password to the
|
||||
// screen and into any transcript. Refuse and name the flag instead.
|
||||
return nil, errors.New("stdin is not a terminal; re-run with -stdin to read the password from the pipe")
|
||||
}
|
||||
|
||||
fmt.Fprint(os.Stderr, "New password: ")
|
||||
first, err := term.ReadPassword(fd)
|
||||
fmt.Fprintln(os.Stderr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read password: %w", err)
|
||||
}
|
||||
|
||||
fmt.Fprint(os.Stderr, "Confirm password: ")
|
||||
second, err := term.ReadPassword(fd)
|
||||
fmt.Fprintln(os.Stderr)
|
||||
if err != nil {
|
||||
wipe(first)
|
||||
return nil, fmt.Errorf("read confirmation: %w", err)
|
||||
}
|
||||
defer wipe(second)
|
||||
|
||||
if string(first) != string(second) {
|
||||
wipe(first)
|
||||
return nil, errors.New("the two entries do not match")
|
||||
}
|
||||
return first, nil
|
||||
}
|
||||
|
||||
// describePolicy turns a policy error into advice, still without quoting the
|
||||
// password or revealing its length.
|
||||
func describePolicy(err error) error {
|
||||
switch {
|
||||
case errors.Is(err, auth.ErrEmptyPassword):
|
||||
return errors.New("the password is empty")
|
||||
case errors.Is(err, auth.ErrPasswordTooShort):
|
||||
return fmt.Errorf("the password is too short; it must be at least %d bytes", auth.MinPasswordLength)
|
||||
case errors.Is(err, auth.ErrPasswordTooLong):
|
||||
return fmt.Errorf("the password is too long; the maximum is %d bytes", auth.MaxPasswordLength)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// trimOneNewline removes a single trailing "\n" or "\r\n", and only one: a
|
||||
// password may legitimately end in whitespace, so this strips the line
|
||||
// terminator a shell adds and nothing more.
|
||||
func trimOneNewline(b []byte) []byte {
|
||||
if n := len(b); n > 0 && b[n-1] == '\n' {
|
||||
b = b[:n-1]
|
||||
if n := len(b); n > 0 && b[n-1] == '\r' {
|
||||
b = b[:n-1]
|
||||
}
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func wipe(b []byte) {
|
||||
for i := range b {
|
||||
b[i] = 0
|
||||
}
|
||||
}
|
||||
18
go-api/go.mod
Normal file
18
go-api/go.mod
Normal file
@@ -0,0 +1,18 @@
|
||||
module github.com/krow/krow-backend/go-api
|
||||
|
||||
go 1.27
|
||||
|
||||
require (
|
||||
github.com/jackc/pgx/v5 v5.10.0
|
||||
golang.org/x/crypto v0.42.0
|
||||
golang.org/x/term v0.35.0
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/jackc/pgpassfile v1.0.0 // indirect
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.2 // indirect
|
||||
golang.org/x/sync v0.17.0 // indirect
|
||||
golang.org/x/sys v0.37.0 // indirect
|
||||
golang.org/x/text v0.29.0 // indirect
|
||||
)
|
||||
32
go-api/go.sum
Normal file
32
go-api/go.sum
Normal file
@@ -0,0 +1,32 @@
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
|
||||
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
|
||||
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
|
||||
github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0=
|
||||
github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
|
||||
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
|
||||
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
|
||||
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
|
||||
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
golang.org/x/crypto v0.42.0 h1:chiH31gIWm57EkTXpwnqf8qeuMUi0yekh6mT2AvFlqI=
|
||||
golang.org/x/crypto v0.42.0/go.mod h1:4+rDnOTJhQCx2q7/j6rAN5XDw8kPjeaXEUR2eL94ix8=
|
||||
golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug=
|
||||
golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ=
|
||||
golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||
golang.org/x/term v0.35.0 h1:bZBVKBudEyhRcajGcNc3jIfWPqV4y/Kt2XcoigOWtDQ=
|
||||
golang.org/x/term v0.35.0/go.mod h1:TPGtkTLesOwf2DE8CgVYiZinHAOuy5AYUYT1lENIZnA=
|
||||
golang.org/x/text v0.29.0 h1:1neNs90w9YzJ9BocxfsQNHKuAT4pkghyXc4nhZ6sJvk=
|
||||
golang.org/x/text v0.29.0/go.mod h1:7MhJOA9CD2qZyOKYazxdYMF85OwPdEr9jTtBpO7ydH4=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
125
go-api/internal/auth/credentials.go
Normal file
125
go-api/internal/auth/credentials.go
Normal file
@@ -0,0 +1,125 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ErrInvalidCredentials is the single answer to every failed sign-in.
|
||||
//
|
||||
// It is returned when the email is unknown, when the password is wrong, when
|
||||
// the account has no password set, and when the account is suspended. The
|
||||
// caller cannot tell those apart from the error, which is the point: an error
|
||||
// that distinguishes them is an account-enumeration oracle, and one that says
|
||||
// "this account is suspended" confirms the address is real.
|
||||
//
|
||||
// The distinction is still made — it goes to the server log, via the reason
|
||||
// returned alongside this error.
|
||||
var ErrInvalidCredentials = errors.New("auth: invalid credentials")
|
||||
|
||||
// Reason is why a sign-in failed. For the log, never for the client.
|
||||
type Reason string
|
||||
|
||||
const (
|
||||
ReasonOK Reason = "ok"
|
||||
ReasonNoSuchUser Reason = "no_such_user"
|
||||
ReasonNoPassword Reason = "no_password_set"
|
||||
ReasonNotActive Reason = "user_not_active"
|
||||
ReasonBadPassword Reason = "wrong_password"
|
||||
)
|
||||
|
||||
// Credentials verifies a password against a stored user.
|
||||
type Credentials struct {
|
||||
users UserStore
|
||||
}
|
||||
|
||||
// NewCredentials builds the verifier.
|
||||
func NewCredentials(users UserStore) *Credentials { return &Credentials{users: users} }
|
||||
|
||||
// decoyHash is a real argon2id hash of a value nobody knows.
|
||||
//
|
||||
// It exists to close a timing side channel. Without it, an unknown email
|
||||
// returns as fast as the database can say "no rows" — a millisecond or two —
|
||||
// while a known email spends the ~100ms that argon2id costs by design. That
|
||||
// difference is trivially measurable over a network and turns the login
|
||||
// endpoint into an account-enumeration oracle no matter how carefully the
|
||||
// error messages are worded.
|
||||
//
|
||||
// So every failure that skips the real password check pays for a decoy one
|
||||
// instead. The comparison always fails; the cost is the entire purpose.
|
||||
//
|
||||
// Built once, lazily: it costs a full argon2id derivation, which is worth
|
||||
// paying on the first failed login rather than on every process start.
|
||||
var decoyHash = sync.OnceValue(func() string {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
// A hash of a fixed string is still a fine decoy — its only job is to
|
||||
// take the right amount of time, and it is never compared against
|
||||
// anything a caller supplies.
|
||||
token = "decoy-password-that-is-never-correct"
|
||||
}
|
||||
hash, err := HashPassword(token)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return hash
|
||||
})
|
||||
|
||||
// burnTime performs a password verification that is guaranteed to fail, so a
|
||||
// rejected sign-in costs the same as an accepted one.
|
||||
func burnTime(password string) {
|
||||
if h := decoyHash(); h != "" {
|
||||
_, _ = VerifyPassword(h, password)
|
||||
}
|
||||
}
|
||||
|
||||
// Verify resolves an email and password to a user.
|
||||
//
|
||||
// On success it returns the user and ReasonOK. On any failure it returns
|
||||
// ErrInvalidCredentials, a zero user, and the reason — which the caller should
|
||||
// log and must not send to the client.
|
||||
//
|
||||
// A non-nil error that is NOT ErrInvalidCredentials is an operational failure
|
||||
// (the database is down, a stored hash is corrupt) and should become a 500
|
||||
// rather than a 401: the caller's credentials were never actually judged.
|
||||
func (c *Credentials) Verify(ctx context.Context, email, password string) (User, Reason, error) {
|
||||
user, err := c.users.FindByEmail(ctx, email)
|
||||
if errors.Is(err, ErrUserNotFound) {
|
||||
burnTime(password)
|
||||
return User{}, ReasonNoSuchUser, ErrInvalidCredentials
|
||||
}
|
||||
if err != nil {
|
||||
return User{}, ReasonNoSuchUser, err
|
||||
}
|
||||
|
||||
if user.PasswordHash == "" {
|
||||
// The seeded user is in this state until `setpassword` is run against
|
||||
// it. Refused exactly like a wrong password, at the same cost.
|
||||
burnTime(password)
|
||||
return User{}, ReasonNoPassword, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
ok, err := VerifyPassword(user.PasswordHash, password)
|
||||
if err != nil {
|
||||
// The stored hash could not be read. That is this server's problem,
|
||||
// not the caller's, and must not be reported as a failed login.
|
||||
return User{}, ReasonBadPassword, err
|
||||
}
|
||||
if !ok {
|
||||
return User{}, ReasonBadPassword, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
// Status is checked AFTER the password, and reported the same way.
|
||||
//
|
||||
// Order matters: checking it first would let anyone learn that an address
|
||||
// belongs to a suspended account without knowing its password, because the
|
||||
// refusal would arrive without paying the argon2 cost. Checking it after
|
||||
// means a suspended account is indistinguishable from a wrong password —
|
||||
// same answer, same timing.
|
||||
if !user.IsActive() {
|
||||
return User{}, ReasonNotActive, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
return user, ReasonOK, nil
|
||||
}
|
||||
230
go-api/internal/auth/password.go
Normal file
230
go-api/internal/auth/password.go
Normal file
@@ -0,0 +1,230 @@
|
||||
// Package auth is the authentication foundation: password hashing, session
|
||||
// tokens, and the server-side session lifecycle.
|
||||
//
|
||||
// It deliberately knows nothing about HTTP. There is no handler, no cookie and
|
||||
// no middleware here — those arrive in a later phase and will be written in
|
||||
// terms of this package, not inside it. What lives here is the part that must
|
||||
// be correct regardless of transport: how a password becomes a hash, how a
|
||||
// session token is generated and stored, and when a session stops being valid.
|
||||
//
|
||||
// Two rules hold throughout, and every function below is written to keep them:
|
||||
//
|
||||
// - A raw session token exists in exactly two places: the response that
|
||||
// created it, and the client's cookie. The database holds SHA-256 of it.
|
||||
// - Neither a password nor a token nor a password hash is ever returned in an
|
||||
// error, formatted into a string, or logged. Nothing in this package logs.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
// Password policy errors. They describe the rule that was broken and never
|
||||
// echo the password back.
|
||||
var (
|
||||
ErrEmptyPassword = errors.New("auth: password is empty")
|
||||
ErrPasswordTooShort = errors.New("auth: password is shorter than the minimum length")
|
||||
ErrPasswordTooLong = errors.New("auth: password is longer than the maximum length")
|
||||
|
||||
// ErrInvalidHash means the stored value is not a hash this package wrote:
|
||||
// wrong prefix, wrong field count, or unparseable parameters.
|
||||
ErrInvalidHash = errors.New("auth: password hash is malformed")
|
||||
|
||||
// ErrIncompatibleVersion means the hash was produced by a future argon2
|
||||
// version this binary cannot verify. Distinguished from ErrInvalidHash
|
||||
// because it is an upgrade problem, not corruption.
|
||||
ErrIncompatibleVersion = errors.New("auth: password hash uses an unsupported argon2 version")
|
||||
)
|
||||
|
||||
const (
|
||||
// MinPasswordLength is measured in bytes, not runes. A byte floor is the
|
||||
// honest one: it is what the KDF consumes, and counting runes would let a
|
||||
// short ASCII password through by way of a generous rune count.
|
||||
MinPasswordLength = 12
|
||||
|
||||
// MaxPasswordLength caps the input. Argon2 has no internal length limit —
|
||||
// unlike bcrypt, it does not silently truncate — so the only reason for a
|
||||
// ceiling is to stop an unbounded body from being hashed at 64 MiB of
|
||||
// memory per attempt. 1 KiB is far above any real passphrase.
|
||||
MaxPasswordLength = 1024
|
||||
)
|
||||
|
||||
// PasswordParams are the argon2id cost parameters.
|
||||
//
|
||||
// They are stored inside every hash this package writes, so a future increase
|
||||
// does not invalidate existing hashes: verification reads the parameters out of
|
||||
// the stored string rather than assuming today's defaults.
|
||||
type PasswordParams struct {
|
||||
// Memory is the KiB of memory the KDF fills. This is the parameter that
|
||||
// makes GPU and ASIC attacks expensive, and the one worth raising first.
|
||||
Memory uint32
|
||||
// Time is the number of passes over that memory.
|
||||
Time uint32
|
||||
// Threads is the parallelism (argon2's `p`).
|
||||
Threads uint8
|
||||
// SaltLength and KeyLength are in bytes.
|
||||
SaltLength uint32
|
||||
KeyLength uint32
|
||||
}
|
||||
|
||||
// DefaultPasswordParams follows the OWASP Password Storage Cheat Sheet's
|
||||
// argon2id recommendation: 64 MiB of memory, 3 iterations, 4 lanes (m=65536,
|
||||
// t=3, p=4). A 16-byte salt and a 32-byte key are the RFC 9106 defaults.
|
||||
//
|
||||
// This costs roughly a tenth of a second per login on developer hardware,
|
||||
// which is the point: it is a cost an attacker pays per guess.
|
||||
var DefaultPasswordParams = PasswordParams{
|
||||
Memory: 64 * 1024,
|
||||
Time: 3,
|
||||
Threads: 4,
|
||||
SaltLength: 16,
|
||||
KeyLength: 32,
|
||||
}
|
||||
|
||||
// HashPassword hashes a plaintext password with the default parameters.
|
||||
//
|
||||
// The returned string is a complete, self-describing PHC record — algorithm,
|
||||
// version, parameters, salt and digest — and is what belongs in
|
||||
// users.password_hash. It is safe to store and unsafe to log.
|
||||
func HashPassword(plain string) (string, error) {
|
||||
return HashPasswordWithParams(plain, DefaultPasswordParams)
|
||||
}
|
||||
|
||||
// HashPasswordWithParams is HashPassword with explicit cost parameters. Tests
|
||||
// use it to run at a cost that does not dominate the test suite; production
|
||||
// code should call HashPassword.
|
||||
func HashPasswordWithParams(plain string, p PasswordParams) (string, error) {
|
||||
if err := ValidatePassword(plain); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if p.SaltLength == 0 || p.KeyLength == 0 || p.Memory == 0 || p.Time == 0 || p.Threads == 0 {
|
||||
return "", fmt.Errorf("auth: argon2id parameters must all be non-zero")
|
||||
}
|
||||
|
||||
salt := make([]byte, p.SaltLength)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
// crypto/rand failing is not recoverable and must never fall back to a
|
||||
// weaker source: a predictable salt defeats the whole construction.
|
||||
return "", fmt.Errorf("auth: read salt: %w", err)
|
||||
}
|
||||
|
||||
key := argon2.IDKey([]byte(plain), salt, p.Time, p.Memory, p.Threads, p.KeyLength)
|
||||
|
||||
// The PHC string format, as produced by the reference implementation:
|
||||
// $argon2id$v=19$m=65536,t=3,p=4$<b64 salt>$<b64 key>
|
||||
// Standard base64 without padding, which is what the format specifies.
|
||||
return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
|
||||
argon2.Version, p.Memory, p.Time, p.Threads,
|
||||
base64.RawStdEncoding.EncodeToString(salt),
|
||||
base64.RawStdEncoding.EncodeToString(key),
|
||||
), nil
|
||||
}
|
||||
|
||||
// VerifyPassword reports whether plain is the password behind encoded.
|
||||
//
|
||||
// A false return with a nil error is the ordinary "wrong password" answer. A
|
||||
// non-nil error means the *stored hash* could not be read, which is an
|
||||
// operational problem rather than a failed login, and callers should tell the
|
||||
// two apart: the first is a 401, the second is a 500.
|
||||
//
|
||||
// The digest comparison is constant-time. The length and parameter checks
|
||||
// before it are not, and do not need to be: they depend only on the stored
|
||||
// hash, never on the supplied password.
|
||||
func VerifyPassword(encoded, plain string) (bool, error) {
|
||||
p, salt, want, err := DecodePasswordHash(encoded)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
// No policy check on `plain` here. A password that predates a tightened
|
||||
// minimum length must still be able to log in; the policy applies when a
|
||||
// password is set, which is where ValidatePassword is called.
|
||||
if len(plain) > MaxPasswordLength {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
got := argon2.IDKey([]byte(plain), salt, p.Time, p.Memory, p.Threads, p.KeyLength)
|
||||
return subtle.ConstantTimeCompare(got, want) == 1, nil
|
||||
}
|
||||
|
||||
// ValidatePassword applies the policy for setting a new password.
|
||||
func ValidatePassword(plain string) error {
|
||||
switch {
|
||||
case len(plain) == 0:
|
||||
return ErrEmptyPassword
|
||||
case len(plain) < MinPasswordLength:
|
||||
return ErrPasswordTooShort
|
||||
case len(plain) > MaxPasswordLength:
|
||||
return ErrPasswordTooLong
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DecodePasswordHash parses a PHC argon2id record back into its parts.
|
||||
//
|
||||
// Exported so that a future re-hash-on-login path can ask whether a stored hash
|
||||
// was written with weaker parameters than today's default and upgrade it. It
|
||||
// returns the salt and digest, never the password.
|
||||
func DecodePasswordHash(encoded string) (p PasswordParams, salt, key []byte, err error) {
|
||||
// $argon2id$v=19$m=65536,t=3,p=4$<salt>$<key> splits into six fields, the
|
||||
// first of which is empty because the string starts with the separator.
|
||||
parts := strings.Split(encoded, "$")
|
||||
if len(parts) != 6 || parts[0] != "" {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if parts[1] != "argon2id" {
|
||||
// bcrypt, argon2i and argon2d all land here. This package writes and
|
||||
// reads argon2id and nothing else; a different algorithm is a
|
||||
// migration decision, not something to guess at during a login.
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
var version int
|
||||
if _, err := fmt.Sscanf(parts[2], "v=%d", &version); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if version != argon2.Version {
|
||||
return p, nil, nil, ErrIncompatibleVersion
|
||||
}
|
||||
|
||||
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &p.Memory, &p.Time, &p.Threads); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if p.Memory == 0 || p.Time == 0 || p.Threads == 0 {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
if salt, err = base64.RawStdEncoding.DecodeString(parts[4]); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if key, err = base64.RawStdEncoding.DecodeString(parts[5]); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if len(salt) == 0 || len(key) == 0 {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
p.SaltLength = uint32(len(salt))
|
||||
p.KeyLength = uint32(len(key))
|
||||
return p, salt, key, nil
|
||||
}
|
||||
|
||||
// NeedsRehash reports whether a stored hash was written with parameters weaker
|
||||
// than want, so a successful login can transparently upgrade it.
|
||||
//
|
||||
// Unused in Phase 3B — there is no login yet — and exported now because the
|
||||
// judgement belongs beside the format that encodes the parameters.
|
||||
func NeedsRehash(encoded string, want PasswordParams) bool {
|
||||
p, _, _, err := DecodePasswordHash(encoded)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return p.Memory < want.Memory || p.Time < want.Time ||
|
||||
p.KeyLength < want.KeyLength || p.SaltLength < want.SaltLength
|
||||
}
|
||||
254
go-api/internal/auth/password_test.go
Normal file
254
go-api/internal/auth/password_test.go
Normal file
@@ -0,0 +1,254 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// testParams runs argon2id at a cost that is still real but does not make the
|
||||
// suite crawl. Every property under test — salting, verification, the encoded
|
||||
// format — is independent of the cost, and DefaultPasswordParams is asserted
|
||||
// separately in TestDefaultPasswordParamsMeetOWASP.
|
||||
var testParams = PasswordParams{Memory: 8 * 1024, Time: 1, Threads: 2, SaltLength: 16, KeyLength: 32}
|
||||
|
||||
const goodPassword = "correct-horse-battery-staple"
|
||||
|
||||
// 1. Hash generation produces a well-formed, self-describing argon2id record.
|
||||
func TestHashPasswordProducesArgon2idRecord(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("HashPasswordWithParams: %v", err)
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(hash, "$argon2id$") {
|
||||
t.Fatalf("hash is not argon2id: %q", firstField(hash))
|
||||
}
|
||||
// bcrypt would be "$2a$"/"$2b$"; argon2i and argon2d are different KDFs.
|
||||
// The decision was argon2id specifically, so assert it rather than merely
|
||||
// "some hash was produced".
|
||||
if parts := strings.Split(hash, "$"); len(parts) != 6 {
|
||||
t.Fatalf("hash has %d fields, want 6 (PHC format)", len(parts))
|
||||
}
|
||||
|
||||
got, salt, key, err := DecodePasswordHash(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodePasswordHash: %v", err)
|
||||
}
|
||||
if got.Memory != testParams.Memory || got.Time != testParams.Time || got.Threads != testParams.Threads {
|
||||
t.Errorf("decoded params = m=%d,t=%d,p=%d, want m=%d,t=%d,p=%d",
|
||||
got.Memory, got.Time, got.Threads, testParams.Memory, testParams.Time, testParams.Threads)
|
||||
}
|
||||
if len(salt) != int(testParams.SaltLength) {
|
||||
t.Errorf("salt is %d bytes, want %d", len(salt), testParams.SaltLength)
|
||||
}
|
||||
if len(key) != int(testParams.KeyLength) {
|
||||
t.Errorf("key is %d bytes, want %d", len(key), testParams.KeyLength)
|
||||
}
|
||||
// The whole point of the format: the hash must not contain the password.
|
||||
if strings.Contains(hash, goodPassword) {
|
||||
t.Error("the encoded hash contains the plaintext password")
|
||||
}
|
||||
}
|
||||
|
||||
// 2. The correct password verifies.
|
||||
func TestVerifyPasswordAcceptsTheCorrectPassword(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
ok, err := VerifyPassword(hash, goodPassword)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyPassword: %v", err)
|
||||
}
|
||||
if !ok {
|
||||
t.Fatal("the correct password did not verify")
|
||||
}
|
||||
}
|
||||
|
||||
// 3. An incorrect password is rejected — including the near misses that a
|
||||
// sloppy comparison would let through.
|
||||
func TestVerifyPasswordRejectsIncorrectPasswords(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
|
||||
wrong := map[string]string{
|
||||
"different": "incorrect-horse-battery-staple",
|
||||
"empty": "",
|
||||
"prefix": goodPassword[:len(goodPassword)-1],
|
||||
"suffix appended": goodPassword + "x",
|
||||
"case flipped": strings.ToUpper(goodPassword),
|
||||
"whitespace": " " + goodPassword,
|
||||
}
|
||||
for name, candidate := range wrong {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
ok, err := VerifyPassword(hash, candidate)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyPassword returned an error for a wrong password: %v", err)
|
||||
}
|
||||
if ok {
|
||||
t.Error("a wrong password verified")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Different passwords produce different hashes — and so does the SAME
|
||||
// password hashed twice, which is the stronger property and the one that
|
||||
// actually depends on the salt being random.
|
||||
func TestHashPasswordIsSaltedPerCall(t *testing.T) {
|
||||
a, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash a: %v", err)
|
||||
}
|
||||
b, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash b: %v", err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("hashing the same password twice produced identical hashes; the salt is not random")
|
||||
}
|
||||
// Both must still verify: a per-call salt is only useful if it travels
|
||||
// with the hash.
|
||||
for i, h := range []string{a, b} {
|
||||
ok, err := VerifyPassword(h, goodPassword)
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("hash %d did not verify its own password (ok=%v err=%v)", i, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
c, err := HashPasswordWithParams("a-completely-different-password", testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash c: %v", err)
|
||||
}
|
||||
if c == a {
|
||||
t.Error("different passwords produced identical hashes")
|
||||
}
|
||||
// And a hash must not verify a password it was not made from.
|
||||
if ok, _ := VerifyPassword(c, goodPassword); ok {
|
||||
t.Error("a hash verified a password it was not derived from")
|
||||
}
|
||||
}
|
||||
|
||||
// 5a. Empty and out-of-policy passwords are refused at hashing time.
|
||||
func TestHashPasswordRejectsInvalidPasswords(t *testing.T) {
|
||||
cases := map[string]struct {
|
||||
password string
|
||||
want error
|
||||
}{
|
||||
"empty": {"", ErrEmptyPassword},
|
||||
"too short": {strings.Repeat("a", MinPasswordLength-1), ErrPasswordTooShort},
|
||||
"too long": {strings.Repeat("a", MaxPasswordLength+1), ErrPasswordTooLong},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(tc.password, testParams)
|
||||
if err != tc.want {
|
||||
t.Fatalf("error = %v, want %v", err, tc.want)
|
||||
}
|
||||
if hash != "" {
|
||||
t.Error("a hash was returned alongside the error")
|
||||
}
|
||||
// The rejection must not quote the input back.
|
||||
if tc.password != "" && err != nil && strings.Contains(err.Error(), tc.password) {
|
||||
t.Error("the error message contains the password")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The boundary itself is allowed: the rule is "at least MinPasswordLength".
|
||||
if _, err := HashPasswordWithParams(strings.Repeat("a", MinPasswordLength), testParams); err != nil {
|
||||
t.Errorf("a password of exactly the minimum length was rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 5b. A malformed *stored* hash is an error, not a silent "wrong password".
|
||||
// The distinction matters: one is a 401, the other is a 500.
|
||||
func TestVerifyPasswordRejectsMalformedHashes(t *testing.T) {
|
||||
valid, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
fields := strings.Split(valid, "$")
|
||||
|
||||
bad := map[string]string{
|
||||
"empty": "",
|
||||
"not a phc string": "not-a-hash",
|
||||
"bcrypt": "$2a$10$N9qo8uLOickgx2ZMRZoMyeIjZAgcfl7p92ldGxad68LJZdL17lhWy",
|
||||
"argon2i": strings.Replace(valid, "argon2id", "argon2i", 1),
|
||||
"too few fields": strings.Join(fields[:5], "$"),
|
||||
"unparseable params": strings.Replace(valid, fields[3], "m=x,t=y,p=z", 1),
|
||||
"zero memory": strings.Replace(valid, fields[3], "m=0,t=1,p=2", 1),
|
||||
"bad version": strings.Replace(valid, fields[2], "v=notanumber", 1),
|
||||
"salt not base64": strings.Replace(valid, fields[4], "!!!!not base64!!!!", 1),
|
||||
}
|
||||
for name, encoded := range bad {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
ok, err := VerifyPassword(encoded, goodPassword)
|
||||
if err == nil {
|
||||
t.Fatal("a malformed hash verified without an error")
|
||||
}
|
||||
if ok {
|
||||
t.Error("a malformed hash reported a successful verification")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// A future argon2 version is reported as its own error, because it is an
|
||||
// upgrade problem rather than corruption.
|
||||
future := strings.Replace(valid, fields[2], "v=99", 1)
|
||||
if _, err := VerifyPassword(future, goodPassword); err != ErrIncompatibleVersion {
|
||||
t.Errorf("error for a future version = %v, want ErrIncompatibleVersion", err)
|
||||
}
|
||||
}
|
||||
|
||||
// The defaults are a security decision, so they are asserted rather than
|
||||
// assumed: OWASP's argon2id recommendation is m=65536 (64 MiB), t=3, p=4.
|
||||
func TestDefaultPasswordParamsMeetOWASP(t *testing.T) {
|
||||
p := DefaultPasswordParams
|
||||
if p.Memory < 64*1024 {
|
||||
t.Errorf("Memory = %d KiB, want at least 65536", p.Memory)
|
||||
}
|
||||
if p.Time < 3 {
|
||||
t.Errorf("Time = %d, want at least 3", p.Time)
|
||||
}
|
||||
if p.Threads < 1 {
|
||||
t.Errorf("Threads = %d, want at least 1", p.Threads)
|
||||
}
|
||||
if p.SaltLength < 16 {
|
||||
t.Errorf("SaltLength = %d, want at least 16", p.SaltLength)
|
||||
}
|
||||
if p.KeyLength < 32 {
|
||||
t.Errorf("KeyLength = %d, want at least 32", p.KeyLength)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNeedsRehash(t *testing.T) {
|
||||
weak, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
if !NeedsRehash(weak, DefaultPasswordParams) {
|
||||
t.Error("a hash below the default cost was not flagged for rehashing")
|
||||
}
|
||||
strong, err := HashPasswordWithParams(goodPassword, DefaultPasswordParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
if NeedsRehash(strong, DefaultPasswordParams) {
|
||||
t.Error("a hash at the default cost was flagged for rehashing")
|
||||
}
|
||||
if !NeedsRehash("not-a-hash", DefaultPasswordParams) {
|
||||
t.Error("an unreadable hash should be flagged for rehashing")
|
||||
}
|
||||
}
|
||||
|
||||
// firstField is used only to report a failure without dumping a whole hash.
|
||||
func firstField(hash string) string {
|
||||
parts := strings.SplitN(hash, "$", 3)
|
||||
if len(parts) < 2 {
|
||||
return hash
|
||||
}
|
||||
return "$" + parts[1] + "$"
|
||||
}
|
||||
376
go-api/internal/auth/schema_test.go
Normal file
376
go-api/internal/auth/schema_test.go
Normal file
@@ -0,0 +1,376 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// These tests drive migration 000004 itself: what it creates, that it can be
|
||||
// rolled back, and that rolling it back and re-applying it leaves the schema
|
||||
// where it started. They also cover the one change 000004 makes to an existing
|
||||
// table — the global unique index on users.email.
|
||||
//
|
||||
// Every one runs in its own throwaway database. Nothing here can reach the
|
||||
// development database: testutil builds the name from its own prefix.
|
||||
|
||||
const migration4Up = "000004_auth_sessions.up.sql"
|
||||
const migration4Down = "000004_auth_sessions.down.sql"
|
||||
|
||||
// 14. Email is unique across the whole table, not merely within one
|
||||
// organization. This is what makes "log in with your email" answerable.
|
||||
func TestUsersEmailIsGloballyUnique(t *testing.T) {
|
||||
f := newFixture(t, "schema_email_unique")
|
||||
|
||||
// A second organization. Under 000001's (org_id, email) key alone, the
|
||||
// insert below would have been perfectly legal.
|
||||
var otherOrg string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Second Org", "second-org").Scan(&otherOrg); err != nil {
|
||||
t.Fatalf("create the second organization: %v", err)
|
||||
}
|
||||
|
||||
_, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "session-owner@example.test", "Impostor")
|
||||
if err == nil {
|
||||
t.Fatal("the same email was accepted in a second organization; login would be ambiguous")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "users_email_global_key") {
|
||||
t.Errorf("the rejection came from %v, want a users_email_global_key violation", err)
|
||||
}
|
||||
|
||||
// citext makes the index case-insensitive, which is what a login form
|
||||
// needs: nobody should be able to register Demo@… beside demo@….
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "SESSION-OWNER@EXAMPLE.TEST", "Impostor"); err == nil {
|
||||
t.Error("a case-variant of an existing email was accepted")
|
||||
}
|
||||
|
||||
// A genuinely different email in the second organization is still fine —
|
||||
// the index constrains duplicates, not multi-tenancy.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "someone-else@example.test", "Colleague"); err != nil {
|
||||
t.Errorf("a distinct email in a second organization was rejected: %v", err)
|
||||
}
|
||||
|
||||
// The index is unique and covers exactly users(email).
|
||||
var isUnique bool
|
||||
var definition string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT i.indisunique, pg_get_indexdef(i.indexrelid)
|
||||
FROM pg_index i
|
||||
JOIN pg_class c ON c.oid = i.indexrelid
|
||||
WHERE c.relname = 'users_email_global_key'`).Scan(&isUnique, &definition); err != nil {
|
||||
t.Fatalf("read users_email_global_key: %v", err)
|
||||
}
|
||||
if !isUnique {
|
||||
t.Error("users_email_global_key is not a unique index")
|
||||
}
|
||||
if !strings.Contains(definition, "(email)") {
|
||||
t.Errorf("users_email_global_key covers %q, want (email)", definition)
|
||||
}
|
||||
}
|
||||
|
||||
// 15 and 17. Applying every migration in order produces the sessions table
|
||||
// this package needs, and leaves everything the earlier migrations built.
|
||||
func TestMigrationUpBuildsTheSessionsSchema(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, "schema_up")
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
// Columns and types.
|
||||
want := map[string]string{
|
||||
"id": "uuid",
|
||||
"user_id": "uuid",
|
||||
"token_hash": "text",
|
||||
"expires_at": "timestamp with time zone",
|
||||
"absolute_expires_at": "timestamp with time zone",
|
||||
"created_date": "timestamp with time zone",
|
||||
"last_seen_at": "timestamp with time zone",
|
||||
}
|
||||
rows, err := pool.Query(ctx,
|
||||
`SELECT column_name, data_type, is_nullable
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = 'public' AND table_name = 'sessions'`)
|
||||
if err != nil {
|
||||
t.Fatalf("read the sessions columns: %v", err)
|
||||
}
|
||||
got := map[string]string{}
|
||||
for rows.Next() {
|
||||
var name, dataType, nullable string
|
||||
if err := rows.Scan(&name, &dataType, &nullable); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
got[name] = dataType
|
||||
// Every column is required. A nullable expires_at would be a session
|
||||
// with no deadline at all.
|
||||
if nullable != "NO" {
|
||||
t.Errorf("sessions.%s is nullable; every session column is required", name)
|
||||
}
|
||||
}
|
||||
rows.Close()
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("read the sessions columns: %v", err)
|
||||
}
|
||||
if len(got) == 0 {
|
||||
t.Fatal("the sessions table does not exist after the migrations")
|
||||
}
|
||||
for name, wantType := range want {
|
||||
if gotType, ok := got[name]; !ok {
|
||||
t.Errorf("sessions.%s is missing", name)
|
||||
} else if gotType != wantType {
|
||||
t.Errorf("sessions.%s is %s, want %s", name, gotType, wantType)
|
||||
}
|
||||
}
|
||||
for name := range got {
|
||||
if _, expected := want[name]; !expected {
|
||||
t.Errorf("sessions has an unexpected column %q", name)
|
||||
}
|
||||
}
|
||||
|
||||
// Constraints: the primary key, the uniqueness that makes a token identify
|
||||
// one session, and the checks that keep a raw token and an immortal
|
||||
// session out of the table.
|
||||
for _, c := range []struct{ name, kind string }{
|
||||
{"sessions_pkey", "p"},
|
||||
{"sessions_token_hash_key", "u"},
|
||||
{"sessions_token_hash_sha256", "c"},
|
||||
{"sessions_absolute_after_created", "c"},
|
||||
{"sessions_within_absolute", "c"},
|
||||
} {
|
||||
var kind string
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT contype::text FROM pg_constraint
|
||||
WHERE conrelid = 'public.sessions'::regclass AND conname = $1`, c.name).Scan(&kind); err != nil {
|
||||
t.Errorf("constraint %s is missing: %v", c.name, err)
|
||||
continue
|
||||
}
|
||||
if kind != c.kind {
|
||||
t.Errorf("constraint %s is of type %q, want %q", c.name, kind, c.kind)
|
||||
}
|
||||
}
|
||||
|
||||
// The foreign key, and that it cascades. ON DELETE NO ACTION here would
|
||||
// mean a deleted user keeps a working session.
|
||||
var fkTarget, onDelete string
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT confrelid::regclass::text, confdeltype::text
|
||||
FROM pg_constraint
|
||||
WHERE conrelid = 'public.sessions'::regclass AND contype = 'f'`).Scan(&fkTarget, &onDelete); err != nil {
|
||||
t.Fatalf("read the sessions foreign key: %v", err)
|
||||
}
|
||||
if fkTarget != "users" {
|
||||
t.Errorf("the foreign key points at %s, want users", fkTarget)
|
||||
}
|
||||
if onDelete != "c" {
|
||||
t.Errorf("the foreign key is ON DELETE %q, want \"c\" (CASCADE)", onDelete)
|
||||
}
|
||||
|
||||
// Indexes. token_hash is indexed by its UNIQUE constraint — that index is
|
||||
// the lookup path — plus the two the sweep and per-user revocation need.
|
||||
indexes := indexNames(ctx, t, pool, "sessions")
|
||||
for _, name := range []string{"sessions_pkey", "sessions_token_hash_key", "sessions_user_idx", "sessions_expires_idx"} {
|
||||
if !contains(indexes, name) {
|
||||
t.Errorf("index %s is missing; sessions has %v", name, indexes)
|
||||
}
|
||||
}
|
||||
|
||||
// 17. The earlier migrations are untouched: the tables they built are all
|
||||
// still here, and so is the org-scoped uniqueness 000001 declared.
|
||||
//
|
||||
// Named rather than counted. This was a count of 18 — the 17 tables from
|
||||
// 000001 plus sessions — which asserted the right thing in the wrong way:
|
||||
// it broke when 000005 added two unrelated tables, and it would have stayed
|
||||
// green if 000004 had dropped one table and added another. Listing them is
|
||||
// both more precise and stable across later migrations.
|
||||
for _, table := range []string{
|
||||
"organizations", "users", "user_preferences", "role_categories",
|
||||
"certifications", "badges", "courses", "learning_paths", "job_postings",
|
||||
"worker_profiles", "job_applications", "ai_interviews", "staff",
|
||||
"assignments", "shift_records", "evidence", "user_activity",
|
||||
"sessions",
|
||||
} {
|
||||
if !tableExists(ctx, t, pool, table) {
|
||||
t.Errorf("%s is missing after the migrations", table)
|
||||
}
|
||||
}
|
||||
var orgScoped int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_constraint
|
||||
WHERE conrelid = 'public.users'::regclass AND conname = 'users_org_email_key'`).Scan(&orgScoped); err != nil {
|
||||
t.Fatalf("read users_org_email_key: %v", err)
|
||||
}
|
||||
if orgScoped != 1 {
|
||||
t.Error("000004 removed users_org_email_key; it should leave the existing table alone")
|
||||
}
|
||||
}
|
||||
|
||||
// 16. 000004 rolls back cleanly, and re-applies afterwards. A migration that
|
||||
// cannot be reversed is a migration nobody can safely deploy.
|
||||
func TestMigration000004IsReversible(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, "schema_down")
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
if !tableExists(ctx, t, pool, "sessions") {
|
||||
t.Fatal("sessions does not exist before the rollback")
|
||||
}
|
||||
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Down); err != nil {
|
||||
t.Fatalf("apply %s: %v", migration4Down, err)
|
||||
}
|
||||
|
||||
if tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("sessions survived the rollback")
|
||||
}
|
||||
if indexExists(ctx, t, pool, "users_email_global_key") {
|
||||
t.Error("users_email_global_key survived the rollback")
|
||||
}
|
||||
|
||||
// The rollback must touch nothing else. users is still here, still has the
|
||||
// constraint that predates 000004, and still has its rows.
|
||||
if !tableExists(ctx, t, pool, "users") {
|
||||
t.Fatal("the rollback dropped the users table")
|
||||
}
|
||||
var orgScoped int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_constraint
|
||||
WHERE conrelid = 'public.users'::regclass AND conname = 'users_org_email_key'`).Scan(&orgScoped); err != nil {
|
||||
t.Fatalf("read users_org_email_key: %v", err)
|
||||
}
|
||||
if orgScoped != 1 {
|
||||
t.Error("the rollback removed users_org_email_key, which it did not create")
|
||||
}
|
||||
// With the global index gone, the pre-000004 rule is back in force: the
|
||||
// same email in two organizations is legal again.
|
||||
var orgA, orgB string
|
||||
if err := pool.QueryRow(ctx, `INSERT INTO organizations (name, slug) VALUES ('A','a') RETURNING id::text`).Scan(&orgA); err != nil {
|
||||
t.Fatalf("create org A: %v", err)
|
||||
}
|
||||
if err := pool.QueryRow(ctx, `INSERT INTO organizations (name, slug) VALUES ('B','b') RETURNING id::text`).Scan(&orgB); err != nil {
|
||||
t.Fatalf("create org B: %v", err)
|
||||
}
|
||||
for _, org := range []string{orgA, orgB} {
|
||||
if _, err := pool.Exec(ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, '')`,
|
||||
org, "shared@example.test"); err != nil {
|
||||
t.Fatalf("insert after the rollback: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Re-applying now must fail, because the data violates the uniqueness the
|
||||
// migration is about to declare — and it must fail without half-applying.
|
||||
// That is the honest behaviour: the operator has duplicates to resolve.
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Up); err == nil {
|
||||
t.Fatal("000004 applied over duplicate emails; the index would not be unique")
|
||||
}
|
||||
if tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("the failed migration left the sessions table behind; it is not atomic")
|
||||
}
|
||||
|
||||
// Resolve the duplicate and it applies cleanly, restoring exactly what the
|
||||
// rollback removed.
|
||||
if _, err := pool.Exec(ctx, `DELETE FROM users WHERE org_id = $1::uuid`, orgB); err != nil {
|
||||
t.Fatalf("remove the duplicate: %v", err)
|
||||
}
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Up); err != nil {
|
||||
t.Fatalf("re-apply %s: %v", migration4Up, err)
|
||||
}
|
||||
if !tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("sessions did not come back")
|
||||
}
|
||||
if !indexExists(ctx, t, pool, "users_email_global_key") {
|
||||
t.Error("users_email_global_key did not come back")
|
||||
}
|
||||
}
|
||||
|
||||
// Every migration has a matching down file, so any of them can be reversed.
|
||||
func TestEveryMigrationHasADownFile(t *testing.T) {
|
||||
ups := testutil.MigrationFiles(t, ".up.sql")
|
||||
downs := testutil.MigrationFiles(t, ".down.sql")
|
||||
if len(ups) != len(downs) {
|
||||
t.Fatalf("%d up migrations and %d down migrations", len(ups), len(downs))
|
||||
}
|
||||
for i, up := range ups {
|
||||
want := strings.TrimSuffix(up, ".up.sql") + ".down.sql"
|
||||
if downs[i] != want {
|
||||
t.Errorf("%s has no matching down migration (found %s)", up, downs[i])
|
||||
}
|
||||
}
|
||||
// 000004 is still present, with its pair. This used to also assert that it
|
||||
// was the NEWEST and that there were exactly four migrations — a snapshot
|
||||
// that this test's own comment predicted would need revisiting, and which
|
||||
// 000005 duly broke. The count belongs to whichever phase added the newest
|
||||
// migration (see TestMigrationPairsIncluding000005), so it is asserted
|
||||
// there and not here. What this phase cares about — that the migration it
|
||||
// added is intact and reversible — is unchanged and still checked.
|
||||
if !contains(ups, migration4Up) {
|
||||
t.Errorf("%s is missing from the migrations directory", migration4Up)
|
||||
}
|
||||
if !contains(downs, migration4Down) {
|
||||
t.Errorf("%s is missing from the migrations directory", migration4Down)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── introspection helpers ──────────────────────────────────────────────── */
|
||||
|
||||
func tableExists(ctx context.Context, t *testing.T, pool *pgxpool.Pool, name string) bool {
|
||||
t.Helper()
|
||||
var reg *string
|
||||
if err := pool.QueryRow(ctx, `SELECT to_regclass('public.' || $1)::text`, name).Scan(®); err != nil {
|
||||
t.Fatalf("to_regclass(%s): %v", name, err)
|
||||
}
|
||||
return reg != nil
|
||||
}
|
||||
|
||||
func indexExists(ctx context.Context, t *testing.T, pool *pgxpool.Pool, name string) bool {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_indexes WHERE schemaname = 'public' AND indexname = $1`,
|
||||
name).Scan(&n); err != nil {
|
||||
t.Fatalf("look up index %s: %v", name, err)
|
||||
}
|
||||
return n > 0
|
||||
}
|
||||
|
||||
func indexNames(ctx context.Context, t *testing.T, pool *pgxpool.Pool, table string) []string {
|
||||
t.Helper()
|
||||
rows, err := pool.Query(ctx,
|
||||
`SELECT indexname FROM pg_indexes WHERE schemaname = 'public' AND tablename = $1`, table)
|
||||
if err != nil {
|
||||
t.Fatalf("list the indexes on %s: %v", table, err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var names []string
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("list the indexes on %s: %v", table, err)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
func contains(haystack []string, needle string) bool {
|
||||
for _, s := range haystack {
|
||||
if s == needle {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
324
go-api/internal/auth/session.go
Normal file
324
go-api/internal/auth/session.go
Normal file
@@ -0,0 +1,324 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Session lifecycle errors.
|
||||
var (
|
||||
// ErrSessionNotFound means no row matched the token hash. It is returned
|
||||
// for an unknown token and for a well-formed token that has been revoked;
|
||||
// a caller must not distinguish the two to the client.
|
||||
ErrSessionNotFound = errors.New("auth: session not found")
|
||||
|
||||
// ErrSessionExpired means a row matched but is no longer valid, by either
|
||||
// the sliding or the absolute deadline.
|
||||
ErrSessionExpired = errors.New("auth: session expired")
|
||||
)
|
||||
|
||||
// Session is one row of the sessions table.
|
||||
//
|
||||
// TokenHash is the SHA-256 of the token, never the token. There is no field
|
||||
// here that can hold the raw secret, by design: the only place it exists after
|
||||
// Issue returns is the caller's cookie.
|
||||
type Session struct {
|
||||
ID string
|
||||
UserID string
|
||||
|
||||
// TokenHash is lowercase hex SHA-256. See HashToken.
|
||||
TokenHash string
|
||||
|
||||
// ExpiresAt is the sliding deadline; it moves forward as the session is
|
||||
// used, never past AbsoluteExpiresAt.
|
||||
ExpiresAt time.Time
|
||||
|
||||
// AbsoluteExpiresAt is fixed when the session is created and never moves.
|
||||
AbsoluteExpiresAt time.Time
|
||||
|
||||
CreatedDate time.Time
|
||||
LastSeenAt time.Time
|
||||
}
|
||||
|
||||
// IsExpired reports whether the session is dead at the given instant, by
|
||||
// either deadline. The absolute one is checked as well as the sliding one
|
||||
// precisely because the sliding one can be moved.
|
||||
func (s Session) IsExpired(now time.Time) bool {
|
||||
return !now.Before(s.ExpiresAt) || !now.Before(s.AbsoluteExpiresAt)
|
||||
}
|
||||
|
||||
// Policy is how long a session lives.
|
||||
//
|
||||
// Two pairs of durations, because "Remember Me" is a different risk than a
|
||||
// session on a shared machine, and because a sliding window alone can be slid
|
||||
// forever.
|
||||
type Policy struct {
|
||||
// IdleLifetime is how long a normal session survives without being used.
|
||||
IdleLifetime time.Duration
|
||||
// AbsoluteLifetime caps a normal session's total life regardless of use.
|
||||
AbsoluteLifetime time.Duration
|
||||
|
||||
// RememberIdleLifetime and RememberAbsoluteLifetime are the same two
|
||||
// bounds for a session created with Remember Me.
|
||||
RememberIdleLifetime time.Duration
|
||||
RememberAbsoluteLifetime time.Duration
|
||||
}
|
||||
|
||||
// DefaultPolicy implements the Phase 3B session-lifetime decision.
|
||||
//
|
||||
// normal login 12 hours idle, capped at 24 hours of total life. Twelve hours
|
||||
// covers a working day; the daily cap means an unattended tab
|
||||
// cannot be slid along indefinitely.
|
||||
// Remember Me 30 days idle, capped at 90 days. The user asked to stay
|
||||
// signed in; 90 days is the point at which they re-prove it.
|
||||
//
|
||||
// The idle bound is what expires an abandoned session. The absolute bound is
|
||||
// what guarantees no session lives forever.
|
||||
var DefaultPolicy = Policy{
|
||||
IdleLifetime: 12 * time.Hour,
|
||||
AbsoluteLifetime: 24 * time.Hour,
|
||||
RememberIdleLifetime: 30 * 24 * time.Hour,
|
||||
RememberAbsoluteLifetime: 90 * 24 * time.Hour,
|
||||
}
|
||||
|
||||
// lifetimes picks the pair that applies to this session.
|
||||
func (p Policy) lifetimes(remember bool) (idle, absolute time.Duration) {
|
||||
if remember {
|
||||
return p.RememberIdleLifetime, p.RememberAbsoluteLifetime
|
||||
}
|
||||
return p.IdleLifetime, p.AbsoluteLifetime
|
||||
}
|
||||
|
||||
func (p Policy) validate() error {
|
||||
pairs := []struct {
|
||||
name string
|
||||
idle, absolute time.Duration
|
||||
}{
|
||||
{"normal", p.IdleLifetime, p.AbsoluteLifetime},
|
||||
{"remember-me", p.RememberIdleLifetime, p.RememberAbsoluteLifetime},
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
if pair.idle <= 0 || pair.absolute <= 0 {
|
||||
return fmt.Errorf("auth: %s session lifetimes must be positive", pair.name)
|
||||
}
|
||||
// An absolute bound below the idle bound would make the idle window
|
||||
// unreachable, which is a configuration mistake rather than a policy.
|
||||
if pair.absolute < pair.idle {
|
||||
return fmt.Errorf("auth: %s absolute lifetime is shorter than its idle lifetime", pair.name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Store is the persistence the session lifecycle needs.
|
||||
//
|
||||
// An interface rather than a concrete type so that the lifecycle rules below
|
||||
// are testable without a database, and so this package does not depend on pgx.
|
||||
// The PostgreSQL implementation is PGStore, in store.go.
|
||||
//
|
||||
// Every method takes a token *hash*. No implementation ever receives a raw
|
||||
// token, which is what makes it structurally impossible to store one.
|
||||
type Store interface {
|
||||
// Create inserts the session and fills in the id the database assigned,
|
||||
// which is why it takes a pointer: the id is generated by the default on
|
||||
// the column, so the caller cannot know it beforehand.
|
||||
Create(ctx context.Context, s *Session) error
|
||||
FindByTokenHash(ctx context.Context, tokenHash string) (Session, error)
|
||||
Touch(ctx context.Context, id string, expiresAt, lastSeenAt time.Time) error
|
||||
Delete(ctx context.Context, id string) error
|
||||
DeleteByTokenHash(ctx context.Context, tokenHash string) error
|
||||
DeleteExpired(ctx context.Context, now time.Time) (int64, error)
|
||||
}
|
||||
|
||||
// Manager applies the session rules over a Store.
|
||||
//
|
||||
// It is the only place that turns a raw token into a hash, and the only place
|
||||
// that decides whether a session is still alive.
|
||||
type Manager struct {
|
||||
store Store
|
||||
policy Policy
|
||||
|
||||
// now is injectable so the expiry rules can be tested at a chosen instant
|
||||
// rather than by sleeping. Production always leaves it as time.Now.
|
||||
now func() time.Time
|
||||
|
||||
// slideThreshold avoids one UPDATE per request. The sliding deadline is
|
||||
// only pushed forward once the session has used up this fraction of its
|
||||
// idle window, so a burst of requests writes at most one row.
|
||||
slideThreshold float64
|
||||
}
|
||||
|
||||
// NewManager builds a Manager. An invalid policy is a programming error and is
|
||||
// reported here rather than at the first login.
|
||||
func NewManager(store Store, policy Policy) (*Manager, error) {
|
||||
if store == nil {
|
||||
return nil, errors.New("auth: session store is required")
|
||||
}
|
||||
if err := policy.validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Manager{store: store, policy: policy, now: time.Now, slideThreshold: 0.5}, nil
|
||||
}
|
||||
|
||||
// WithClock replaces the clock. For tests.
|
||||
func (m *Manager) WithClock(now func() time.Time) *Manager {
|
||||
if now != nil {
|
||||
m.now = now
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// Policy is the lifetime policy in force.
|
||||
func (m *Manager) Policy() Policy { return m.policy }
|
||||
|
||||
// Issue creates a session for a user and returns the raw token exactly once.
|
||||
//
|
||||
// The token is the return value and is never stored: what reaches the database
|
||||
// is HashToken(token). The caller's only job is to put the raw token straight
|
||||
// into an HttpOnly cookie and then forget it — not log it, not echo it in a
|
||||
// JSON body, not put it in a URL.
|
||||
func (m *Manager) Issue(ctx context.Context, userID string, remember bool) (string, Session, error) {
|
||||
if userID == "" {
|
||||
return "", Session{}, errors.New("auth: user id is required")
|
||||
}
|
||||
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
return "", Session{}, err
|
||||
}
|
||||
|
||||
now := m.now().UTC()
|
||||
idle, absolute := m.policy.lifetimes(remember)
|
||||
s := Session{
|
||||
UserID: userID,
|
||||
TokenHash: HashToken(token),
|
||||
ExpiresAt: now.Add(idle),
|
||||
AbsoluteExpiresAt: now.Add(absolute),
|
||||
CreatedDate: now,
|
||||
LastSeenAt: now,
|
||||
}
|
||||
// The idle window can be the longer of the two only through a bad policy,
|
||||
// which validate() rejects; clamping anyway keeps the database CHECK
|
||||
// (expires_at <= absolute_expires_at) from being the thing that notices.
|
||||
if s.ExpiresAt.After(s.AbsoluteExpiresAt) {
|
||||
s.ExpiresAt = s.AbsoluteExpiresAt
|
||||
}
|
||||
|
||||
if err := m.store.Create(ctx, &s); err != nil {
|
||||
return "", Session{}, err
|
||||
}
|
||||
return token, s, nil
|
||||
}
|
||||
|
||||
// Authenticate resolves a raw token to a live session, sliding its expiry.
|
||||
//
|
||||
// It returns ErrSessionNotFound for an unknown token and ErrSessionExpired for
|
||||
// a dead one. Callers must answer the client identically in both cases: which
|
||||
// of the two it was tells an attacker whether a guessed token ever existed.
|
||||
//
|
||||
// An expired row is deleted as it is found, so a session that times out is
|
||||
// gone rather than waiting for the sweep.
|
||||
func (m *Manager) Authenticate(ctx context.Context, token string) (Session, error) {
|
||||
if token == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
|
||||
hash := HashToken(token)
|
||||
s, err := m.store.FindByTokenHash(ctx, hash)
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
|
||||
now := m.now().UTC()
|
||||
if s.IsExpired(now) {
|
||||
// Best effort: failing to delete does not make the session valid.
|
||||
_ = m.store.Delete(ctx, s.ID)
|
||||
return Session{}, ErrSessionExpired
|
||||
}
|
||||
|
||||
if err := m.slide(ctx, &s, now); err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// slide moves the sliding deadline forward, bounded by the absolute one.
|
||||
//
|
||||
// Only once the session is past slideThreshold of its idle window, so a page
|
||||
// that fires ten requests does not fire ten UPDATEs. The idle window is
|
||||
// recovered from the row rather than taken from the policy, so a session keeps
|
||||
// the lifetime it was issued under even if the policy changes underneath it.
|
||||
func (m *Manager) slide(ctx context.Context, s *Session, now time.Time) error {
|
||||
idle := s.ExpiresAt.Sub(s.LastSeenAt)
|
||||
if idle <= 0 {
|
||||
return nil
|
||||
}
|
||||
if now.Sub(s.LastSeenAt) < time.Duration(float64(idle)*m.slideThreshold) {
|
||||
return nil
|
||||
}
|
||||
|
||||
next := now.Add(idle)
|
||||
if next.After(s.AbsoluteExpiresAt) {
|
||||
next = s.AbsoluteExpiresAt
|
||||
}
|
||||
if err := m.store.Touch(ctx, s.ID, next, now); err != nil {
|
||||
return err
|
||||
}
|
||||
s.ExpiresAt = next
|
||||
s.LastSeenAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
// Lookup resolves a raw token without sliding the expiry or deleting anything.
|
||||
//
|
||||
// A read-only Authenticate, for callers that need to inspect a session without
|
||||
// treating the call as activity.
|
||||
func (m *Manager) Lookup(ctx context.Context, token string) (Session, error) {
|
||||
if token == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
s, err := m.store.FindByTokenHash(ctx, HashToken(token))
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
if s.IsExpired(m.now().UTC()) {
|
||||
return Session{}, ErrSessionExpired
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Revoke deletes the session behind a raw token. This is what logout calls.
|
||||
//
|
||||
// Deleting an already-absent session is not an error: logging out twice, or
|
||||
// logging out with a stale cookie, should succeed rather than fail loudly.
|
||||
func (m *Manager) Revoke(ctx context.Context, token string) error {
|
||||
if token == "" {
|
||||
return ErrEmptyToken
|
||||
}
|
||||
err := m.store.DeleteByTokenHash(ctx, HashToken(token))
|
||||
if errors.Is(err, ErrSessionNotFound) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// RevokeID deletes a session by its row id, for callers that already hold one.
|
||||
func (m *Manager) RevokeID(ctx context.Context, id string) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
err := m.store.Delete(ctx, id)
|
||||
if errors.Is(err, ErrSessionNotFound) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Sweep deletes every session that is past either deadline, and reports how
|
||||
// many rows went. Authenticate already removes the expired sessions it meets;
|
||||
// this collects the ones nobody comes back for.
|
||||
func (m *Manager) Sweep(ctx context.Context) (int64, error) {
|
||||
return m.store.DeleteExpired(ctx, m.now().UTC())
|
||||
}
|
||||
488
go-api/internal/auth/session_test.go
Normal file
488
go-api/internal/auth/session_test.go
Normal file
@@ -0,0 +1,488 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The session tests run against a real PostgreSQL database, because what they
|
||||
// are checking is largely the schema: the unique index that makes lookup work,
|
||||
// the CHECK that refuses a raw token, and the ON DELETE CASCADE that stops a
|
||||
// deleted user leaving a live session behind. None of that can be exercised
|
||||
// against an in-memory fake.
|
||||
//
|
||||
// Each test gets its own throwaway database, migrated but not seeded — the
|
||||
// seed fixture has nothing to say about sessions, and skipping it keeps these
|
||||
// tests fast. testutil skips rather than fails when PostgreSQL is absent.
|
||||
|
||||
type fixture struct {
|
||||
pool *pgxpool.Pool
|
||||
orgID string
|
||||
userID string
|
||||
store *PGStore
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func newFixture(t *testing.T, label string) *fixture {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, label)
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
f := &fixture{pool: pool, ctx: ctx, store: NewPGStore(pool)}
|
||||
if err := pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Auth Test Org", "auth-test-org").Scan(&f.orgID); err != nil {
|
||||
t.Fatalf("create organization: %v", err)
|
||||
}
|
||||
f.userID = f.newUser(t, "session-owner@example.test")
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fixture) newUser(t *testing.T, email string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name, role)
|
||||
VALUES ($1::uuid, $2::citext, $3, 'admin') RETURNING id::text`,
|
||||
f.orgID, email, "Session Owner").Scan(&id); err != nil {
|
||||
t.Fatalf("create user %s: %v", email, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func (f *fixture) sessionCount(t *testing.T) int {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := f.pool.QueryRow(f.ctx, `SELECT count(*)::int FROM sessions`).Scan(&n); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// manager builds a Manager over the fixture's store with a clock the test drives.
|
||||
func (f *fixture) manager(t *testing.T, p Policy, clock *time.Time) *Manager {
|
||||
t.Helper()
|
||||
m, err := NewManager(f.store, p)
|
||||
if err != nil {
|
||||
t.Fatalf("NewManager: %v", err)
|
||||
}
|
||||
return m.WithClock(func() time.Time { return *clock })
|
||||
}
|
||||
|
||||
// shortPolicy keeps the arithmetic in these tests small and legible. The
|
||||
// production values are asserted separately, in TestDefaultPolicy.
|
||||
var shortPolicy = Policy{
|
||||
IdleLifetime: time.Hour,
|
||||
AbsoluteLifetime: 3 * time.Hour,
|
||||
RememberIdleLifetime: 24 * time.Hour,
|
||||
RememberAbsoluteLifetime: 72 * time.Hour,
|
||||
}
|
||||
|
||||
// 9. A session is created, and what lands in the database is the hash.
|
||||
func TestSessionCreation(t *testing.T) {
|
||||
f := newFixture(t, "session_create")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
if token == "" {
|
||||
t.Fatal("Issue returned an empty token")
|
||||
}
|
||||
if sess.ID == "" {
|
||||
t.Error("the session was not given an id")
|
||||
}
|
||||
if sess.UserID != f.userID {
|
||||
t.Errorf("UserID = %q, want %q", sess.UserID, f.userID)
|
||||
}
|
||||
if sess.TokenHash != HashToken(token) {
|
||||
t.Error("the session's token hash is not the hash of the returned token")
|
||||
}
|
||||
if sess.TokenHash == token {
|
||||
t.Fatal("the raw token was stored as the hash")
|
||||
}
|
||||
if want := now.Add(shortPolicy.IdleLifetime); !sess.ExpiresAt.Equal(want) {
|
||||
t.Errorf("ExpiresAt = %v, want %v", sess.ExpiresAt, want)
|
||||
}
|
||||
if want := now.Add(shortPolicy.AbsoluteLifetime); !sess.AbsoluteExpiresAt.Equal(want) {
|
||||
t.Errorf("AbsoluteExpiresAt = %v, want %v", sess.AbsoluteExpiresAt, want)
|
||||
}
|
||||
|
||||
// The row itself: exactly one, holding the hash and never the token.
|
||||
var stored string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT token_hash FROM sessions WHERE id = $1::uuid`, sess.ID).Scan(&stored); err != nil {
|
||||
t.Fatalf("read the stored session: %v", err)
|
||||
}
|
||||
if stored != HashToken(token) {
|
||||
t.Error("the stored token_hash is not the hash of the token")
|
||||
}
|
||||
|
||||
// Nothing anywhere in the table equals the token. This is the property the
|
||||
// whole design exists for, so it is asserted against the database and not
|
||||
// against the struct.
|
||||
var leaked int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`, token).Scan(&leaked); err != nil {
|
||||
t.Fatalf("scan for a leaked token: %v", err)
|
||||
}
|
||||
if leaked != 0 {
|
||||
t.Fatal("the raw token is present in the database")
|
||||
}
|
||||
|
||||
// Remember Me gets the longer pair of deadlines.
|
||||
_, remembered, err := m.Issue(f.ctx, f.userID, true)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue with remember: %v", err)
|
||||
}
|
||||
if want := now.Add(shortPolicy.RememberIdleLifetime); !remembered.ExpiresAt.Equal(want) {
|
||||
t.Errorf("Remember Me ExpiresAt = %v, want %v", remembered.ExpiresAt, want)
|
||||
}
|
||||
if want := now.Add(shortPolicy.RememberAbsoluteLifetime); !remembered.AbsoluteExpiresAt.Equal(want) {
|
||||
t.Errorf("Remember Me AbsoluteExpiresAt = %v, want %v", remembered.AbsoluteExpiresAt, want)
|
||||
}
|
||||
}
|
||||
|
||||
// 10. A live session is found by the token, and only by the right token.
|
||||
func TestSessionLookup(t *testing.T) {
|
||||
f := newFixture(t, "session_lookup")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, issued, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
got, err := m.Authenticate(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate: %v", err)
|
||||
}
|
||||
if got.ID != issued.ID || got.UserID != f.userID {
|
||||
t.Errorf("Authenticate returned session %q for user %q, want %q / %q",
|
||||
got.ID, got.UserID, issued.ID, f.userID)
|
||||
}
|
||||
|
||||
// Lookup is the read-only form and must agree.
|
||||
if looked, err := m.Lookup(f.ctx, token); err != nil || looked.ID != issued.ID {
|
||||
t.Errorf("Lookup = %q, %v; want %q, nil", looked.ID, err, issued.ID)
|
||||
}
|
||||
|
||||
// A different, valid-looking token must not resolve. This is the case the
|
||||
// unique index and the hash lookup exist to make hopeless.
|
||||
other, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, other); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("an unknown token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
|
||||
// A malformed token must be refused without a round trip, and an empty one
|
||||
// must be refused before it is hashed at all.
|
||||
if _, err := m.Authenticate(f.ctx, "not-a-real-token"); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a malformed token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, ""); !errors.Is(err, ErrEmptyToken) {
|
||||
t.Errorf("an empty token returned %v, want ErrEmptyToken", err)
|
||||
}
|
||||
|
||||
// The store is keyed by hash and refuses a raw token outright, so a caller
|
||||
// that forgets to hash gets an error rather than a silent miss.
|
||||
if _, err := f.store.FindByTokenHash(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("FindByTokenHash with a raw token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 11. An expired session is rejected, and is cleaned up as it is found.
|
||||
func TestExpiredSessionIsRejected(t *testing.T) {
|
||||
f := newFixture(t, "session_expiry")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
// One second before the deadline it is still good.
|
||||
now = sess.ExpiresAt.Add(-time.Second)
|
||||
if _, err := m.Authenticate(f.ctx, token); err != nil {
|
||||
t.Fatalf("a session one second from expiry was rejected: %v", err)
|
||||
}
|
||||
|
||||
// Exactly at the deadline it is not. The boundary is closed, not open: a
|
||||
// session that expires at 12:00 is dead at 12:00.
|
||||
fresh, err := m.Lookup(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Lookup: %v", err)
|
||||
}
|
||||
now = fresh.ExpiresAt
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionExpired) {
|
||||
t.Fatalf("an expired session returned %v, want ErrSessionExpired", err)
|
||||
}
|
||||
|
||||
// Finding it expired removes it, so the row does not linger until a sweep.
|
||||
if n := f.sessionCount(t); n != 0 {
|
||||
t.Errorf("%d sessions remain after an expired one was authenticated, want 0", n)
|
||||
}
|
||||
// And the second attempt cannot tell the client anything different.
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("re-authenticating a swept session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// The absolute ceiling is what stops a sliding window being slid forever.
|
||||
func TestSessionCannotOutliveItsAbsoluteDeadline(t *testing.T) {
|
||||
f := newFixture(t, "session_absolute")
|
||||
start := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
now := start
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
ceiling := sess.AbsoluteExpiresAt
|
||||
|
||||
// Keep using the session steadily, well inside the idle window each time,
|
||||
// right up to the ceiling. The sliding deadline must never cross it.
|
||||
for _, at := range []time.Duration{50 * time.Minute, 105 * time.Minute, 160 * time.Minute, 175 * time.Minute} {
|
||||
now = start.Add(at)
|
||||
got, err := m.Authenticate(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate at +%v: %v", at, err)
|
||||
}
|
||||
if got.ExpiresAt.After(ceiling) {
|
||||
t.Fatalf("at +%v the sliding deadline %v passed the absolute ceiling %v",
|
||||
at, got.ExpiresAt, ceiling)
|
||||
}
|
||||
if !got.AbsoluteExpiresAt.Equal(ceiling) {
|
||||
t.Fatalf("at +%v the absolute ceiling moved to %v, want %v", at, got.AbsoluteExpiresAt, ceiling)
|
||||
}
|
||||
}
|
||||
|
||||
// At the ceiling the session is over, however recently it was used.
|
||||
now = ceiling
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionExpired) {
|
||||
t.Fatalf("at the absolute ceiling Authenticate returned %v, want ErrSessionExpired", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 12. Revoking a session deletes it — and revoking twice is not an error,
|
||||
// because logging out with a stale cookie should succeed.
|
||||
func TestSessionDeletion(t *testing.T) {
|
||||
f := newFixture(t, "session_delete")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
if err := m.Revoke(f.ctx, token); err != nil {
|
||||
t.Fatalf("Revoke: %v", err)
|
||||
}
|
||||
if n := f.sessionCount(t); n != 0 {
|
||||
t.Errorf("%d sessions remain after revocation, want 0", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a revoked token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if err := m.Revoke(f.ctx, token); err != nil {
|
||||
t.Errorf("revoking twice returned %v, want nil", err)
|
||||
}
|
||||
if err := m.RevokeID(f.ctx, sess.ID); err != nil {
|
||||
t.Errorf("revoking an absent session by id returned %v, want nil", err)
|
||||
}
|
||||
|
||||
// The store's own contract is stricter: it reports what it did.
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, false); err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
var id string
|
||||
if err := f.pool.QueryRow(f.ctx, `SELECT id::text FROM sessions LIMIT 1`).Scan(&id); err != nil {
|
||||
t.Fatalf("read the session id: %v", err)
|
||||
}
|
||||
if err := f.store.Delete(f.ctx, id); err != nil {
|
||||
t.Fatalf("store.Delete: %v", err)
|
||||
}
|
||||
if err := f.store.Delete(f.ctx, id); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("deleting an absent session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 13. Deleting a user deletes their sessions. Without this, a removed account
|
||||
// keeps working until its cookie happens to expire.
|
||||
func TestUserDeletionCascadesToSessions(t *testing.T) {
|
||||
f := newFixture(t, "session_cascade")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
// A second user, so the test can prove the cascade removes one user's
|
||||
// sessions and leaves the other's alone.
|
||||
otherID := f.newUser(t, "other-user@example.test")
|
||||
|
||||
doomedToken, _, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue for the doomed user: %v", err)
|
||||
}
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, true); err != nil {
|
||||
t.Fatalf("second Issue for the doomed user: %v", err)
|
||||
}
|
||||
survivingToken, _, err := m.Issue(f.ctx, otherID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue for the surviving user: %v", err)
|
||||
}
|
||||
if n := f.sessionCount(t); n != 3 {
|
||||
t.Fatalf("%d sessions before the delete, want 3", n)
|
||||
}
|
||||
|
||||
if _, err := f.pool.Exec(f.ctx, `DELETE FROM users WHERE id = $1::uuid`, f.userID); err != nil {
|
||||
// A RESTRICT foreign key would fail here, which is exactly the design
|
||||
// this test rules out.
|
||||
t.Fatalf("delete the user: %v", err)
|
||||
}
|
||||
|
||||
if n := f.sessionCount(t); n != 1 {
|
||||
t.Fatalf("%d sessions after deleting one of two users, want 1", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, doomedToken); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a deleted user's session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, survivingToken); err != nil {
|
||||
t.Errorf("the other user's session was destroyed too: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Sweep collects the sessions nobody comes back for, by either deadline.
|
||||
func TestSweepDeletesExpiredSessions(t *testing.T) {
|
||||
f := newFixture(t, "session_sweep")
|
||||
start := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
now := start
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, false); err != nil { // dies at +1h
|
||||
t.Fatalf("Issue short: %v", err)
|
||||
}
|
||||
longToken, _, err := m.Issue(f.ctx, f.userID, true) // dies at +24h
|
||||
if err != nil {
|
||||
t.Fatalf("Issue long: %v", err)
|
||||
}
|
||||
|
||||
now = start.Add(90 * time.Minute)
|
||||
n, err := m.Sweep(f.ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Sweep: %v", err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Errorf("Sweep removed %d sessions, want 1", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, longToken); err != nil {
|
||||
t.Errorf("Sweep removed a live session: %v", err)
|
||||
}
|
||||
|
||||
now = start.Add(25 * time.Hour)
|
||||
if n, err = m.Sweep(f.ctx); err != nil || n != 1 {
|
||||
t.Errorf("second Sweep removed %d sessions (err %v), want 1", n, err)
|
||||
}
|
||||
if got := f.sessionCount(t); got != 0 {
|
||||
t.Errorf("%d sessions remain after the sweep, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
// The database is the last line of defence against storing a raw token: the
|
||||
// CHECK constraint refuses anything that is not a SHA-256 hex digest, even if
|
||||
// the Go guard were bypassed.
|
||||
func TestDatabaseRefusesARawToken(t *testing.T) {
|
||||
f := newFixture(t, "session_rawtoken")
|
||||
now := time.Now().UTC()
|
||||
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
// Through the store: a Go error naming the mistake.
|
||||
err = f.store.Create(f.ctx, &Session{
|
||||
UserID: f.userID, TokenHash: token,
|
||||
ExpiresAt: now.Add(time.Hour), AbsoluteExpiresAt: now.Add(2 * time.Hour),
|
||||
CreatedDate: now, LastSeenAt: now,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("the store accepted a raw token as a token hash")
|
||||
}
|
||||
|
||||
// Straight past the store, in SQL: the constraint still refuses it.
|
||||
_, err = f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, token)
|
||||
if err == nil {
|
||||
t.Fatal("the sessions table accepted a raw token; the CHECK constraint is not doing its job")
|
||||
}
|
||||
|
||||
// The same insert with a proper hash succeeds, so the constraint is not
|
||||
// simply rejecting everything.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, HashToken(token)); err != nil {
|
||||
t.Fatalf("a well-formed session was refused: %v", err)
|
||||
}
|
||||
|
||||
// And the same hash cannot be stored twice: UNIQUE is what makes a token
|
||||
// identify exactly one session.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, HashToken(token)); err == nil {
|
||||
t.Fatal("two sessions were stored with the same token hash")
|
||||
}
|
||||
}
|
||||
|
||||
// The lifetimes are a decision, so they are asserted rather than assumed.
|
||||
func TestDefaultPolicy(t *testing.T) {
|
||||
p := DefaultPolicy
|
||||
if p.IdleLifetime != 12*time.Hour {
|
||||
t.Errorf("IdleLifetime = %v, want 12h", p.IdleLifetime)
|
||||
}
|
||||
if p.RememberIdleLifetime != 30*24*time.Hour {
|
||||
t.Errorf("RememberIdleLifetime = %v, want 720h (30 days)", p.RememberIdleLifetime)
|
||||
}
|
||||
// The point of the absolute bound: it must exist and must exceed the
|
||||
// window it caps, or a session could live forever.
|
||||
if p.AbsoluteLifetime <= 0 || p.RememberAbsoluteLifetime <= 0 {
|
||||
t.Fatal("an absolute lifetime is unset; a session could live forever")
|
||||
}
|
||||
if err := p.validate(); err != nil {
|
||||
t.Errorf("the default policy is not valid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewManagerRejectsBadInput(t *testing.T) {
|
||||
if _, err := NewManager(nil, DefaultPolicy); err == nil {
|
||||
t.Error("NewManager accepted a nil store")
|
||||
}
|
||||
bad := map[string]Policy{
|
||||
"zero": {},
|
||||
"absolute below idle": {IdleLifetime: 2 * time.Hour, AbsoluteLifetime: time.Hour, RememberIdleLifetime: time.Hour, RememberAbsoluteLifetime: time.Hour},
|
||||
"remember-me unset": {IdleLifetime: time.Hour, AbsoluteLifetime: time.Hour},
|
||||
"negative idle lifetime": {IdleLifetime: -time.Hour, AbsoluteLifetime: time.Hour, RememberIdleLifetime: time.Hour, RememberAbsoluteLifetime: time.Hour},
|
||||
}
|
||||
for name, p := range bad {
|
||||
if _, err := NewManager(NewPGStore(nil), p); err == nil {
|
||||
t.Errorf("NewManager accepted the %s policy", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
204
go-api/internal/auth/store.go
Normal file
204
go-api/internal/auth/store.go
Normal file
@@ -0,0 +1,204 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
)
|
||||
|
||||
// Querier is satisfied by *pgxpool.Pool and by pgx.Tx, so every method below
|
||||
// works inside or outside a transaction.
|
||||
//
|
||||
// Declared here rather than imported from internal/repo: that package's
|
||||
// Querier is identical, but this one keeps the authentication foundation
|
||||
// independent of the resource/descriptor layer, which it otherwise shares
|
||||
// nothing with. Go interfaces are structural, so both are satisfied by the
|
||||
// same values.
|
||||
type Querier interface {
|
||||
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
|
||||
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
||||
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
|
||||
}
|
||||
|
||||
// PGStore is the sessions table.
|
||||
//
|
||||
// Every statement here is a constant string with bind parameters. Nothing —
|
||||
// not an id, not a token hash, not a timestamp — is ever formatted into SQL.
|
||||
// There is no identifier taken from a caller, so there is nothing to quote and
|
||||
// nothing to escape.
|
||||
type PGStore struct {
|
||||
db Querier
|
||||
}
|
||||
|
||||
// NewPGStore builds the store over a pool or a transaction.
|
||||
func NewPGStore(db Querier) *PGStore { return &PGStore{db: db} }
|
||||
|
||||
// Compile-time check that the persistence layer satisfies the lifecycle's
|
||||
// expectations. If Store gains a method, this line is where it is noticed.
|
||||
var _ Store = (*PGStore)(nil)
|
||||
|
||||
// sessionColumns is the projection every read below shares, in the order the
|
||||
// scan expects.
|
||||
const sessionColumns = `id::text, user_id::text, token_hash,
|
||||
expires_at, absolute_expires_at, created_date, last_seen_at`
|
||||
|
||||
// Create inserts a session.
|
||||
//
|
||||
// The id is left to the database's gen_random_uuid() default when the caller
|
||||
// did not choose one, and returned so the caller's Session is complete.
|
||||
// created_date and last_seen_at come from the caller rather than now(), so the
|
||||
// row agrees with the deadlines the Manager computed from the same instant.
|
||||
func (s *PGStore) Create(ctx context.Context, sess *Session) error {
|
||||
if sess.UserID == "" {
|
||||
return errors.New("auth: session user id is required")
|
||||
}
|
||||
// The last line of defence against writing a raw token to disk. The
|
||||
// database CHECK enforces the same shape; this turns it into a Go error
|
||||
// naming the actual mistake instead of a constraint violation.
|
||||
if !IsTokenHash(sess.TokenHash) {
|
||||
return errors.New("auth: session token_hash is not a SHA-256 hex digest")
|
||||
}
|
||||
|
||||
const q = `INSERT INTO sessions
|
||||
(id, user_id, token_hash, expires_at, absolute_expires_at, created_date, last_seen_at)
|
||||
VALUES
|
||||
(COALESCE($1::uuid, gen_random_uuid()), $2::uuid, $3::text,
|
||||
$4::timestamptz, $5::timestamptz, $6::timestamptz, $7::timestamptz)
|
||||
RETURNING id::text`
|
||||
|
||||
var id *string
|
||||
if sess.ID != "" {
|
||||
id = &sess.ID
|
||||
}
|
||||
if err := s.db.QueryRow(ctx, q, id, sess.UserID, sess.TokenHash,
|
||||
sess.ExpiresAt, sess.AbsoluteExpiresAt, sess.CreatedDate, sess.LastSeenAt,
|
||||
).Scan(&sess.ID); err != nil {
|
||||
return fmt.Errorf("auth: create session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindByTokenHash reads one session by the hash of its token.
|
||||
//
|
||||
// The parameter is a hash, never a token: the Manager hashes before it calls
|
||||
// here, so a raw secret never reaches the query layer at all. Returns
|
||||
// ErrSessionNotFound when no row matches, which callers must not distinguish
|
||||
// from an expired session when answering a client.
|
||||
//
|
||||
// Expiry is deliberately not filtered in SQL. The caller decides what an
|
||||
// expired row means — Authenticate deletes it, a diagnostic might report it —
|
||||
// and a WHERE clause here would collapse "revoked" and "timed out" into one
|
||||
// indistinguishable answer at the wrong layer.
|
||||
func (s *PGStore) FindByTokenHash(ctx context.Context, tokenHash string) (Session, error) {
|
||||
if tokenHash == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
if !IsTokenHash(tokenHash) {
|
||||
// A value of the wrong shape cannot match any row, and querying with
|
||||
// it would be an unnecessary round trip on every malformed cookie.
|
||||
return Session{}, ErrSessionNotFound
|
||||
}
|
||||
|
||||
const q = `SELECT ` + sessionColumns + ` FROM sessions WHERE token_hash = $1::text`
|
||||
|
||||
var out Session
|
||||
err := s.db.QueryRow(ctx, q, tokenHash).Scan(
|
||||
&out.ID, &out.UserID, &out.TokenHash,
|
||||
&out.ExpiresAt, &out.AbsoluteExpiresAt, &out.CreatedDate, &out.LastSeenAt)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return Session{}, ErrSessionNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return Session{}, fmt.Errorf("auth: find session: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Touch moves the sliding deadline and records the activity.
|
||||
//
|
||||
// The UPDATE is guarded by `expires_at <= absolute_expires_at` in the database
|
||||
// CHECK; the Manager clamps before calling, so a violation here would mean a
|
||||
// bug rather than a race.
|
||||
func (s *PGStore) Touch(ctx context.Context, id string, expiresAt, lastSeenAt time.Time) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
|
||||
const q = `UPDATE sessions
|
||||
SET expires_at = $2::timestamptz, last_seen_at = $3::timestamptz
|
||||
WHERE id = $1::uuid`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, id, expiresAt, lastSeenAt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: touch session: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
// The session was revoked between the read and this write. Reporting
|
||||
// it as absent is honest; the caller treats it as a failed lookup.
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete removes one session by id. Returns ErrSessionNotFound if there was
|
||||
// nothing to remove; Manager.RevokeID absorbs that, because logging out of a
|
||||
// session that is already gone is a success.
|
||||
func (s *PGStore) Delete(ctx context.Context, id string) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
|
||||
const q = `DELETE FROM sessions WHERE id = $1::uuid`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: delete session: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteByTokenHash removes one session by the hash of its token. This is the
|
||||
// logout path: the cookie is all the client has.
|
||||
func (s *PGStore) DeleteByTokenHash(ctx context.Context, tokenHash string) error {
|
||||
if tokenHash == "" {
|
||||
return ErrEmptyToken
|
||||
}
|
||||
if !IsTokenHash(tokenHash) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
const q = `DELETE FROM sessions WHERE token_hash = $1::text`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, tokenHash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: delete session by token: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteExpired removes every session past either deadline and reports the
|
||||
// count.
|
||||
//
|
||||
// Both deadlines are tested. Filtering on expires_at alone would leave behind
|
||||
// a session whose sliding window is still open but whose absolute ceiling has
|
||||
// passed — precisely the row the absolute bound exists to kill.
|
||||
func (s *PGStore) DeleteExpired(ctx context.Context, now time.Time) (int64, error) {
|
||||
const q = `DELETE FROM sessions
|
||||
WHERE expires_at <= $1::timestamptz OR absolute_expires_at <= $1::timestamptz`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, now)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("auth: delete expired sessions: %w", err)
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
74
go-api/internal/auth/token.go
Normal file
74
go-api/internal/auth/token.go
Normal file
@@ -0,0 +1,74 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
)
|
||||
|
||||
// ErrEmptyToken is returned when a token is required and none was supplied.
|
||||
// It is deliberately distinct from "no such session": an absent cookie is a
|
||||
// different situation from a cookie that no longer matches a row.
|
||||
var ErrEmptyToken = errors.New("auth: session token is empty")
|
||||
|
||||
// TokenBytes is the entropy behind a session token.
|
||||
//
|
||||
// 32 bytes — 256 bits — is the size of the SHA-256 output it is hashed to, so
|
||||
// nothing is wasted at either end, and it puts guessing a live session far
|
||||
// beyond reach: an attacker who could test a billion candidates a second would
|
||||
// still need on the order of 10^60 years.
|
||||
//
|
||||
// This is a raw byte count, not a character count. The encoded token is 43
|
||||
// characters of base64url.
|
||||
const TokenBytes = 32
|
||||
|
||||
// tokenHashPattern is the exact shape stored in sessions.token_hash, and the
|
||||
// same pattern the sessions_token_hash_sha256 CHECK constraint enforces in
|
||||
// migration 000004. Validating here turns a database constraint violation into
|
||||
// a clear Go error at the point the mistake was made.
|
||||
var tokenHashPattern = regexp.MustCompile(`^[0-9a-f]{64}$`)
|
||||
|
||||
// GenerateToken returns a new, cryptographically random session token.
|
||||
//
|
||||
// The returned string is the secret itself. It is what goes into the HttpOnly
|
||||
// cookie and it must never be written to the database, to a log, or to an
|
||||
// error message. Only its hash is persisted — see HashToken.
|
||||
//
|
||||
// base64url without padding, so the value is safe in a cookie, a header and a
|
||||
// URL without escaping, and contains no '=' to be mangled by a cookie parser.
|
||||
func GenerateToken() (string, error) {
|
||||
buf := make([]byte, TokenBytes)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
// There is no fallback. math/rand here would produce tokens an
|
||||
// attacker can predict from a handful of observed sessions.
|
||||
return "", fmt.Errorf("auth: read random bytes: %w", err)
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
// HashToken returns the lowercase hex SHA-256 of a session token.
|
||||
//
|
||||
// This is what the database stores. A plain hash — not argon2 — is the right
|
||||
// choice here and the wrong one for a password, and the difference is entropy:
|
||||
// a session token is 256 uniformly random bits, so there is no dictionary to
|
||||
// run against it and no work factor worth paying on every single request. A
|
||||
// password is chosen by a human and needs argon2id precisely because it is not.
|
||||
//
|
||||
// The function is pure and deterministic: the same token always hashes to the
|
||||
// same string, which is what makes lookup by hash possible at all.
|
||||
func HashToken(token string) string {
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// IsTokenHash reports whether s has the shape HashToken produces.
|
||||
//
|
||||
// Used to catch the one mistake that would be catastrophic and silent: passing
|
||||
// a raw token where a hash is expected, and storing the secret in plaintext.
|
||||
// A raw token is base64url and contains characters outside [0-9a-f], or is the
|
||||
// wrong length, so it always fails this test.
|
||||
func IsTokenHash(s string) bool { return tokenHashPattern.MatchString(s) }
|
||||
158
go-api/internal/auth/token_test.go
Normal file
158
go-api/internal/auth/token_test.go
Normal file
@@ -0,0 +1,158 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// 6. Session tokens carry real entropy from crypto/rand.
|
||||
//
|
||||
// Randomness cannot be proved by a test, so this asserts the properties whose
|
||||
// absence would mean the generator is broken: the full 256 bits are present,
|
||||
// the output is not a constant, and the bytes are not all the same value —
|
||||
// which is what a zeroed or unseeded buffer looks like.
|
||||
func TestGenerateTokenIsCryptographicallyRandom(t *testing.T) {
|
||||
const runs = 512
|
||||
|
||||
seen := make(map[string]struct{}, runs)
|
||||
bitsSet := make([]int, 8*TokenBytes) // how often each bit position was 1
|
||||
|
||||
for i := 0; i < runs; i++ {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
raw, err := base64.RawURLEncoding.DecodeString(token)
|
||||
if err != nil {
|
||||
t.Fatalf("token is not base64url: %v", err)
|
||||
}
|
||||
if len(raw) != TokenBytes {
|
||||
t.Fatalf("token decodes to %d bytes, want %d", len(raw), TokenBytes)
|
||||
}
|
||||
// A cookie value must survive a round trip untouched: no '=' padding,
|
||||
// no '+' or '/' to be re-encoded.
|
||||
if strings.ContainsAny(token, "=+/") {
|
||||
t.Fatalf("token contains a character that is unsafe in a cookie or URL")
|
||||
}
|
||||
|
||||
if _, dup := seen[token]; dup {
|
||||
t.Fatalf("GenerateToken returned a duplicate within %d calls", runs)
|
||||
}
|
||||
seen[token] = struct{}{}
|
||||
|
||||
for bit := 0; bit < 8*TokenBytes; bit++ {
|
||||
if raw[bit/8]&(1<<(bit%8)) != 0 {
|
||||
bitsSet[bit]++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Each bit should be 1 about half the time. A bit that is *always* 0 or
|
||||
// always 1 across 512 draws has a chance of roughly 2^-511 of being random
|
||||
// and is far more likely a stuck generator. The bound is deliberately
|
||||
// loose — this is a smoke test for a broken source, not a statistical
|
||||
// suite, and it must never flake.
|
||||
for bit, count := range bitsSet {
|
||||
if count == 0 || count == runs {
|
||||
t.Errorf("bit %d was constant across %d tokens; the entropy source is broken", bit, runs)
|
||||
}
|
||||
}
|
||||
|
||||
// 256 bits is the size that makes guessing a live session hopeless.
|
||||
if TokenBytes < 32 {
|
||||
t.Errorf("TokenBytes = %d, want at least 32", TokenBytes)
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Two generated tokens differ.
|
||||
func TestGenerateTokenReturnsDistinctValues(t *testing.T) {
|
||||
a, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("first: %v", err)
|
||||
}
|
||||
b, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("second: %v", err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("two consecutive tokens were identical")
|
||||
}
|
||||
if HashToken(a) == HashToken(b) {
|
||||
t.Fatal("two distinct tokens hashed to the same value")
|
||||
}
|
||||
}
|
||||
|
||||
// 8. Hashing is deterministic, and is genuinely SHA-256 rather than something
|
||||
// that merely looks like it.
|
||||
func TestHashTokenIsDeterministicSHA256(t *testing.T) {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
first, second := HashToken(token), HashToken(token)
|
||||
if first != second {
|
||||
t.Fatal("hashing the same token twice produced different values")
|
||||
}
|
||||
|
||||
// Checked against the standard library directly: lookup only works if the
|
||||
// value stored is exactly this.
|
||||
want := sha256.Sum256([]byte(token))
|
||||
if first != hex.EncodeToString(want[:]) {
|
||||
t.Fatal("HashToken does not agree with crypto/sha256")
|
||||
}
|
||||
|
||||
if len(first) != 64 {
|
||||
t.Fatalf("hash is %d characters, want 64 hex characters", len(first))
|
||||
}
|
||||
if first != strings.ToLower(first) {
|
||||
t.Error("hash is not lowercase; the database CHECK requires lowercase hex")
|
||||
}
|
||||
// The stored value must not be the secret.
|
||||
if strings.Contains(first, token) || first == token {
|
||||
t.Error("the hash contains the token")
|
||||
}
|
||||
if HashToken(token+"x") == first {
|
||||
t.Error("a different token hashed to the same value")
|
||||
}
|
||||
|
||||
// The empty string has a hash too — that is a property of SHA-256, not a
|
||||
// licence to store one. Guarding against an empty token is the Manager's
|
||||
// job, and TestManagerRejectsEmptyToken covers it.
|
||||
if HashToken("") == "" {
|
||||
t.Error("HashToken returned an empty string")
|
||||
}
|
||||
}
|
||||
|
||||
// IsTokenHash is the guard that stops a raw token being written where a hash
|
||||
// belongs, so it must reject every raw token and accept every real hash.
|
||||
func TestIsTokenHash(t *testing.T) {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
if !IsTokenHash(HashToken(token)) {
|
||||
t.Error("a real hash was not recognised as one")
|
||||
}
|
||||
if IsTokenHash(token) {
|
||||
t.Error("a raw token was accepted as a hash; this is the check that prevents storing the secret")
|
||||
}
|
||||
|
||||
for name, s := range map[string]string{
|
||||
"empty": "",
|
||||
"too short": strings.Repeat("a", 63),
|
||||
"too long": strings.Repeat("a", 65),
|
||||
"uppercase": strings.ToUpper(HashToken(token)),
|
||||
"non-hex": strings.Repeat("g", 64),
|
||||
"trailing": HashToken(token) + "\n",
|
||||
} {
|
||||
if IsTokenHash(s) {
|
||||
t.Errorf("%s was accepted as a token hash", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
142
go-api/internal/auth/users.go
Normal file
142
go-api/internal/auth/users.go
Normal file
@@ -0,0 +1,142 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
// ErrUserNotFound means no user matched. Callers authenticating someone must
|
||||
// answer this identically to a wrong password — see the note on Credentials.
|
||||
var ErrUserNotFound = errors.New("auth: user not found")
|
||||
|
||||
// StatusActive is the only users.status that may hold a session. The column's
|
||||
// CHECK constraint (migration 000001) permits 'active' and 'suspended'.
|
||||
const StatusActive = "active"
|
||||
|
||||
// User is the part of a users row that authentication needs.
|
||||
//
|
||||
// Note what is absent: preferences, timestamps, legacy ids. This is not a
|
||||
// general user model — the resource layer already has one — it is the set of
|
||||
// facts required to answer "may this person hold a session, and whose data do
|
||||
// they see".
|
||||
type User struct {
|
||||
ID string
|
||||
OrgID string
|
||||
Email string
|
||||
FullName string
|
||||
Role string
|
||||
AccountType string
|
||||
Status string
|
||||
|
||||
// PasswordHash is empty when the user has never had a password set.
|
||||
// migration 000001 leaves the column nullable and NULL, and the seeded
|
||||
// demo user is in exactly that state until `setpassword` is run.
|
||||
PasswordHash string
|
||||
}
|
||||
|
||||
// IsActive reports whether this user may authenticate or hold a session.
|
||||
func (u User) IsActive() bool { return u.Status == StatusActive }
|
||||
|
||||
// CanAuthenticate reports whether a password check is even possible. A user
|
||||
// with no password hash cannot sign in, and must be refused in the same way
|
||||
// and with the same timing as a wrong password.
|
||||
func (u User) CanAuthenticate() bool { return u.IsActive() && u.PasswordHash != "" }
|
||||
|
||||
// UserStore is the read side of authentication.
|
||||
//
|
||||
// Deliberately read-only apart from MarkLoggedIn: creating and editing users is
|
||||
// the resource layer's job, and nothing in the sign-in path should be able to
|
||||
// write a role, a status or an organization.
|
||||
type UserStore interface {
|
||||
// FindByEmail resolves the login identifier. Email is citext and globally
|
||||
// unique (migration 000004), so this returns at most one row.
|
||||
FindByEmail(ctx context.Context, email string) (User, error)
|
||||
// FindByID resolves the user behind a session on every request.
|
||||
FindByID(ctx context.Context, id string) (User, error)
|
||||
// MarkLoggedIn records a successful sign-in.
|
||||
MarkLoggedIn(ctx context.Context, id string, at time.Time) error
|
||||
}
|
||||
|
||||
// PGUserStore reads users from PostgreSQL.
|
||||
type PGUserStore struct {
|
||||
db Querier
|
||||
}
|
||||
|
||||
// NewPGUserStore builds the store over a pool or a transaction.
|
||||
func NewPGUserStore(db Querier) *PGUserStore { return &PGUserStore{db: db} }
|
||||
|
||||
var _ UserStore = (*PGUserStore)(nil)
|
||||
|
||||
// userColumns is the projection both lookups share.
|
||||
//
|
||||
// password_hash is COALESCEd to the empty string rather than scanned into a
|
||||
// *string: a NULL hash and an empty hash mean the same thing here — no password
|
||||
// is set — and collapsing them at the edge means no caller has to remember to
|
||||
// nil-check before handing the value to VerifyPassword.
|
||||
const userColumns = `id::text, org_id::text, email::text, full_name,
|
||||
role, account_type, status, COALESCE(password_hash, '')`
|
||||
|
||||
func scanUser(row pgx.Row) (User, error) {
|
||||
var u User
|
||||
err := row.Scan(&u.ID, &u.OrgID, &u.Email, &u.FullName,
|
||||
&u.Role, &u.AccountType, &u.Status, &u.PasswordHash)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return User{}, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// FindByEmail looks a user up by their login identifier.
|
||||
//
|
||||
// The comparison is against the citext column, so it is case-insensitive:
|
||||
// "Demo@Krow.app" finds the same row as "demo@krow.app", which is what a person
|
||||
// typing their own address at a login form expects. Trimming is the caller's
|
||||
// job and is done in the handler, where the raw input is.
|
||||
func (s *PGUserStore) FindByEmail(ctx context.Context, email string) (User, error) {
|
||||
if email == "" {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
const q = `SELECT ` + userColumns + ` FROM users WHERE email = $1::citext`
|
||||
u, err := scanUser(s.db.QueryRow(ctx, q, email))
|
||||
if err != nil && !errors.Is(err, ErrUserNotFound) {
|
||||
return User{}, fmt.Errorf("auth: find user by email: %w", err)
|
||||
}
|
||||
return u, err
|
||||
}
|
||||
|
||||
// FindByID resolves the user behind a session.
|
||||
//
|
||||
// This runs on every authenticated request, which is why status is read here
|
||||
// rather than cached in the session row: suspending an account must take effect
|
||||
// on the next request, not whenever the session happens to expire.
|
||||
func (s *PGUserStore) FindByID(ctx context.Context, id string) (User, error) {
|
||||
if id == "" {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
const q = `SELECT ` + userColumns + ` FROM users WHERE id = $1::uuid`
|
||||
u, err := scanUser(s.db.QueryRow(ctx, q, id))
|
||||
if err != nil && !errors.Is(err, ErrUserNotFound) {
|
||||
return User{}, fmt.Errorf("auth: find user by id: %w", err)
|
||||
}
|
||||
return u, err
|
||||
}
|
||||
|
||||
// MarkLoggedIn stamps last_login_at.
|
||||
//
|
||||
// Deliberately not part of the transaction that creates the session: a failure
|
||||
// to record the timestamp is a lost diagnostic, not a reason to refuse a
|
||||
// sign-in that has already succeeded on its merits.
|
||||
func (s *PGUserStore) MarkLoggedIn(ctx context.Context, id string, at time.Time) error {
|
||||
const q = `UPDATE users SET last_login_at = $2::timestamptz WHERE id = $1::uuid`
|
||||
if _, err := s.db.Exec(ctx, q, id, at); err != nil {
|
||||
return fmt.Errorf("auth: mark logged in: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
68
go-api/internal/authctx/authctx.go
Normal file
68
go-api/internal/authctx/authctx.go
Normal file
@@ -0,0 +1,68 @@
|
||||
// Package authctx carries the authenticated identity of a request.
|
||||
//
|
||||
// It is the successor to the development identity that used to be injected by
|
||||
// httpserver.devOrgMiddleware. The difference is not the shape — both put a
|
||||
// value on the request context — but the provenance: everything here was read
|
||||
// out of a server-side session row, and nothing in it can be influenced by the
|
||||
// request that carries it.
|
||||
//
|
||||
// That is the whole point of the package existing separately from the handlers.
|
||||
// A handler that wants to know who is calling has exactly one place to ask, and
|
||||
// that place cannot be reached from a request body, a query string or a header.
|
||||
// There is deliberately no setter that takes a user id from a client.
|
||||
package authctx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
type key struct{}
|
||||
|
||||
// ErrNoIdentity means a protected operation was reached without an
|
||||
// authenticated identity. That is a routing or middleware bug rather than a
|
||||
// client error: an unauthenticated request should have been refused before it
|
||||
// got this far.
|
||||
var ErrNoIdentity = errors.New("no authenticated identity in context")
|
||||
|
||||
// Identity is who the request is, as resolved from the session row.
|
||||
//
|
||||
// Role is carried because Phase 3D will need it, and because carrying it now
|
||||
// means the middleware reads it once per request instead of every future
|
||||
// authorization check re-querying the user. It is NOT consulted anywhere in
|
||||
// Phase 3C: authentication only.
|
||||
type Identity struct {
|
||||
UserID string
|
||||
OrgID string
|
||||
Email string
|
||||
FullName string
|
||||
Role string
|
||||
AccountType string
|
||||
Status string
|
||||
|
||||
// SessionID is the row this identity came from, so logout and per-session
|
||||
// diagnostics do not have to re-hash the cookie.
|
||||
SessionID string
|
||||
// ExpiresAt is the session's sliding deadline as of this request.
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// With returns a context carrying the authenticated identity.
|
||||
func With(ctx context.Context, id Identity) context.Context {
|
||||
return context.WithValue(ctx, key{}, id)
|
||||
}
|
||||
|
||||
// From reads the identity, reporting whether one was present.
|
||||
func From(ctx context.Context) (Identity, bool) {
|
||||
v, ok := ctx.Value(key{}).(Identity)
|
||||
return v, ok && v.UserID != ""
|
||||
}
|
||||
|
||||
// MustFrom reads the identity or returns ErrNoIdentity.
|
||||
func MustFrom(ctx context.Context) (Identity, error) {
|
||||
if v, ok := From(ctx); ok {
|
||||
return v, nil
|
||||
}
|
||||
return Identity{}, ErrNoIdentity
|
||||
}
|
||||
329
go-api/internal/config/config.go
Normal file
329
go-api/internal/config/config.go
Normal file
@@ -0,0 +1,329 @@
|
||||
// Package config loads and validates the backend's runtime configuration.
|
||||
//
|
||||
// Configuration comes from the process environment. A .env file in the
|
||||
// repository root is read first as a convenience for local development, and
|
||||
// never overrides a variable that is already set — so an explicit
|
||||
// `DATABASE_PASSWORD=… go run ./cmd/api` always wins over the file.
|
||||
//
|
||||
// Nothing here has a credential baked in. Load fails loudly rather than
|
||||
// falling back to a default host, database or user, because a silent default
|
||||
// is how a development process ends up pointed at the wrong database.
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config is the whole of the Phase 1 configuration surface.
|
||||
type Config struct {
|
||||
AppEnv string
|
||||
Log LogConfig
|
||||
HTTP HTTPConfig
|
||||
DB DBConfig
|
||||
Seed SeedConfig
|
||||
}
|
||||
|
||||
// SeedConfig locates the demo fixture. The file is generated from the frontend
|
||||
// repository, so it lives beside the migrations rather than inside the Go
|
||||
// module: regenerating it must not require rebuilding the binary.
|
||||
type SeedConfig struct {
|
||||
FixturePath string
|
||||
}
|
||||
|
||||
type LogConfig struct {
|
||||
Level string
|
||||
}
|
||||
|
||||
type HTTPConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
ReadTimeout time.Duration
|
||||
WriteTimeout time.Duration
|
||||
IdleTimeout time.Duration
|
||||
ShutdownTimeout time.Duration
|
||||
|
||||
// CORSOrigins is the exact set of browser origins allowed to call the API.
|
||||
//
|
||||
// It exists for one reason: in local development the Vite dev server is an
|
||||
// origin of its own (http://localhost:5173) and the API is another
|
||||
// (http://127.0.0.1:8080), so every fetch from the frontend is
|
||||
// cross-origin. Empty means CORS is off and the API answers only
|
||||
// same-origin callers, which is the correct posture everywhere the
|
||||
// frontend is served from the same host as the API.
|
||||
//
|
||||
// Origins are matched exactly and echoed back one at a time. There is no
|
||||
// wildcard and no pattern: "*" would let any page on the internet read
|
||||
// this API, and once authentication exists that becomes a real hole rather
|
||||
// than a theoretical one.
|
||||
CORSOrigins []string
|
||||
}
|
||||
|
||||
type DBConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Name string
|
||||
User string
|
||||
Password string
|
||||
Schema string
|
||||
SSLMode string
|
||||
MaxOpenConns int32
|
||||
MinIdleConns int32
|
||||
ConnMaxLifetime time.Duration
|
||||
ConnectTimeout time.Duration
|
||||
StatementTimeout time.Duration
|
||||
}
|
||||
|
||||
// DSN builds a libpq-style connection URL.
|
||||
//
|
||||
// Every component is URL-escaped: the local database is called "Krow-force",
|
||||
// which is both mixed-case and hyphenated, and a password may contain anything
|
||||
// at all. Escaping is not optional here.
|
||||
func (d DBConfig) DSN() string {
|
||||
u := &url.URL{
|
||||
Scheme: "postgres",
|
||||
User: url.UserPassword(d.User, d.Password),
|
||||
Host: fmt.Sprintf("%s:%d", d.Host, d.Port),
|
||||
Path: "/" + d.Name,
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("sslmode", d.SSLMode)
|
||||
// Pin the schema on every connection so no query can accidentally resolve
|
||||
// against a different one, and so nothing reaches for a system schema.
|
||||
q.Set("search_path", d.Schema)
|
||||
q.Set("connect_timeout", strconv.Itoa(int(d.ConnectTimeout.Seconds())))
|
||||
q.Set("statement_timeout", strconv.Itoa(int(d.StatementTimeout.Milliseconds())))
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// Redacted returns the DSN with the password replaced, for logs.
|
||||
func (d DBConfig) Redacted() string {
|
||||
u, err := url.Parse(d.DSN())
|
||||
if err != nil {
|
||||
return "postgres://<unparseable>"
|
||||
}
|
||||
if _, hasPassword := u.User.Password(); hasPassword {
|
||||
u.User = url.UserPassword(u.User.Username(), "xxxxx")
|
||||
}
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// Load reads the environment, applies defaults and validates the result.
|
||||
func Load() (*Config, error) {
|
||||
loadDotEnv(".env")
|
||||
|
||||
var missing []string
|
||||
required := func(key string) string {
|
||||
v := strings.TrimSpace(os.Getenv(key))
|
||||
if v == "" {
|
||||
missing = append(missing, key)
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
cfg := &Config{
|
||||
AppEnv: withDefault("APP_ENV", "development"),
|
||||
Log: LogConfig{Level: withDefault("LOG_LEVEL", "info")},
|
||||
HTTP: HTTPConfig{
|
||||
Host: withDefault("HTTP_HOST", "127.0.0.1"),
|
||||
Port: intDefault("HTTP_PORT", 8080),
|
||||
ReadTimeout: durationDefault("HTTP_READ_TIMEOUT", 15*time.Second),
|
||||
WriteTimeout: durationDefault("HTTP_WRITE_TIMEOUT", 30*time.Second),
|
||||
IdleTimeout: durationDefault("HTTP_IDLE_TIMEOUT", 60*time.Second),
|
||||
ShutdownTimeout: durationDefault("HTTP_SHUTDOWN_TIMEOUT", 10*time.Second),
|
||||
CORSOrigins: corsOrigins(withDefault("APP_ENV", "development")),
|
||||
},
|
||||
Seed: SeedConfig{
|
||||
FixturePath: withDefault("SEED_FIXTURE_PATH", "./seed/fixtures/seed.json"),
|
||||
},
|
||||
DB: DBConfig{
|
||||
Host: required("DATABASE_HOST"),
|
||||
Port: intDefault("DATABASE_PORT", 5432),
|
||||
Name: required("DATABASE_NAME"),
|
||||
User: required("DATABASE_USER"),
|
||||
Password: os.Getenv("DATABASE_PASSWORD"), // may legitimately be empty (trust/peer auth)
|
||||
Schema: withDefault("DATABASE_SCHEMA", "public"),
|
||||
SSLMode: withDefault("DATABASE_SSLMODE", "disable"),
|
||||
MaxOpenConns: int32(intDefault("DATABASE_MAX_OPEN_CONNS", 25)),
|
||||
MinIdleConns: int32(intDefault("DATABASE_MIN_IDLE_CONNS", 2)),
|
||||
ConnMaxLifetime: durationDefault("DATABASE_CONN_MAX_LIFETIME", 30*time.Minute),
|
||||
ConnectTimeout: durationDefault("DATABASE_CONNECT_TIMEOUT", 5*time.Second),
|
||||
StatementTimeout: durationDefault("DATABASE_STATEMENT_TIMEOUT", 10*time.Second),
|
||||
},
|
||||
}
|
||||
|
||||
if len(missing) > 0 {
|
||||
return nil, fmt.Errorf("missing required environment variables: %s "+
|
||||
"(copy .env.example to .env and fill them in)", strings.Join(missing, ", "))
|
||||
}
|
||||
if err := cfg.validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func (c *Config) validate() error {
|
||||
switch c.AppEnv {
|
||||
case "development", "staging", "production":
|
||||
default:
|
||||
return fmt.Errorf("APP_ENV must be development, staging or production, got %q", c.AppEnv)
|
||||
}
|
||||
if c.HTTP.Port < 1 || c.HTTP.Port > 65535 {
|
||||
return fmt.Errorf("HTTP_PORT out of range: %d", c.HTTP.Port)
|
||||
}
|
||||
if c.DB.Port < 1 || c.DB.Port > 65535 {
|
||||
return fmt.Errorf("DATABASE_PORT out of range: %d", c.DB.Port)
|
||||
}
|
||||
// The application owns exactly one schema and it is never a system schema.
|
||||
switch c.DB.Schema {
|
||||
case "pg_catalog", "pg_toast", "information_schema":
|
||||
return fmt.Errorf("DATABASE_SCHEMA must not be a PostgreSQL system schema, got %q", c.DB.Schema)
|
||||
}
|
||||
if strings.HasPrefix(c.DB.Schema, "pg_") {
|
||||
return fmt.Errorf("DATABASE_SCHEMA must not start with \"pg_\", got %q", c.DB.Schema)
|
||||
}
|
||||
if c.DB.MinIdleConns > c.DB.MaxOpenConns {
|
||||
return fmt.Errorf("DATABASE_MIN_IDLE_CONNS (%d) exceeds DATABASE_MAX_OPEN_CONNS (%d)",
|
||||
c.DB.MinIdleConns, c.DB.MaxOpenConns)
|
||||
}
|
||||
if c.AppEnv == "production" && c.DB.SSLMode == "disable" {
|
||||
return fmt.Errorf("DATABASE_SSLMODE=disable is not allowed when APP_ENV=production")
|
||||
}
|
||||
for _, origin := range c.HTTP.CORSOrigins {
|
||||
// "*" is rejected rather than quietly honoured. The middleware echoes a
|
||||
// single matched origin, so a wildcard could only ever be a
|
||||
// misunderstanding of what this setting does.
|
||||
if origin == "*" {
|
||||
return fmt.Errorf("HTTP_CORS_ORIGINS must list explicit origins; \"*\" is not accepted")
|
||||
}
|
||||
if !strings.HasPrefix(origin, "http://") && !strings.HasPrefix(origin, "https://") {
|
||||
return fmt.Errorf("HTTP_CORS_ORIGINS entry %q must be a full origin including the scheme", origin)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// devCORSOrigins are the origins the Vite dev server can occupy. Vite binds
|
||||
// localhost by default and 127.0.0.1 when asked, and a browser treats those two
|
||||
// as different origins, so both are listed. 4173 is `vite preview`.
|
||||
var devCORSOrigins = []string{
|
||||
"http://localhost:5173", "http://127.0.0.1:5173",
|
||||
"http://localhost:4173", "http://127.0.0.1:4173",
|
||||
}
|
||||
|
||||
// corsOrigins reads HTTP_CORS_ORIGINS, a comma-separated allowlist.
|
||||
//
|
||||
// The development default is the Vite dev server, because that is the whole
|
||||
// point of the setting in Phase 2D. Outside development the default is empty:
|
||||
// a staging or production deployment that genuinely serves its frontend from
|
||||
// another origin has to say so explicitly, rather than inheriting a list of
|
||||
// localhost origins nobody reviewed.
|
||||
func corsOrigins(appEnv string) []string {
|
||||
raw, set := os.LookupEnv("HTTP_CORS_ORIGINS")
|
||||
if !set {
|
||||
if appEnv == "development" {
|
||||
return devCORSOrigins
|
||||
}
|
||||
return nil
|
||||
}
|
||||
var out []string
|
||||
for _, part := range strings.Split(raw, ",") {
|
||||
// A trailing slash makes the string unequal to the Origin header the
|
||||
// browser actually sends, which fails in a way that looks like a
|
||||
// server bug rather than a typo.
|
||||
if o := strings.TrimRight(strings.TrimSpace(part), "/"); o != "" {
|
||||
out = append(out, o)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func withDefault(key, fallback string) string {
|
||||
if v := strings.TrimSpace(os.Getenv(key)); v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func intDefault(key string, fallback int) int {
|
||||
v := strings.TrimSpace(os.Getenv(key))
|
||||
if v == "" {
|
||||
return fallback
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func durationDefault(key string, fallback time.Duration) time.Duration {
|
||||
v := strings.TrimSpace(os.Getenv(key))
|
||||
if v == "" {
|
||||
return fallback
|
||||
}
|
||||
d, err := time.ParseDuration(v)
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// loadDotEnv reads KEY=VALUE lines, walking up from the working directory so
|
||||
// `go run ./cmd/api` finds the repository-root .env. Existing environment
|
||||
// variables always win. Absence of the file is not an error.
|
||||
func loadDotEnv(name string) {
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
for i := 0; i < 5; i++ {
|
||||
path := dir + string(os.PathSeparator) + name
|
||||
if data, err := os.ReadFile(path); err == nil {
|
||||
applyDotEnv(string(data))
|
||||
return
|
||||
}
|
||||
parent := parentDir(dir)
|
||||
if parent == dir {
|
||||
return
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
}
|
||||
|
||||
func parentDir(dir string) string {
|
||||
i := strings.LastIndex(dir, string(os.PathSeparator))
|
||||
if i <= 0 {
|
||||
return dir
|
||||
}
|
||||
return dir[:i]
|
||||
}
|
||||
|
||||
func applyDotEnv(content string) {
|
||||
for _, line := range strings.Split(content, "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
key, value, ok := strings.Cut(line, "=")
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
key = strings.TrimSpace(strings.TrimPrefix(key, "export "))
|
||||
value = strings.TrimSpace(value)
|
||||
if len(value) >= 2 {
|
||||
if (value[0] == '"' && value[len(value)-1] == '"') ||
|
||||
(value[0] == '\'' && value[len(value)-1] == '\'') {
|
||||
value = value[1 : len(value)-1]
|
||||
}
|
||||
}
|
||||
if _, present := os.LookupEnv(key); !present {
|
||||
_ = os.Setenv(key, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
122
go-api/internal/db/db.go
Normal file
122
go-api/internal/db/db.go
Normal file
@@ -0,0 +1,122 @@
|
||||
// Package db owns the PostgreSQL connection pool and the database health check.
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
)
|
||||
|
||||
// DB wraps the pgx pool together with the schema the application is pinned to.
|
||||
type DB struct {
|
||||
Pool *pgxpool.Pool
|
||||
Schema string
|
||||
}
|
||||
|
||||
// Open builds the pool and verifies it can actually reach the database.
|
||||
//
|
||||
// pgxpool.New is lazy — it returns a usable pool without having connected —
|
||||
// so a bad host or a wrong password would otherwise not surface until the
|
||||
// first request. Acquiring and pinging once here turns a misconfiguration into
|
||||
// a startup failure instead of a runtime surprise.
|
||||
func Open(ctx context.Context, cfg config.DBConfig) (*DB, error) {
|
||||
poolCfg, err := pgxpool.ParseConfig(cfg.DSN())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse database config: %w", err)
|
||||
}
|
||||
poolCfg.MaxConns = cfg.MaxOpenConns
|
||||
poolCfg.MinIdleConns = cfg.MinIdleConns
|
||||
poolCfg.MaxConnLifetime = cfg.ConnMaxLifetime
|
||||
|
||||
pool, err := pgxpool.NewWithConfig(ctx, poolCfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create connection pool: %w", err)
|
||||
}
|
||||
|
||||
pingCtx, cancel := context.WithTimeout(ctx, cfg.ConnectTimeout)
|
||||
defer cancel()
|
||||
if err := pool.Ping(pingCtx); err != nil {
|
||||
pool.Close()
|
||||
return nil, fmt.Errorf("connect to %s: %w", cfg.Redacted(), err)
|
||||
}
|
||||
|
||||
return &DB{Pool: pool, Schema: cfg.Schema}, nil
|
||||
}
|
||||
|
||||
// Close releases every pooled connection.
|
||||
func (d *DB) Close() {
|
||||
if d != nil && d.Pool != nil {
|
||||
d.Pool.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// Health is what the /health endpoint reports about the database.
|
||||
type Health struct {
|
||||
Reachable bool `json:"reachable"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Version string `json:"version,omitempty"`
|
||||
Database string `json:"database,omitempty"`
|
||||
Schema string `json:"schema,omitempty"`
|
||||
SchemaPresent bool `json:"schema_present"`
|
||||
AppliedMigration *int64 `json:"applied_migration,omitempty"`
|
||||
MigrationDirty bool `json:"migration_dirty"`
|
||||
TableCount int `json:"table_count"`
|
||||
LatencyMS int64 `json:"latency_ms"`
|
||||
}
|
||||
|
||||
// Check answers "can the API serve requests against this database right now".
|
||||
//
|
||||
// It reports more than a ping because a reachable database with no schema in it
|
||||
// is a different failure from an unreachable one, and both are worth telling
|
||||
// apart at a glance during Phase 1. Reads are confined to the configured schema
|
||||
// via to_regclass and a count over information_schema, which is the standard
|
||||
// SQL view rather than a pg_catalog table.
|
||||
func (d *DB) Check(ctx context.Context) Health {
|
||||
started := time.Now()
|
||||
h := Health{Schema: d.Schema}
|
||||
|
||||
conn, err := d.Pool.Acquire(ctx)
|
||||
if err != nil {
|
||||
h.Error = err.Error()
|
||||
h.LatencyMS = time.Since(started).Milliseconds()
|
||||
return h
|
||||
}
|
||||
defer conn.Release()
|
||||
|
||||
if err := conn.QueryRow(ctx,
|
||||
`SELECT current_database(), current_setting('server_version')`,
|
||||
).Scan(&h.Database, &h.Version); err != nil {
|
||||
h.Error = err.Error()
|
||||
h.LatencyMS = time.Since(started).Milliseconds()
|
||||
return h
|
||||
}
|
||||
h.Reachable = true
|
||||
|
||||
if err := conn.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM information_schema.tables
|
||||
WHERE table_schema = $1 AND table_type = 'BASE TABLE'`,
|
||||
d.Schema,
|
||||
).Scan(&h.TableCount); err != nil {
|
||||
h.Error = err.Error()
|
||||
h.LatencyMS = time.Since(started).Milliseconds()
|
||||
return h
|
||||
}
|
||||
h.SchemaPresent = h.TableCount > 0
|
||||
|
||||
// golang-migrate's bookkeeping table. Absent before the first migration,
|
||||
// which is a legitimate state and not an error.
|
||||
var version int64
|
||||
var dirty bool
|
||||
err = conn.QueryRow(ctx, `SELECT version, dirty FROM schema_migrations LIMIT 1`).Scan(&version, &dirty)
|
||||
if err == nil {
|
||||
h.AppliedMigration = &version
|
||||
h.MigrationDirty = dirty
|
||||
}
|
||||
|
||||
h.LatencyMS = time.Since(started).Milliseconds()
|
||||
return h
|
||||
}
|
||||
543
go-api/internal/definition/agent.go
Normal file
543
go-api/internal/definition/agent.go
Normal file
@@ -0,0 +1,543 @@
|
||||
package definition
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// One Markdown definition → one agent.
|
||||
//
|
||||
// A port of parseAgent / validateAgentSource in src/lib/agents/registry.js and
|
||||
// normalizeAgent in src/lib/agents/agentConfig.js.
|
||||
//
|
||||
// The contract normalizeAgent holds, and this holds with it:
|
||||
//
|
||||
// - Everything is optional. A definition declaring only an id and a name
|
||||
// normalizes to a working agent with documented defaults.
|
||||
// - Nothing unknown survives. Statuses, reasoning modes, pages, icons,
|
||||
// knowledge kinds and permission roles are checked against the closed
|
||||
// tables in vocabulary.go; an unrecognised value is a named error rather
|
||||
// than a dropped key.
|
||||
// - What validates is kept. One bad entry costs its author that entry and a
|
||||
// message, never the rest of the file.
|
||||
//
|
||||
// One rule is deliberately absent, and it is absent on the frontend for the
|
||||
// same reason: an agent with NO SKILLS is not refused. Five of this product's
|
||||
// pages have no Owliver skills and answer from their own page responder, so
|
||||
// refusing a skill-less agent would mean inventing placeholder skills to make
|
||||
// those pages configurable.
|
||||
|
||||
// Starter is one conversation starter.
|
||||
type Starter struct {
|
||||
Label string `json:"label"`
|
||||
Prompt string `json:"prompt"`
|
||||
}
|
||||
|
||||
// Knowledge is one thing an agent has been told, as distinct from something it
|
||||
// can do. Modelled as a document with an id and a body because that is the
|
||||
// shape a retrieval layer reads.
|
||||
type Knowledge struct {
|
||||
ID string `json:"id"`
|
||||
Label string `json:"label"`
|
||||
Kind string `json:"kind"`
|
||||
Body string `json:"body"`
|
||||
URL string `json:"url"`
|
||||
}
|
||||
|
||||
// Person is one named grant on an agent.
|
||||
type Person struct {
|
||||
User string `json:"user"`
|
||||
Role string `json:"role"`
|
||||
}
|
||||
|
||||
// Permissions is who owns an agent, who may reach it, and what they may do.
|
||||
//
|
||||
// Parsed and NOT enforced. Migration 000005 deliberately has no
|
||||
// definition_permissions table: the block stays inside the Markdown until its
|
||||
// semantics are defined.
|
||||
type Permissions struct {
|
||||
Owner string `json:"owner"`
|
||||
Access string `json:"access"`
|
||||
People []Person `json:"people"`
|
||||
}
|
||||
|
||||
// Agent is a definition as the backend reads it.
|
||||
type Agent struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Version int `json:"version"`
|
||||
|
||||
// Pages as CANONICAL surface ids.
|
||||
//
|
||||
// Unlike Skill.Pages, which keeps what the author wrote. The two are
|
||||
// genuinely different on the frontend — normalizeAgent maps every page
|
||||
// through canonicalPage and parseSkill does not — so agent_definitions.pages
|
||||
// and skill_definitions.pages hold different vocabularies for the same
|
||||
// concept. Reproduced rather than reconciled: making them agree here would
|
||||
// make each one disagree with its own editor.
|
||||
Pages []string `json:"pages"`
|
||||
|
||||
Icon string `json:"icon"`
|
||||
Reasoning string `json:"reasoning"`
|
||||
Trigger string `json:"trigger"`
|
||||
WebSearch bool `json:"webSearch"`
|
||||
|
||||
Skills []string `json:"skills"`
|
||||
Subagents []string `json:"subagents"`
|
||||
Starters []Starter `json:"starters"`
|
||||
Knowledge []Knowledge `json:"knowledge"`
|
||||
Permissions Permissions `json:"permissions"`
|
||||
|
||||
// Instructions is the body's `## Instructions` section. Prose belongs under
|
||||
// a heading where it can be written and read as prose, not in a
|
||||
// frontmatter string.
|
||||
Instructions string `json:"instructions"`
|
||||
|
||||
// Errors is what this definition lost on the way in, in the order
|
||||
// normalizeAgent produces them. Carried on the record rather than thrown,
|
||||
// so one bad entry costs its author that entry and a message.
|
||||
Errors []string `json:"errors"`
|
||||
|
||||
Body string `json:"-"`
|
||||
}
|
||||
|
||||
// asList is agentConfig.js's own coercion: an array stays an array, nothing
|
||||
// becomes nothing, and anything else becomes a list of one.
|
||||
//
|
||||
// This is why `pages: candidates` is accepted for an AGENT and refused for a
|
||||
// SKILL — parseSkill requires a real sequence and normalizeAgent coerces.
|
||||
func asList(v any) []any {
|
||||
switch x := v.(type) {
|
||||
case []any:
|
||||
return x
|
||||
case nil:
|
||||
return []any{}
|
||||
case string:
|
||||
if x == "" {
|
||||
return []any{}
|
||||
}
|
||||
}
|
||||
return []any{v}
|
||||
}
|
||||
|
||||
// uniqueStrings keeps order and drops repeats; a blank entry is an error rather
|
||||
// than a silent gap, because a blank id is an address that points nowhere.
|
||||
func uniqueStrings(raw any, where, label string, errs *[]string) []string {
|
||||
seen := map[string]bool{}
|
||||
out := []string{}
|
||||
for i, entry := range asList(raw) {
|
||||
value := jsTrimmed(entry)
|
||||
if value == "" {
|
||||
*errs = append(*errs, fmt.Sprintf("%s[%d]: %s cannot be blank.", where, i, label))
|
||||
continue
|
||||
}
|
||||
if seen[value] {
|
||||
continue
|
||||
}
|
||||
seen[value] = true
|
||||
out = append(out, value)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// normalizePages resolves every declared page to a canonical surface key.
|
||||
//
|
||||
// Through CanonicalPage, so a definition may write an alias — `university` for
|
||||
// `krow-forge` — exactly as a skill may. An unknown page is an error rather
|
||||
// than a silently dropped entry, because a page nobody recognises is an agent
|
||||
// that will never appear anywhere and give no reason why.
|
||||
func normalizePages(raw any, errs *[]string) []string {
|
||||
seen := map[string]bool{}
|
||||
pages := []string{}
|
||||
for i, entry := range asList(raw) {
|
||||
written := jsTrimmed(entry)
|
||||
if written == "" {
|
||||
*errs = append(*errs, fmt.Sprintf("pages[%d]: a page cannot be blank.", i))
|
||||
continue
|
||||
}
|
||||
canonical := CanonicalPage(written)
|
||||
if canonical == "" {
|
||||
*errs = append(*errs, fmt.Sprintf(
|
||||
"pages[%d]: `%s` is not a page this product has.", i, written))
|
||||
continue
|
||||
}
|
||||
if seen[canonical] {
|
||||
continue
|
||||
}
|
||||
seen[canonical] = true
|
||||
pages = append(pages, canonical)
|
||||
}
|
||||
return pages
|
||||
}
|
||||
|
||||
// normalizeStarter reads one starter, in either the plain-string or the mapping
|
||||
// form. A starter with no prompt of its own asks what it says.
|
||||
func normalizeStarter(raw any, index int, errs *[]string) *Starter {
|
||||
where := fmt.Sprintf("starters[%d]", index)
|
||||
|
||||
switch v := raw.(type) {
|
||||
case string, float64:
|
||||
label := jsTrimmed(v)
|
||||
if label == "" {
|
||||
*errs = append(*errs, where+": a starter needs text.")
|
||||
return nil
|
||||
}
|
||||
return &Starter{Label: label, Prompt: label}
|
||||
case map[string]any:
|
||||
// `raw.label ?? raw.prompt` — nullish, so an absent or null label
|
||||
// falls through to the prompt and a starter written as a bare prompt
|
||||
// still has something to show.
|
||||
source := v["label"]
|
||||
if source == nil {
|
||||
source = v["prompt"]
|
||||
}
|
||||
label := jsTrimmed(source)
|
||||
if label == "" {
|
||||
*errs = append(*errs, where+": a starter needs a `label`.")
|
||||
return nil
|
||||
}
|
||||
prompt := jsTrimmed(v["prompt"])
|
||||
if prompt == "" {
|
||||
prompt = label
|
||||
}
|
||||
return &Starter{Label: label, Prompt: prompt}
|
||||
}
|
||||
|
||||
*errs = append(*errs, where+": a starter must be a line of text, or a mapping of options.")
|
||||
return nil
|
||||
}
|
||||
|
||||
// normalizeKnowledge reads one knowledge entry.
|
||||
func normalizeKnowledge(raw any, index int, errs *[]string) *Knowledge {
|
||||
where := fmt.Sprintf("knowledge[%d]", index)
|
||||
|
||||
switch v := raw.(type) {
|
||||
case string, float64:
|
||||
body := jsTrimmed(v)
|
||||
if body == "" {
|
||||
*errs = append(*errs, where+": a knowledge entry needs text.")
|
||||
return nil
|
||||
}
|
||||
id := slugify(runeSlice(body, 40))
|
||||
if id == "" {
|
||||
id = fmt.Sprintf("k%d", index+1)
|
||||
}
|
||||
return &Knowledge{
|
||||
ID: id, Label: runeSlice(body, 60), Kind: DefaultKnowledgeKind, Body: body,
|
||||
}
|
||||
case map[string]any:
|
||||
label := jsTrimmed(v["label"])
|
||||
body := jsTrimmed(v["body"])
|
||||
url := jsTrimmed(v["url"])
|
||||
|
||||
if label == "" && body == "" {
|
||||
*errs = append(*errs, where+": a knowledge entry needs a `label` or a `body`.")
|
||||
return nil
|
||||
}
|
||||
|
||||
kind := jsTrimmed(v["kind"])
|
||||
if kind == "" {
|
||||
kind = DefaultKnowledgeKind
|
||||
}
|
||||
if !contains(KnowledgeKinds, kind) {
|
||||
*errs = append(*errs, fmt.Sprintf(
|
||||
"%s: `%s` is not a knowledge kind. Use one of %s.",
|
||||
where, kind, strings.Join(KnowledgeKinds, ", ")))
|
||||
return nil
|
||||
}
|
||||
if kind == "link" && url == "" {
|
||||
*errs = append(*errs, where+": a `link` needs a `url`.")
|
||||
return nil
|
||||
}
|
||||
|
||||
id := jsTrimmed(v["id"])
|
||||
if id == "" {
|
||||
id = slugify(label)
|
||||
}
|
||||
if id == "" {
|
||||
id = fmt.Sprintf("k%d", index+1)
|
||||
}
|
||||
if label == "" {
|
||||
label = runeSlice(body, 60)
|
||||
}
|
||||
return &Knowledge{ID: id, Label: label, Kind: kind, Body: body, URL: url}
|
||||
}
|
||||
|
||||
*errs = append(*errs, where+": a knowledge entry must be a line of text, or a mapping of options.")
|
||||
return nil
|
||||
}
|
||||
|
||||
// runeSlice is JavaScript's String.prototype.slice(0, n), which counts UTF-16
|
||||
// units. Counting runes instead differs only for astral characters, and cutting
|
||||
// a surrogate pair in half — which the frontend can do — would produce a label
|
||||
// no comparison could match. Runes are used deliberately; the conformance suite
|
||||
// carries no case that distinguishes them.
|
||||
func runeSlice(s string, n int) string {
|
||||
r := []rune(s)
|
||||
if len(r) <= n {
|
||||
return s
|
||||
}
|
||||
return string(r[:n])
|
||||
}
|
||||
|
||||
// normalizePermissions reads the `permissions:` block.
|
||||
func normalizePermissions(raw any, errs *[]string) Permissions {
|
||||
none := Permissions{Access: DefaultAgentAccess, People: []Person{}}
|
||||
if raw == nil {
|
||||
return none
|
||||
}
|
||||
|
||||
mapping, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
*errs = append(*errs, "permissions: must be a mapping of `owner`, `access` and `people`.")
|
||||
return none
|
||||
}
|
||||
|
||||
access := jsTrimmed(mapping["access"])
|
||||
if access == "" {
|
||||
access = DefaultAgentAccess
|
||||
}
|
||||
if !contains(AgentAccess, access) {
|
||||
*errs = append(*errs, fmt.Sprintf(
|
||||
"permissions.access: `%s` is not an access mode. Use one of %s.",
|
||||
access, strings.Join(AgentAccess, ", ")))
|
||||
}
|
||||
|
||||
people := []Person{}
|
||||
for i, entry := range asList(mapping["people"]) {
|
||||
where := fmt.Sprintf("permissions.people[%d]", i)
|
||||
person, ok := entry.(map[string]any)
|
||||
if !ok {
|
||||
*errs = append(*errs, where+": must be a mapping of `user` and `role`.")
|
||||
continue
|
||||
}
|
||||
user := jsTrimmed(person["user"])
|
||||
if user == "" {
|
||||
*errs = append(*errs, where+": needs a `user`.")
|
||||
continue
|
||||
}
|
||||
role := jsTrimmed(person["role"])
|
||||
if role == "" {
|
||||
role = DefaultPermission
|
||||
}
|
||||
if !contains(PermissionRole, role) {
|
||||
*errs = append(*errs, fmt.Sprintf(
|
||||
"%s: `%s` is not a role. Use one of %s.",
|
||||
where, role, strings.Join(PermissionRole, ", ")))
|
||||
continue
|
||||
}
|
||||
people = append(people, Person{User: user, Role: role})
|
||||
}
|
||||
|
||||
result := Permissions{Owner: jsTrimmed(mapping["owner"]), Access: access, People: people}
|
||||
if !contains(AgentAccess, access) {
|
||||
result.Access = DefaultAgentAccess
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// ParseAgent reads an agent definition.
|
||||
//
|
||||
// The order in which errors accumulate is part of the contract: validateAgent
|
||||
// reports the FIRST one, so a definition with two problems must name the same
|
||||
// one the editor names. That order is status, reasoning, icon, version,
|
||||
// subagents, starters, knowledge, pages, skills, permissions — which is
|
||||
// evaluation order in normalizeAgent, counting the object literal it returns.
|
||||
func ParseAgent(raw string, opts Options) (*Agent, error) {
|
||||
doc, err := ParseFrontmatter(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data := doc.Data
|
||||
errs := []string{}
|
||||
|
||||
id := jsTrim(jsString(data["id"]))
|
||||
if !jsTruthy(data["id"]) {
|
||||
id = slugify(data["name"])
|
||||
if id == "" {
|
||||
id = fileStem(opts.path())
|
||||
}
|
||||
}
|
||||
|
||||
status := jsTrimmed(data["status"])
|
||||
if status == "" {
|
||||
status = DefaultAgentStatus
|
||||
}
|
||||
if !contains(AgentStatuses, status) {
|
||||
errs = append(errs, fmt.Sprintf("status: `%s` is not a status. Use one of %s.",
|
||||
status, strings.Join(AgentStatuses, ", ")))
|
||||
}
|
||||
|
||||
reasoning := jsTrimmed(data["reasoning"])
|
||||
if reasoning == "" {
|
||||
reasoning = DefaultReasoning
|
||||
}
|
||||
if !contains(ReasoningModes, reasoning) {
|
||||
errs = append(errs, fmt.Sprintf("reasoning: `%s` is not a reasoning mode. Use one of %s.",
|
||||
reasoning, strings.Join(ReasoningModes, ", ")))
|
||||
}
|
||||
|
||||
icon := jsTrimmed(data["icon"])
|
||||
if icon == "" {
|
||||
icon = DefaultAgentIcon
|
||||
}
|
||||
if !contains(AgentIcons, icon) {
|
||||
errs = append(errs, fmt.Sprintf("icon: `%s` is not an icon this product has.", icon))
|
||||
}
|
||||
|
||||
// A version is an integer that only ever goes up. Anything else is an
|
||||
// authoring slip, and reading it as 1 is kinder than refusing the file —
|
||||
// but it is still reported, because a definition that thinks it is v3 and
|
||||
// registers as v1 will publish over something.
|
||||
version := 1
|
||||
if v, present := data["version"]; present && v != nil && v != "" {
|
||||
parsed := jsNumber(v)
|
||||
if math.IsNaN(parsed) || parsed != math.Trunc(parsed) || math.IsInf(parsed, 0) || parsed < 1 {
|
||||
errs = append(errs, fmt.Sprintf(
|
||||
"version: `%s` is not a whole number of 1 or more.", jsString(v)))
|
||||
} else if parsed > maxExactInteger {
|
||||
// Beyond 2^53-1 a float64 no longer names one integer, so there is
|
||||
// no value to carry. Saturating keeps the conversion below defined,
|
||||
// and ValidateAgent refuses everything above MaxVersion anyway, so
|
||||
// a saturated version can never reach a column.
|
||||
version = maxExactInteger
|
||||
} else {
|
||||
version = int(parsed)
|
||||
}
|
||||
}
|
||||
|
||||
subagents := uniqueStrings(data["subagents"], "subagents", "a subagent id", &errs)
|
||||
kept := subagents[:0]
|
||||
for _, s := range subagents {
|
||||
if id != "" && s == id {
|
||||
errs = append(errs, "subagents: an agent cannot be its own subagent.")
|
||||
continue
|
||||
}
|
||||
kept = append(kept, s)
|
||||
}
|
||||
subagents = kept
|
||||
|
||||
starters := []Starter{}
|
||||
for i, entry := range asList(data["starters"]) {
|
||||
if s := normalizeStarter(entry, i, &errs); s != nil {
|
||||
starters = append(starters, *s)
|
||||
}
|
||||
}
|
||||
|
||||
knowledge := []Knowledge{}
|
||||
for i, entry := range asList(data["knowledge"]) {
|
||||
if k := normalizeKnowledge(entry, i, &errs); k != nil {
|
||||
knowledge = append(knowledge, *k)
|
||||
}
|
||||
}
|
||||
|
||||
// From here the order follows the object literal normalizeAgent returns.
|
||||
pages := normalizePages(data["pages"], &errs)
|
||||
skills := uniqueStrings(data["skills"], "skills", "a skill id", &errs)
|
||||
permissions := normalizePermissions(data["permissions"], &errs)
|
||||
|
||||
instructions, _ := sectionSource(doc.Body, "Instructions")
|
||||
|
||||
agent := &Agent{
|
||||
ID: id,
|
||||
Name: "Untitled agent",
|
||||
Status: status,
|
||||
Version: version,
|
||||
Pages: pages,
|
||||
Icon: icon,
|
||||
Reasoning: reasoning,
|
||||
Trigger: jsTrimmed(data["trigger"]),
|
||||
WebSearch: data["webSearch"] == true || data["web_search"] == true,
|
||||
Skills: skills,
|
||||
Subagents: subagents,
|
||||
Starters: starters,
|
||||
Knowledge: knowledge,
|
||||
Permissions: permissions,
|
||||
Instructions: jsTrim(instructions),
|
||||
Errors: errs,
|
||||
Body: doc.Body,
|
||||
}
|
||||
|
||||
if jsTruthy(data["name"]) {
|
||||
agent.Name = jsString(data["name"])
|
||||
}
|
||||
if jsTruthy(data["description"]) {
|
||||
agent.Description = jsString(data["description"])
|
||||
}
|
||||
if !contains(AgentStatuses, status) {
|
||||
agent.Status = DefaultAgentStatus
|
||||
}
|
||||
if !contains(ReasoningModes, reasoning) {
|
||||
agent.Reasoning = DefaultReasoning
|
||||
}
|
||||
if !contains(AgentIcons, icon) {
|
||||
agent.Icon = DefaultAgentIcon
|
||||
}
|
||||
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// ValidateAgent decides whether an agent definition may be stored.
|
||||
//
|
||||
// Returns nil when it may. The order is the order an author would fix things
|
||||
// in, which is why it reads the same way ValidateSkill does. Note that the
|
||||
// `id` message differs from the skill one by two words — that difference is
|
||||
// the frontend's, and it is reproduced rather than tidied.
|
||||
func ValidateAgent(raw string) error {
|
||||
if jsTrim(raw) == "" {
|
||||
return &Rejection{Message: "Paste or upload a Markdown definition."}
|
||||
}
|
||||
|
||||
if n := len([]rune(raw)); n > MaxMarkdownLength {
|
||||
return &Rejection{
|
||||
BackendOnly: true,
|
||||
Message: fmt.Sprintf(
|
||||
"That definition is %d characters. The limit is %d.", n, MaxMarkdownLength),
|
||||
}
|
||||
}
|
||||
|
||||
agent, err := ParseAgent(raw, Options{})
|
||||
if err != nil {
|
||||
return &Rejection{Message: jsTrim("That definition could not be parsed. " + err.Error())}
|
||||
}
|
||||
|
||||
if agent.ID == "" {
|
||||
return &Rejection{Message: "The frontmatter needs an `id`."}
|
||||
}
|
||||
if !isDefinitionID(agent.ID) {
|
||||
return &Rejection{Message: "The `id` must be lower-case letters, numbers and dashes."}
|
||||
}
|
||||
// `!raw.includes('name:') || agent.name === 'Untitled agent'` — the literal
|
||||
// substring test is the frontend's, and it is why a definition whose name
|
||||
// resolves to the fallback is refused even when some other key happens to
|
||||
// spell `name:`.
|
||||
if !strings.Contains(raw, "name:") || agent.Name == "Untitled agent" {
|
||||
return &Rejection{Message: "The frontmatter needs a `name`."}
|
||||
}
|
||||
if len(agent.Pages) == 0 {
|
||||
return &Rejection{Message: "An agent needs at least one `pages:` entry, or it can never be offered anywhere."}
|
||||
}
|
||||
if len(agent.Errors) > 0 {
|
||||
return &Rejection{Message: agent.Errors[0]}
|
||||
}
|
||||
|
||||
// The backend's own bound, the companion to the size rule above:
|
||||
// agent_definitions.version is a PostgreSQL `integer`, and the frontend
|
||||
// accepts any whole number of 1 or more. It is checked here rather than in
|
||||
// ParseAgent so the normalized record stays identical to the frontend's for
|
||||
// every definition the frontend accepts, and last among the rules so a
|
||||
// definition the frontend also refuses is refused with the frontend's own
|
||||
// message.
|
||||
if agent.Version > MaxVersion {
|
||||
return &Rejection{
|
||||
BackendOnly: true,
|
||||
Message: fmt.Sprintf(
|
||||
"version: `%d` is larger than %d.", agent.Version, MaxVersion),
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
928
go-api/internal/definition/conformance_test.go
Normal file
928
go-api/internal/definition/conformance_test.go
Normal file
@@ -0,0 +1,928 @@
|
||||
package definition_test
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/definition"
|
||||
)
|
||||
|
||||
// Phase 4D — JS/Go parser conformance.
|
||||
//
|
||||
// The fixture these tests read (testdata/oracle.json) is not written by hand.
|
||||
// It is captured by scripts/oracle.mjs, which loads the REAL frontend module
|
||||
// graph through Vite — import.meta.glob, the `@/` alias and raw Markdown
|
||||
// loading all behave exactly as they do in the app — and records what the
|
||||
// JavaScript parser did with every shipped definition and every adversarial
|
||||
// case. So the assertion below is not "Go agrees with a description of the
|
||||
// frontend"; it is "Go agrees with the frontend", replayed.
|
||||
//
|
||||
// Regenerate after any change to src/lib/skills or src/lib/agents:
|
||||
//
|
||||
// node scripts/oracle.mjs go-api/internal/definition/testdata/oracle.json
|
||||
//
|
||||
// A frontend change that alters parsing therefore fails these tests, which is
|
||||
// the point: the contract cannot drift silently in either direction.
|
||||
|
||||
type oracle struct {
|
||||
Vocabulary struct {
|
||||
Pages []struct {
|
||||
ID string `json:"id"`
|
||||
Aliases []string `json:"aliases"`
|
||||
} `json:"pages"`
|
||||
AgentStatuses []string `json:"agentStatuses"`
|
||||
Reasoning []string `json:"reasoning"`
|
||||
Icons []string `json:"icons"`
|
||||
KnowledgeKinds []string `json:"knowledgeKinds"`
|
||||
Access []string `json:"access"`
|
||||
Roles []string `json:"roles"`
|
||||
} `json:"vocabulary"`
|
||||
Corpus []observation `json:"corpus"`
|
||||
Cases []observation `json:"cases"`
|
||||
}
|
||||
|
||||
// observation is one definition as the JavaScript saw it, end to end.
|
||||
type observation struct {
|
||||
// Exactly one of these identifies the row.
|
||||
Path string `json:"path"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
|
||||
Kind string `json:"kind"` // agent | skill
|
||||
RawBase64 string `json:"rawBase64"`
|
||||
|
||||
HasFrontmatter bool `json:"hasFrontmatter"`
|
||||
|
||||
Frontmatter struct {
|
||||
OK bool `json:"ok"`
|
||||
Data map[string]any `json:"data"`
|
||||
Body string `json:"body"`
|
||||
Error string `json:"error"`
|
||||
} `json:"frontmatter"`
|
||||
|
||||
Parse struct {
|
||||
OK bool `json:"ok"`
|
||||
Error string `json:"error"`
|
||||
} `json:"parse"`
|
||||
|
||||
Normalized map[string]any `json:"normalized"`
|
||||
|
||||
Accepted bool `json:"accepted"`
|
||||
Rejection *string `json:"rejection"`
|
||||
}
|
||||
|
||||
func (o observation) id() string {
|
||||
if o.Path != "" {
|
||||
return o.Path
|
||||
}
|
||||
return o.Name
|
||||
}
|
||||
|
||||
func (o observation) raw(t *testing.T) string {
|
||||
t.Helper()
|
||||
b, err := base64.StdEncoding.DecodeString(o.RawBase64)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: undecodable fixture: %v", o.id(), err)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func load(t *testing.T) *oracle {
|
||||
t.Helper()
|
||||
b, err := os.ReadFile("testdata/oracle.json")
|
||||
if err != nil {
|
||||
t.Fatalf("read fixture: %v", err)
|
||||
}
|
||||
var o oracle
|
||||
if err := json.Unmarshal(b, &o); err != nil {
|
||||
t.Fatalf("parse fixture: %v", err)
|
||||
}
|
||||
if len(o.Corpus) == 0 || len(o.Cases) == 0 {
|
||||
t.Fatal("fixture is empty; regenerate with scripts/oracle.mjs")
|
||||
}
|
||||
return &o
|
||||
}
|
||||
|
||||
func all(o *oracle) []observation { return append(append([]observation{}, o.Corpus...), o.Cases...) }
|
||||
|
||||
/* ── 1. The corpus is the corpus ──────────────────────────────────────────── */
|
||||
|
||||
// The shipped definition count, asserted rather than assumed. A definition
|
||||
// added to or removed from the product without regenerating the fixture leaves
|
||||
// these tests passing against a corpus that no longer exists, which is the one
|
||||
// way this suite could quietly stop meaning anything.
|
||||
func TestCorpusShape(t *testing.T) {
|
||||
o := load(t)
|
||||
|
||||
counts := map[string]int{}
|
||||
for _, c := range o.Corpus {
|
||||
counts[c.Type]++
|
||||
}
|
||||
|
||||
for _, want := range []struct {
|
||||
kind string
|
||||
n int
|
||||
}{{"agent", 9}, {"skill", 23}, {"example", 5}} {
|
||||
if counts[want.kind] != want.n {
|
||||
t.Errorf("%s definitions: got %d, want %d", want.kind, counts[want.kind], want.n)
|
||||
}
|
||||
}
|
||||
if len(o.Corpus) != 37 {
|
||||
t.Errorf("shipped definitions: got %d, want 37", len(o.Corpus))
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 2. The vocabulary has not drifted ────────────────────────────────────── */
|
||||
|
||||
// Every closed table in vocabulary.go, checked against the table the frontend
|
||||
// actually exports. A page added to surfaces.js fails here rather than becoming
|
||||
// a definition the editor accepts and the API rejects.
|
||||
func TestVocabularyMatchesFrontend(t *testing.T) {
|
||||
o := load(t)
|
||||
|
||||
wantPages := make([]string, len(o.Vocabulary.Pages))
|
||||
for i, p := range o.Vocabulary.Pages {
|
||||
wantPages[i] = p.ID
|
||||
}
|
||||
if !reflect.DeepEqual(definition.SupportedPages, wantPages) {
|
||||
t.Errorf("supported pages differ\n go %v\n js %v", definition.SupportedPages, wantPages)
|
||||
}
|
||||
|
||||
// Aliases resolve, and resolve to the same canonical id.
|
||||
for _, p := range o.Vocabulary.Pages {
|
||||
for _, alias := range append([]string{p.ID}, p.Aliases...) {
|
||||
if got := definition.CanonicalPage(alias); got != p.ID {
|
||||
t.Errorf("CanonicalPage(%q) = %q, want %q", alias, got, p.ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, table := range []struct {
|
||||
name string
|
||||
got []string
|
||||
wanted []string
|
||||
}{
|
||||
{"agent statuses", definition.AgentStatuses, o.Vocabulary.AgentStatuses},
|
||||
{"reasoning modes", definition.ReasoningModes, o.Vocabulary.Reasoning},
|
||||
{"icons", definition.AgentIcons, o.Vocabulary.Icons},
|
||||
{"knowledge kinds", definition.KnowledgeKinds, o.Vocabulary.KnowledgeKinds},
|
||||
{"access modes", definition.AgentAccess, o.Vocabulary.Access},
|
||||
{"permission roles", definition.PermissionRole, o.Vocabulary.Roles},
|
||||
} {
|
||||
if !reflect.DeepEqual(table.got, table.wanted) {
|
||||
t.Errorf("%s differ\n go %v\n js %v", table.name, table.got, table.wanted)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 3. The frontmatter tree ──────────────────────────────────────────────── */
|
||||
|
||||
// The deepest parity check available: the YAML subset must produce the same
|
||||
// data structure the JavaScript produced, for every definition and every
|
||||
// adversarial case. Not a projection of it — the whole tree.
|
||||
func TestFrontmatterTreeParity(t *testing.T) {
|
||||
for _, c := range all(load(t)) {
|
||||
t.Run(c.id(), func(t *testing.T) {
|
||||
raw := c.raw(t)
|
||||
doc, err := definition.ParseFrontmatter(raw)
|
||||
|
||||
if !c.Frontmatter.OK {
|
||||
if err == nil {
|
||||
t.Fatalf("JS refused this frontmatter (%s); Go accepted it", c.Frontmatter.Error)
|
||||
}
|
||||
if err.Error() != c.Frontmatter.Error {
|
||||
t.Errorf("error text differs\n go %q\n js %q", err.Error(), c.Frontmatter.Error)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("JS read this frontmatter; Go refused it: %v", err)
|
||||
}
|
||||
|
||||
if got, want := normalizeTree(doc.Data), normalizeTree(c.Frontmatter.Data); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("frontmatter differs\n go %s\n js %s", show(got), show(want))
|
||||
}
|
||||
if doc.Body != c.Frontmatter.Body {
|
||||
t.Errorf("body differs\n go %q\n js %q", doc.Body, c.Frontmatter.Body)
|
||||
}
|
||||
if got := definition.HasFrontmatter(raw); got != c.HasFrontmatter {
|
||||
t.Errorf("HasFrontmatter = %v, JS said %v", got, c.HasFrontmatter)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// normalizeTree puts a parsed tree into the shape `encoding/json` would have
|
||||
// produced, so the Go value and the value round-tripped through the fixture's
|
||||
// JSON are comparable. Numbers become float64 on both sides, which is what
|
||||
// JavaScript had in the first place.
|
||||
func normalizeTree(v any) any {
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return fmt.Sprintf("unmarshalable: %v", err)
|
||||
}
|
||||
var out any
|
||||
if err := json.Unmarshal(b, &out); err != nil {
|
||||
return fmt.Sprintf("unmarshalable: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func show(v any) string {
|
||||
b, _ := json.Marshal(v)
|
||||
return string(b)
|
||||
}
|
||||
|
||||
/* ── 4. Accept / reject parity ────────────────────────────────────────────── */
|
||||
|
||||
// knownDivergence is the complete list of definitions where the two parsers
|
||||
// disagree, each with the reason. It is a CLOSED list: anything not on it that
|
||||
// disagrees fails, and anything on it that stops disagreeing fails too, so the
|
||||
// list cannot quietly grow and cannot quietly go stale.
|
||||
//
|
||||
// Every entry is a case where Go is stricter, except the first — and the first
|
||||
// is the one asymmetry this package documents as deferred.
|
||||
var knownDivergence = map[string]string{
|
||||
"skill-examples/board-invalid-context.md": "" +
|
||||
"JS rejects on `ui:` placement/source semantics, which this package defers " +
|
||||
"to the frontend. Go accepts and reports Deferred: [ui].",
|
||||
|
||||
"oversized-markdown": "" +
|
||||
"JS accepts; the database refuses it (markdown_size CHECK). Go refuses it " +
|
||||
"first, so an author gets a message instead of a constraint violation.",
|
||||
"oversized-agent": "" +
|
||||
"JS accepts; the database refuses it (markdown_size CHECK). Go refuses it " +
|
||||
"first, so an author gets a message instead of a constraint violation.",
|
||||
|
||||
"version-above-int32-agent": "" +
|
||||
"JS accepts any whole number of 1 or more; agent_definitions.version is a " +
|
||||
"PostgreSQL `integer`, so the database refuses this one. Go refuses it " +
|
||||
"first, for the same reason as the size bound.",
|
||||
}
|
||||
|
||||
func TestAcceptanceParity(t *testing.T) {
|
||||
seen := map[string]bool{}
|
||||
|
||||
for _, c := range all(load(t)) {
|
||||
t.Run(c.id(), func(t *testing.T) {
|
||||
raw := c.raw(t)
|
||||
|
||||
var err error
|
||||
if c.Kind == "agent" {
|
||||
err = definition.ValidateAgent(raw)
|
||||
} else {
|
||||
err = definition.ValidateSkill(raw)
|
||||
}
|
||||
accepted := err == nil
|
||||
|
||||
if reason, expected := knownDivergence[c.id()]; expected {
|
||||
seen[c.id()] = true
|
||||
if accepted == c.Accepted {
|
||||
t.Errorf("listed as a known divergence but the two now agree (%v).\n"+
|
||||
"Remove it from knownDivergence.\n reason on file: %s", accepted, reason)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if accepted != c.Accepted {
|
||||
t.Fatalf("acceptance differs: go=%v js=%v\n go said: %v\n js said: %v",
|
||||
accepted, c.Accepted, err, deref(c.Rejection))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for id := range knownDivergence {
|
||||
if !seen[id] {
|
||||
t.Errorf("knownDivergence names %q, which is not in the fixture", id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Where both refuse a definition, they must refuse it for the same stated
|
||||
// reason. A parser that rejects the right definitions with the wrong messages
|
||||
// sends an author to the wrong line.
|
||||
func TestRejectionMessageParity(t *testing.T) {
|
||||
for _, c := range all(load(t)) {
|
||||
if c.Accepted || c.Rejection == nil {
|
||||
continue
|
||||
}
|
||||
if _, skip := knownDivergence[c.id()]; skip {
|
||||
continue
|
||||
}
|
||||
t.Run(c.id(), func(t *testing.T) {
|
||||
raw := c.raw(t)
|
||||
var err error
|
||||
if c.Kind == "agent" {
|
||||
err = definition.ValidateAgent(raw)
|
||||
} else {
|
||||
err = definition.ValidateSkill(raw)
|
||||
}
|
||||
if err == nil {
|
||||
t.Fatalf("JS rejected this; Go accepted it")
|
||||
}
|
||||
if r, ok := err.(*definition.Rejection); ok && r.BackendOnly {
|
||||
t.Fatalf("refused by a backend-only rule where JS refused it too: %q", r.Message)
|
||||
}
|
||||
if err.Error() != *c.Rejection {
|
||||
t.Errorf("rejection differs\n go %q\n js %q", err.Error(), *c.Rejection)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func deref(s *string) string {
|
||||
if s == nil {
|
||||
return "<accepted>"
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
/* ── 5. Normalized projection parity ──────────────────────────────────────── */
|
||||
|
||||
// textCoercion is the complete list of fixture cases where the JavaScript
|
||||
// record holds a value that is NOT a string in a field migration 000005
|
||||
// projects into a `text` or `text[]` column, named field by field.
|
||||
//
|
||||
// The frontend can afford this and the backend cannot. `name: [a, b]` leaves a
|
||||
// JavaScript ARRAY on skill.name, and every consumer stringifies it at the
|
||||
// point of use — the trigger list on that very record reads "a,b", which is
|
||||
// String(["a","b"]). A text column has no such option: something must be
|
||||
// written, once, at the boundary. This package writes String(x), which is the
|
||||
// string the frontend's own consumers produce.
|
||||
//
|
||||
// Listing them rather than coercing everywhere is the point. Coercion is
|
||||
// applied ONLY to the fields named here, so a genuine difference between two
|
||||
// strings still fails; each entry is checked to be still necessary, so the list
|
||||
// cannot go stale; and an unlisted case that needs coercion fails outright, so
|
||||
// the list cannot quietly grow. It is the same closed-list discipline
|
||||
// knownDivergence has, for the same reason.
|
||||
var textCoercion = map[string][]string{
|
||||
"invalid-field-type-name-list": {"name"},
|
||||
"name-list-agent": {"name"},
|
||||
"name-numeric": {"name"},
|
||||
"name-boolean": {"name"},
|
||||
"description-list": {"description"},
|
||||
"description-numeric": {"description"},
|
||||
"pages-numeric-entry": {"pages"},
|
||||
"pages-mapping-entry": {"pages"},
|
||||
}
|
||||
|
||||
// The two shapes a text projection can have, named rather than inferred.
|
||||
//
|
||||
// Which one applies is a property of the COLUMN, not of what the author
|
||||
// happened to write. `name` is `text`, so a sequence written there becomes one
|
||||
// string — String(["a","b"]) is "a,b". `pages` is `text[]`, so a sequence stays
|
||||
// a sequence and each entry becomes a string of its own. Inferring the shape
|
||||
// from whatever Go produced would make the test agree with the parser by
|
||||
// construction, which is the one thing it must not do.
|
||||
var textScalarFields = map[string]bool{
|
||||
"name": true, "description": true, "category": true, "trigger": true, "prompt": true,
|
||||
}
|
||||
|
||||
var textListFields = map[string]bool{
|
||||
"pages": true, "actions": true, "triggers": true, "skills": true, "subagents": true,
|
||||
}
|
||||
|
||||
// jsText is String(x) for a value decoded from the fixture's JSON — the same
|
||||
// conversion jsvalue.go performs inside the parser, restated here so the test
|
||||
// does not have to reach into the package it is testing to check it.
|
||||
func jsText(t *testing.T, field string, v any) any {
|
||||
t.Helper()
|
||||
|
||||
switch {
|
||||
case textScalarFields[field]:
|
||||
return jsScalarText(v)
|
||||
case textListFields[field]:
|
||||
list, ok := v.([]any)
|
||||
if !ok {
|
||||
return jsScalarText(v)
|
||||
}
|
||||
out := make([]any, len(list))
|
||||
for i, item := range list {
|
||||
out[i] = jsScalarText(item)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
t.Fatalf("textCoercion names %q, which is not a text-projected field", field)
|
||||
return nil
|
||||
}
|
||||
|
||||
func jsScalarText(v any) string {
|
||||
switch x := v.(type) {
|
||||
case nil:
|
||||
return "null"
|
||||
case bool:
|
||||
if x {
|
||||
return "true"
|
||||
}
|
||||
return "false"
|
||||
case float64:
|
||||
if x == float64(int64(x)) {
|
||||
return strconv.FormatInt(int64(x), 10)
|
||||
}
|
||||
return strconv.FormatFloat(x, 'g', -1, 64)
|
||||
case string:
|
||||
return x
|
||||
case []any:
|
||||
// Array.prototype.toString: nil renders as the empty string, not
|
||||
// "null", which is the one place the two differ.
|
||||
parts := make([]string, len(x))
|
||||
for i, item := range x {
|
||||
if item == nil {
|
||||
continue
|
||||
}
|
||||
parts[i] = jsScalarText(item)
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
case map[string]any:
|
||||
return "[object Object]"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// The fields migration 000005 projects into columns, compared for every
|
||||
// definition both parsers accept. These are the values that reach the
|
||||
// database, so a difference here is a row the frontend would render wrongly.
|
||||
func TestProjectionParity(t *testing.T) {
|
||||
fixture := map[string]bool{}
|
||||
|
||||
for _, c := range all(load(t)) {
|
||||
fixture[c.id()] = true
|
||||
if c.Normalized == nil {
|
||||
continue // JS could not parse it; covered by the tree test
|
||||
}
|
||||
coerce := map[string]bool{}
|
||||
for _, field := range textCoercion[c.id()] {
|
||||
coerce[field] = true
|
||||
}
|
||||
|
||||
t.Run(c.id(), func(t *testing.T) {
|
||||
raw := c.raw(t)
|
||||
|
||||
if c.Kind == "agent" {
|
||||
agent, err := definition.ParseAgent(raw, definition.Options{})
|
||||
if err != nil {
|
||||
t.Fatalf("JS parsed this; Go refused it: %v", err)
|
||||
}
|
||||
compare(t, map[string]any{
|
||||
"id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"description": agent.Description,
|
||||
"status": agent.Status,
|
||||
"version": agent.Version,
|
||||
"pages": agent.Pages,
|
||||
"icon": agent.Icon,
|
||||
"reasoning": agent.Reasoning,
|
||||
"trigger": agent.Trigger,
|
||||
"webSearch": agent.WebSearch,
|
||||
"skills": agent.Skills,
|
||||
"subagents": agent.Subagents,
|
||||
"starters": agent.Starters,
|
||||
"permissions": agent.Permissions,
|
||||
"errors": agent.Errors,
|
||||
}, c.Normalized, coerce)
|
||||
return
|
||||
}
|
||||
|
||||
skill, err := definition.ParseSkill(raw, definition.Options{})
|
||||
if err != nil {
|
||||
t.Fatalf("JS parsed this; Go refused it: %v", err)
|
||||
}
|
||||
compare(t, map[string]any{
|
||||
"id": skill.ID,
|
||||
"name": skill.Name,
|
||||
"description": skill.Description,
|
||||
"status": skill.Status,
|
||||
"pages": skill.Pages,
|
||||
"kind": skill.Kind,
|
||||
"category": skill.Category,
|
||||
"actions": skill.Actions,
|
||||
"triggers": skill.Triggers,
|
||||
"declaredTriggers": skill.DeclaredTriggers,
|
||||
"prompt": skill.Prompt,
|
||||
"skillId": skill.SkillID,
|
||||
}, c.Normalized, coerce)
|
||||
})
|
||||
}
|
||||
|
||||
for id := range textCoercion {
|
||||
if !fixture[id] {
|
||||
t.Errorf("textCoercion names %q, which is not in the fixture", id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// compare checks every field Go produced against the JS record, field by field
|
||||
// so a failure names the field rather than dumping two objects.
|
||||
func compare(t *testing.T, got map[string]any, want map[string]any, coerce map[string]bool) {
|
||||
t.Helper()
|
||||
|
||||
keys := make([]string, 0, len(got))
|
||||
for k := range got {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
|
||||
for _, k := range keys {
|
||||
wantValue, present := want[k]
|
||||
if !present {
|
||||
t.Errorf("%s: absent from the JS record", k)
|
||||
continue
|
||||
}
|
||||
if coerce[k] {
|
||||
// Listed in textCoercion. Check the entry is still earning its
|
||||
// place before honouring it: if the JS value is already the string
|
||||
// Go produced, the coercion is doing nothing and the list has gone
|
||||
// stale.
|
||||
coerced := jsText(t, k, normalizeTree(wantValue))
|
||||
if reflect.DeepEqual(normalizeTree(wantValue), normalizeTree(coerced)) {
|
||||
t.Errorf("%s: listed in textCoercion, but the JS value is already "+
|
||||
"a string. Remove the entry.", k)
|
||||
}
|
||||
wantValue = coerced
|
||||
}
|
||||
|
||||
g, w := normalizeTree(got[k]), normalizeTree(wantValue)
|
||||
// An empty list and a missing one are the same thing to both parsers.
|
||||
if isEmptyList(g) && isEmptyList(w) {
|
||||
continue
|
||||
}
|
||||
if !reflect.DeepEqual(g, w) {
|
||||
t.Errorf("%s differs\n go %s\n js %s", k, show(g), show(w))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func isEmptyList(v any) bool {
|
||||
if v == nil {
|
||||
return true
|
||||
}
|
||||
l, ok := v.([]any)
|
||||
return ok && len(l) == 0
|
||||
}
|
||||
|
||||
/* ── 6. The parser never rewrites what is stored ──────────────────────────── */
|
||||
|
||||
// Migration 000005 keeps `markdown` verbatim and derives every other column
|
||||
// from it. Parsing must therefore be a read: normalization exists to
|
||||
// INTERPRET a definition, never to rewrite it.
|
||||
func TestParsingDoesNotMutateSource(t *testing.T) {
|
||||
for _, c := range all(load(t)) {
|
||||
raw := c.raw(t)
|
||||
before := string(append([]byte{}, raw...))
|
||||
|
||||
_, _ = definition.ParseSkill(raw, definition.Options{})
|
||||
_, _ = definition.ParseAgent(raw, definition.Options{})
|
||||
_ = definition.ValidateSkill(raw)
|
||||
_ = definition.ValidateAgent(raw)
|
||||
|
||||
if raw != before {
|
||||
t.Fatalf("%s: the source changed under the parser", c.id())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Normalize is what the parser reads THROUGH; what it returns must never be
|
||||
// what gets stored. Asserted directly, because the whole separation rests on
|
||||
// it: the corpus contains definitions whose normalized form differs from their
|
||||
// stored form, and storing the normalized one would silently rewrite an
|
||||
// author's file.
|
||||
func TestNormalizationIsNotStorage(t *testing.T) {
|
||||
rewritten := 0
|
||||
for _, c := range all(load(t)) {
|
||||
raw := c.raw(t)
|
||||
if definition.Normalize(raw) != raw {
|
||||
rewritten++
|
||||
}
|
||||
}
|
||||
if rewritten == 0 {
|
||||
t.Fatal("no case in the corpus is changed by Normalize; " +
|
||||
"this test can no longer tell storage and interpretation apart")
|
||||
}
|
||||
t.Logf("%d of %d definitions normalize to something other than their stored bytes", rewritten, len(all(load(t))))
|
||||
}
|
||||
|
||||
/* ── 7. Adversarial coverage is real ──────────────────────────────────────── */
|
||||
|
||||
// The adversarial cases Phase 4D requires, each mapped to the fixture rows that
|
||||
// exercise it. A case list that drifts away from the requirement is a suite
|
||||
// that looks thorough and tests something else.
|
||||
func TestAdversarialCoverage(t *testing.T) {
|
||||
required := map[string][]string{
|
||||
"UTF-8 BOM": {"utf8-bom", "utf8-bom-agent", "bom-crlf-blankline"},
|
||||
"CRLF": {"crlf", "crlf-agent", "crlf-inside-frontmatter-only"},
|
||||
"CR": {"cr-only"},
|
||||
"leading blank line": {"leading-blank-line"},
|
||||
"multiple leading blank lines": {"multiple-leading-blank-lines", "leading-spaces-then-blank-lines"},
|
||||
"trailing spaces": {"trailing-spaces-on-values"},
|
||||
"trailing newline": {"many-trailing-newlines", "no-trailing-newline"},
|
||||
"trailing ws after fence": {"trailing-ws-after-open-fence", "trailing-tab-after-close-fence"},
|
||||
"quoted scalar": {"double-quoted-scalar", "doubled-quote-escape"},
|
||||
"single-quoted scalar": {"single-quoted-scalar"},
|
||||
"colon inside quoted string": {"colon-in-quoted-string", "colon-in-unquoted-string"},
|
||||
"hash inside quoted string": {"hash-in-quoted-string", "hash-unquoted-trailing-comment", "hash-unquoted-midword"},
|
||||
"empty scalar": {"empty-scalar", "tilde-scalar", "null-scalar"},
|
||||
"empty array": {"empty-array"},
|
||||
"inline array": {"inline-flow-array", "inline-flow-map"},
|
||||
"multiline scalar": {"block-scalar-literal", "block-scalar-folded"},
|
||||
"duplicate key": {"duplicate-key", "duplicate-key-array"},
|
||||
"malformed YAML": {"malformed-yaml-bare-line", "key-with-space", "ragged-indent"},
|
||||
"malformed opening fence": {"malformed-open-fence-two-dashes", "malformed-open-fence-four-dashes",
|
||||
"malformed-open-fence-indented", "malformed-open-fence-text-after"},
|
||||
"malformed closing fence": {"malformed-close-fence-two-dashes", "malformed-close-fence-missing",
|
||||
"malformed-close-fence-four-dashes"},
|
||||
"missing frontmatter": {"missing-frontmatter", "empty-fence-pair", "frontmatter-is-a-sequence"},
|
||||
"unsupported frontmatter field": {"unsupported-frontmatter-field", "unsupported-field-agent", "uppercase-key"},
|
||||
"invalid field type": {"invalid-field-type-pages-scalar", "invalid-field-type-pages-scalar-agent",
|
||||
"invalid-field-type-name-list"},
|
||||
"invalid definition id": {"invalid-definition-id-uppercase", "invalid-definition-id-leading-dash",
|
||||
"invalid-definition-id-underscore"},
|
||||
"invalid status": {"invalid-status-skill", "invalid-status-agent", "inactive-status-skill"},
|
||||
"invalid visibility": {"visibility-field-personal", "visibility-field-invalid"},
|
||||
"oversized markdown": {"oversized-markdown", "oversized-agent", "at-size-bound"},
|
||||
"empty markdown": {"empty-markdown", "whitespace-only-markdown"},
|
||||
}
|
||||
|
||||
present := map[string]bool{}
|
||||
for _, c := range load(t).Cases {
|
||||
present[c.Name] = true
|
||||
}
|
||||
|
||||
for requirement, names := range required {
|
||||
for _, n := range names {
|
||||
if !present[n] {
|
||||
t.Errorf("%q: the fixture has no case named %q", requirement, n)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8. Deferred blocks are reported, not assumed ─────────────────────────── */
|
||||
|
||||
// Every skill carrying a `ui:` or `owliver:` block must say so, because that is
|
||||
// the one part of validation this package does not do. A block that stopped
|
||||
// being reported would be a gap nobody could see.
|
||||
func TestDeferredBlocksAreReported(t *testing.T) {
|
||||
o := load(t)
|
||||
found := 0
|
||||
|
||||
for _, c := range append(append([]observation{}, o.Corpus...), o.Cases...) {
|
||||
if c.Kind != "skill" || !c.Frontmatter.OK {
|
||||
continue
|
||||
}
|
||||
want := []string{}
|
||||
for _, key := range []string{"ui", "owliver"} {
|
||||
if _, present := c.Frontmatter.Data[key]; present {
|
||||
want = append(want, key)
|
||||
}
|
||||
}
|
||||
|
||||
skill, err := definition.ParseSkill(c.raw(t), definition.Options{})
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if len(want) == 0 {
|
||||
if len(skill.Deferred) != 0 {
|
||||
t.Errorf("%s: reported Deferred %v with no such block", c.id(), skill.Deferred)
|
||||
}
|
||||
continue
|
||||
}
|
||||
found++
|
||||
if !reflect.DeepEqual(skill.Deferred, want) {
|
||||
t.Errorf("%s: Deferred = %v, want %v", c.id(), skill.Deferred, want)
|
||||
}
|
||||
}
|
||||
|
||||
if found != 19 {
|
||||
t.Errorf("definitions carrying a deferred block: got %d, want 19", found)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 9. Mutation checks ───────────────────────────────────────────────────── */
|
||||
|
||||
// Tests that pass against a broken parser are not tests. Each mutation below
|
||||
// is a plausible mistake in this package; every one must be caught by a real
|
||||
// definition changing its meaning, not by an assertion written to notice it.
|
||||
func TestMutationsWouldBeCaught(t *testing.T) {
|
||||
base := strings.Join([]string{
|
||||
"---",
|
||||
"id: sample-skill",
|
||||
"name: Sample Skill",
|
||||
"description: A sample.",
|
||||
"pages:",
|
||||
" - candidates",
|
||||
"---",
|
||||
"",
|
||||
"# Sample Skill",
|
||||
}, "\n")
|
||||
|
||||
mutations := []struct {
|
||||
name string
|
||||
raw string
|
||||
check func(t *testing.T, s *definition.Skill, err error)
|
||||
}{
|
||||
{
|
||||
// Dropping the BOM strip: the fence stops matching and every field
|
||||
// empties out.
|
||||
name: "BOM before the fence still fences",
|
||||
raw: "\uFEFF" + base,
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.ID != "sample-skill" || len(s.Pages) != 1 {
|
||||
t.Errorf("got id=%q pages=%v err=%v", s.ID, s.Pages, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Dropping CR normalization: `candidates\r` is not a page.
|
||||
name: "CRLF endings do not leak into values",
|
||||
raw: strings.ReplaceAll(base, "\n", "\r\n"),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || len(s.Pages) != 1 || s.Pages[0] != "candidates" {
|
||||
t.Errorf("got pages=%v err=%v", s.Pages, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Trimming the closing fence too eagerly, or not at all.
|
||||
name: "trailing tab after the closing fence still closes it",
|
||||
raw: strings.Replace(base, "\n---\n", "\n---\t\n", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.Name != "Sample Skill" {
|
||||
t.Errorf("got name=%q err=%v", s.Name, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// A greedy fence would swallow the second document and lose the id.
|
||||
name: "a second --- document is body, not frontmatter",
|
||||
raw: base + "\n\n---\nid: second\n---\n",
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.ID != "sample-skill" {
|
||||
t.Errorf("got id=%q err=%v", s.ID, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Treating `#` as always starting a comment.
|
||||
name: "a hash inside a word is part of the word",
|
||||
raw: strings.Replace(base, "description: A sample.", "category: ops#1", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.Category != "ops#1" {
|
||||
t.Errorf("got category=%q err=%v", s.Category, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Treating a spaced `#` as part of the value.
|
||||
name: "a spaced hash starts a comment",
|
||||
raw: strings.Replace(base, "description: A sample.", "category: ops # note", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.Category != "ops" {
|
||||
t.Errorf("got category=%q err=%v", s.Category, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Splitting a quoted value on its colon.
|
||||
name: "a colon inside quotes stays in the value",
|
||||
raw: strings.Replace(base, "name: Sample Skill", `name: "Sample: Skill"`, 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.Name != "Sample: Skill" {
|
||||
t.Errorf("got name=%q err=%v", s.Name, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Keeping the first duplicate rather than the last.
|
||||
name: "a duplicate key takes the last value",
|
||||
raw: strings.Replace(base, "name: Sample Skill", "name: First\nname: Second", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || s.Name != "Second" {
|
||||
t.Errorf("got name=%q err=%v", s.Name, err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Accepting ragged indentation instead of refusing it.
|
||||
name: "ragged indentation is refused with its line",
|
||||
raw: strings.Replace(base, " - candidates", " - candidates\n - positions", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
var pe *definition.Error
|
||||
if err == nil {
|
||||
t.Fatalf("accepted ragged indentation: %+v", s)
|
||||
}
|
||||
if !asError(err, &pe) || pe.Line != 6 {
|
||||
t.Errorf("got %v, want an *Error on line 6", err)
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
// Canonicalising a skill's pages, which the frontend does not do.
|
||||
name: "a skill keeps the page name as written",
|
||||
raw: strings.Replace(base, " - candidates", " - Talent Pool", 1),
|
||||
check: func(t *testing.T, s *definition.Skill, err error) {
|
||||
if err != nil || len(s.Pages) != 1 || s.Pages[0] != "Talent Pool" {
|
||||
t.Errorf("got pages=%v err=%v", s.Pages, err)
|
||||
}
|
||||
if err := definition.ValidateSkill(strings.Replace(base, " - candidates", " - Talent Pool", 1)); err != nil {
|
||||
t.Errorf("an aliased page should still validate: %v", err)
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, m := range mutations {
|
||||
t.Run(m.name, func(t *testing.T) {
|
||||
skill, err := definition.ParseSkill(m.raw, definition.Options{})
|
||||
if skill == nil {
|
||||
skill = &definition.Skill{}
|
||||
}
|
||||
m.check(t, skill, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// asError is errors.As, spelled out for the one concrete type this package
|
||||
// returns.
|
||||
func asError(err error, target **definition.Error) bool {
|
||||
e, ok := err.(*definition.Error)
|
||||
if ok {
|
||||
*target = e
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
/* ── 10. Bounds ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// The size bound is the backend's, and both directions of it matter: a
|
||||
// definition at the limit must be storable and one character more must not.
|
||||
// The companion to TestSizeBound, for the other backend-only bound. The
|
||||
// fixture pins a version well above the bound and one exactly at it, which
|
||||
// leaves the step between them untested — an off-by-one there would refuse a
|
||||
// version PostgreSQL can store, or accept one it cannot. Both sides of the
|
||||
// step are named here so that cannot happen.
|
||||
func TestVersionBound(t *testing.T) {
|
||||
agent := func(version string) string {
|
||||
return "---\nid: sample-agent\nname: Sample Agent\npages:\n - candidates\n" +
|
||||
"version: " + version + "\n---\n\n# Sample Agent\n"
|
||||
}
|
||||
|
||||
at := strconv.Itoa(definition.MaxVersion)
|
||||
if err := definition.ValidateAgent(agent(at)); err != nil {
|
||||
t.Errorf("version %s, exactly at the bound, was refused: %v", at, err)
|
||||
}
|
||||
|
||||
over := strconv.FormatInt(int64(definition.MaxVersion)+1, 10)
|
||||
err := definition.ValidateAgent(agent(over))
|
||||
if err == nil {
|
||||
t.Fatalf("version %s, one past the bound, was accepted", over)
|
||||
}
|
||||
r, ok := err.(*definition.Rejection)
|
||||
if !ok || !r.BackendOnly {
|
||||
t.Errorf("the version bound should be reported as a backend-only rule, got %v", err)
|
||||
}
|
||||
|
||||
// The bound belongs to validation, not to parsing: a version the database
|
||||
// cannot store must still normalize to the number the author wrote, or the
|
||||
// record the editor shows and the record Go builds would disagree.
|
||||
parsed, err := definition.ParseAgent(agent(over), definition.Options{})
|
||||
if err != nil {
|
||||
t.Fatalf("parsing a too-large version failed: %v", err)
|
||||
}
|
||||
if got := strconv.Itoa(parsed.Version); got != over {
|
||||
t.Errorf("parse clamped the version to %s; it should carry %s", got, over)
|
||||
}
|
||||
if len(parsed.Errors) != 0 {
|
||||
t.Errorf("parse reported a backend-only bound as an authoring error: %v", parsed.Errors)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSizeBound(t *testing.T) {
|
||||
head := "---\nid: sample-skill\nname: Sample Skill\npages:\n - candidates\n---\n\n"
|
||||
|
||||
at := head + strings.Repeat("y", definition.MaxMarkdownLength-len(head))
|
||||
if n := len([]rune(at)); n != definition.MaxMarkdownLength {
|
||||
t.Fatalf("fixture is %d characters, wanted exactly %d", n, definition.MaxMarkdownLength)
|
||||
}
|
||||
if err := definition.ValidateSkill(at); err != nil {
|
||||
t.Errorf("a definition exactly at the bound was refused: %v", err)
|
||||
}
|
||||
|
||||
over := at + "y"
|
||||
err := definition.ValidateSkill(over)
|
||||
if err == nil {
|
||||
t.Fatal("a definition one character over the bound was accepted")
|
||||
}
|
||||
r, ok := err.(*definition.Rejection)
|
||||
if !ok || !r.BackendOnly {
|
||||
t.Errorf("the size bound should be reported as a backend-only rule, got %v", err)
|
||||
}
|
||||
}
|
||||
102
go-api/internal/definition/definition.go
Normal file
102
go-api/internal/definition/definition.go
Normal file
@@ -0,0 +1,102 @@
|
||||
// Package definition reads Krow agent and skill definitions — Markdown with
|
||||
// YAML frontmatter — the way the frontend reads them.
|
||||
//
|
||||
// # Why this exists
|
||||
//
|
||||
// A definition is authored in the browser and stored by the server, so two
|
||||
// parsers see it: the JavaScript in src/lib/skills and src/lib/agents, and
|
||||
// this one. If they disagree, one of two things happens, and both are silent:
|
||||
//
|
||||
// - A definition the editor accepts and this package rejects looks valid
|
||||
// while it is being written and fails when it is saved.
|
||||
// - A definition this package accepts and the editor rejects is stored and
|
||||
// then cannot be rendered by the product that owns it.
|
||||
//
|
||||
// Compatibility is therefore a contract rather than an aspiration, and it is
|
||||
// enforced by a conformance suite (conformance_test.go) that replays the
|
||||
// ACTUAL output of the JavaScript parser — captured from the real frontend
|
||||
// module graph — against this one, over all 37 shipped definitions and every
|
||||
// adversarial case in testdata/oracle.json. The corpus count is pinned by a
|
||||
// test; the case count is deliberately not restated here, because a number
|
||||
// kept in a comment is a number that goes stale.
|
||||
//
|
||||
// # What is in the contract
|
||||
//
|
||||
// - The document layer: byte-order mark, line endings, leading blank lines,
|
||||
// fence recognition, body extraction. See frontmatter.go.
|
||||
// - The YAML subset: block maps and sequences, scalars, quoting, comments.
|
||||
// A port of yaml.js, with no YAML dependency, deliberately — see yaml.go.
|
||||
// - Definition-level normalization and validation: id, name, description,
|
||||
// status, version, pages, icons, reasoning, permissions, starters,
|
||||
// knowledge, subagents.
|
||||
//
|
||||
// Those cover every column migration 000005 projects out of a definition:
|
||||
// definition_id, status, version, name, description, pages.
|
||||
//
|
||||
// # What is deferred, and why
|
||||
//
|
||||
// A skill may carry a `ui:` block (declarative page sections) or an `owliver:`
|
||||
// block (assistant capabilities). Validating those means reproducing roughly
|
||||
// 1,500 lines of closed vocabulary describing what the FRONTEND can render —
|
||||
// placements, data sources, section types, periods — none of which the backend
|
||||
// stores, projects, or acts on.
|
||||
//
|
||||
// This package therefore does not check them. It records their presence on
|
||||
// Skill.Deferred instead, so the gap is a value a caller can see rather than
|
||||
// an assumption. The one consequence is stated exactly:
|
||||
//
|
||||
// skill-examples/board-invalid-context.md is rejected by the frontend, on a
|
||||
// rule about which placement can supply which data source, and accepted
|
||||
// here. It is the only definition in the corpus where the two disagree, and
|
||||
// the conformance suite asserts that it stays the only one.
|
||||
//
|
||||
// # What is not the parser's job
|
||||
//
|
||||
// Normalization never rewrites what is stored. Migration 000005 keeps
|
||||
// `markdown` verbatim and every other column is derived from it; this package
|
||||
// only ever reads. The Markdown handed in is the Markdown that goes to the
|
||||
// database, byte for byte, and a test asserts it.
|
||||
//
|
||||
// Visibility (personal or organization) is deliberately absent. It is not a
|
||||
// frontmatter field — the frontend ignores `visibility:` in a definition
|
||||
// entirely — it is a storage tier chosen by the request and checked by the
|
||||
// visibility CHECK in migration 000005. A definition cannot name its own
|
||||
// tenancy.
|
||||
package definition
|
||||
|
||||
// MaxMarkdownLength is the markdown_size CHECK from migration 000005, in
|
||||
// CHARACTERS — `length()` in PostgreSQL counts characters, not bytes.
|
||||
//
|
||||
// The frontend does NOT enforce this, so a definition longer than this is one
|
||||
// the editor accepts and the database refuses. This package refuses it first,
|
||||
// which turns a constraint violation into a message an author can act on.
|
||||
const MaxMarkdownLength = 65536
|
||||
|
||||
// MaxVersion is the range of agent_definitions.version, a PostgreSQL
|
||||
// `integer`.
|
||||
//
|
||||
// The frontend accepts any integer of 1 or more, so a version above this is
|
||||
// another value the editor accepts and the database cannot store.
|
||||
const MaxVersion = 2147483647
|
||||
|
||||
// maxExactInteger is 2^53-1, the largest integer a float64 names exactly and so
|
||||
// the largest a JavaScript number carries without loss. It bounds the version
|
||||
// conversion in ParseAgent; it is not a rule about what may be stored, which is
|
||||
// MaxVersion's job.
|
||||
const maxExactInteger = 1<<53 - 1
|
||||
|
||||
// Rejection is a definition that parses but may not be stored.
|
||||
//
|
||||
// Message is the frontend's own wording wherever the rule is shared, so the
|
||||
// editor and the API describe the same problem the same way.
|
||||
type Rejection struct {
|
||||
Message string
|
||||
|
||||
// BackendOnly marks a rule the frontend does not have — a bound the
|
||||
// database imposes that the editor never checks. These are the only
|
||||
// messages that can differ from what an author would see in the browser,
|
||||
// and each one is listed in docs/phase-4d-parser-contract.md.
|
||||
BackendOnly bool
|
||||
}
|
||||
|
||||
func (r *Rejection) Error() string { return r.Message }
|
||||
332
go-api/internal/definition/frontmatter.go
Normal file
332
go-api/internal/definition/frontmatter.go
Normal file
@@ -0,0 +1,332 @@
|
||||
package definition
|
||||
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// The document layer: what is frontmatter, what is body, and what a `## Heading`
|
||||
// section contains.
|
||||
//
|
||||
// A port of the four exported readers in src/lib/skills/registry.js —
|
||||
// normalizeDefinition, hasFrontmatter, parseFrontmatter and the section
|
||||
// readers. Agent and skill definitions are read by the same code on the
|
||||
// frontend, deliberately, so that the two formats cannot drift; the same is
|
||||
// true here.
|
||||
//
|
||||
// The regular expressions the JavaScript uses are hand-rolled rather than
|
||||
// translated, because two of them rely on lookahead and lazy matching that RE2
|
||||
// does not have. Each is written out below with the JavaScript it reproduces.
|
||||
|
||||
// Normalize is the frontend's `normalizeDefinition`: a definition's text as the
|
||||
// parser needs to see it.
|
||||
//
|
||||
// Files arrive from editors, from Windows, from copy-paste and from downloads,
|
||||
// and four of the things they arrive with used to take the whole frontmatter
|
||||
// block down — a UTF-8 byte-order mark before the opening fence, blank lines
|
||||
// above it, CRLF endings, and trailing spaces after `---`. In each case the
|
||||
// fence did not match and the definition registered as untitled with no pages.
|
||||
//
|
||||
// This is NOT a lenient parser. The subset inside the fences is exactly as
|
||||
// strict as it was. This is only about recognising that a fence is a fence.
|
||||
//
|
||||
// String(raw ?? '')
|
||||
// .replace(/^\uFEFF/, '')
|
||||
// .replace(/\r\n?/g, '\n')
|
||||
// .replace(/^\s*\n+/, '')
|
||||
//
|
||||
// The order is load bearing: the BOM goes first so it cannot be counted as the
|
||||
// leading whitespace, and CR normalisation goes before the blank-line strip so
|
||||
// that a CRLF blank line is one.
|
||||
func Normalize(raw string) string {
|
||||
text := strings.TrimPrefix(raw, "\uFEFF")
|
||||
|
||||
// `\r\n?` → `\n`: a CRLF pair and a lone CR both become one newline.
|
||||
if strings.IndexByte(text, '\r') >= 0 {
|
||||
var b strings.Builder
|
||||
b.Grow(len(text))
|
||||
for i := 0; i < len(text); i++ {
|
||||
if text[i] != '\r' {
|
||||
b.WriteByte(text[i])
|
||||
continue
|
||||
}
|
||||
b.WriteByte('\n')
|
||||
if i+1 < len(text) && text[i+1] == '\n' {
|
||||
i++
|
||||
}
|
||||
}
|
||||
text = b.String()
|
||||
}
|
||||
|
||||
// `^\s*\n+` → ``. Greedy `\s*` then at least one newline: the effect is to
|
||||
// drop the leading whitespace run up to and including its LAST newline, and
|
||||
// to drop nothing at all when that run contains no newline. A definition
|
||||
// indented by one space is therefore still unfenced, which is what the
|
||||
// editor decides too.
|
||||
end, last := 0, -1
|
||||
for i, r := range text {
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
if r == '\n' {
|
||||
last = i
|
||||
}
|
||||
end = i + len(string(r))
|
||||
}
|
||||
_ = end
|
||||
if last >= 0 {
|
||||
text = text[last+1:]
|
||||
}
|
||||
|
||||
return text
|
||||
}
|
||||
|
||||
// fence locates the frontmatter block in already-normalized text.
|
||||
//
|
||||
// /^---[ \t]*\n([\s\S]*?)\n---[ \t]*(?=\n|$)/
|
||||
//
|
||||
// Returns the YAML source, the offset just past the closing fence, and whether
|
||||
// there was one. Lazy: the FIRST closing fence wins, which is why a definition
|
||||
// carrying a second `---` document keeps only the first and reads the rest as
|
||||
// body.
|
||||
func fence(text string) (yaml string, end int, ok bool) {
|
||||
if !strings.HasPrefix(text, "---") {
|
||||
return "", 0, false
|
||||
}
|
||||
i := 3
|
||||
for i < len(text) && (text[i] == ' ' || text[i] == '\t') {
|
||||
i++
|
||||
}
|
||||
if i >= len(text) || text[i] != '\n' {
|
||||
return "", 0, false
|
||||
}
|
||||
start := i + 1
|
||||
|
||||
for at := start - 1; at >= 0 && at < len(text); {
|
||||
nl := strings.IndexByte(text[at+1:], '\n')
|
||||
if nl < 0 {
|
||||
return "", 0, false
|
||||
}
|
||||
at = at + 1 + nl // index of the newline that must precede the fence
|
||||
|
||||
rest := text[at+1:]
|
||||
if !strings.HasPrefix(rest, "---") {
|
||||
continue
|
||||
}
|
||||
j := 3
|
||||
for j < len(rest) && (rest[j] == ' ' || rest[j] == '\t') {
|
||||
j++
|
||||
}
|
||||
// `(?=\n|$)` — end of the document, or the end of this line. `$` has no
|
||||
// multiline flag on the frontend either, so it means end of document.
|
||||
if j < len(rest) && rest[j] != '\n' {
|
||||
continue
|
||||
}
|
||||
return text[start:at], at + 1 + j, true
|
||||
}
|
||||
return "", 0, false
|
||||
}
|
||||
|
||||
// HasFrontmatter reports whether this text opens with a frontmatter block at
|
||||
// all. It does not say whether that block parses.
|
||||
func HasFrontmatter(raw string) bool {
|
||||
_, _, ok := fence(Normalize(raw))
|
||||
return ok
|
||||
}
|
||||
|
||||
// Document is a definition split into its two halves.
|
||||
type Document struct {
|
||||
// Data is the frontmatter as plain data. Always a mapping: a frontmatter
|
||||
// block that parses to a sequence is discarded, exactly as the frontend
|
||||
// discards it, because every reader downstream indexes it by key.
|
||||
Data map[string]any
|
||||
|
||||
// Body is everything after the closing fence, trimmed. A document with no
|
||||
// frontmatter is all body.
|
||||
Body string
|
||||
|
||||
// Fenced records whether a frontmatter block was found, which Data alone
|
||||
// cannot express — an empty fence pair and a missing one both give an
|
||||
// empty mapping.
|
||||
Fenced bool
|
||||
}
|
||||
|
||||
// ParseFrontmatter splits a definition and reads its frontmatter.
|
||||
//
|
||||
// Returns a *Error when the YAML subset refuses a line. A document with no
|
||||
// recognisable fence is NOT an error: it is a document with no frontmatter,
|
||||
// and what happens to it is the validator's decision — the same division the
|
||||
// frontend makes.
|
||||
func ParseFrontmatter(raw string) (Document, error) {
|
||||
text := Normalize(raw)
|
||||
|
||||
yaml, end, ok := fence(text)
|
||||
if !ok {
|
||||
return Document{Data: map[string]any{}, Body: text, Fenced: false}, nil
|
||||
}
|
||||
|
||||
value, err := ParseYAML(yaml)
|
||||
if err != nil {
|
||||
return Document{}, err
|
||||
}
|
||||
|
||||
data, _ := value.(map[string]any)
|
||||
if data == nil {
|
||||
data = map[string]any{}
|
||||
}
|
||||
|
||||
return Document{Data: data, Body: jsTrim(text[end:]), Fenced: true}, nil
|
||||
}
|
||||
|
||||
// sectionSource is the text under a `## Heading`, up to the next one.
|
||||
//
|
||||
// new RegExp(`##\\s+${escaped}\\s*\\n([\\s\\S]*?)(?=\\n##\\s|$)`, 'i')
|
||||
//
|
||||
// Case-insensitive, and deliberately not anchored to the start of a line —
|
||||
// that is what the frontend does. The heading is compared literally rather
|
||||
// than compiled into a pattern, which is the same protection the frontend gets
|
||||
// by escaping it: a heading containing regular-expression punctuation must
|
||||
// match the words it is built from.
|
||||
//
|
||||
// Returns ok=false for a section that is not there, which is a different thing
|
||||
// from a section that is there and empty.
|
||||
func sectionSource(body, heading string) (string, bool) {
|
||||
lower := strings.ToLower(body)
|
||||
want := strings.ToLower(heading)
|
||||
|
||||
for at := 0; ; {
|
||||
h := strings.Index(lower[at:], "##")
|
||||
if h < 0 {
|
||||
return "", false
|
||||
}
|
||||
h += at
|
||||
at = h + 2
|
||||
|
||||
// `##` then `\s+` then the heading.
|
||||
i := h + 2
|
||||
gap := i
|
||||
for i < len(body) {
|
||||
r, size := decodeRune(body[i:])
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
i += size
|
||||
}
|
||||
if i == gap {
|
||||
continue // `\s+` needs at least one
|
||||
}
|
||||
if !strings.HasPrefix(lower[i:], want) {
|
||||
continue
|
||||
}
|
||||
i += len(want)
|
||||
|
||||
// `\s*\n`: a whitespace run that ends in a newline.
|
||||
j, nl := i, -1
|
||||
for j < len(body) {
|
||||
r, size := decodeRune(body[j:])
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
if r == '\n' {
|
||||
nl = j
|
||||
break
|
||||
}
|
||||
j += size
|
||||
}
|
||||
if nl < 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
start := nl + 1
|
||||
// `(?=\n##\s|$)`, lazily: the first following line that opens a new
|
||||
// `##` heading. `###` does not, because the character after `##` must
|
||||
// be whitespace.
|
||||
for k := start; ; {
|
||||
n := strings.Index(body[k:], "\n##")
|
||||
if n < 0 {
|
||||
return body[start:], true
|
||||
}
|
||||
n += k
|
||||
after := n + 3
|
||||
if after < len(body) {
|
||||
r, _ := decodeRune(body[after:])
|
||||
if jsIsSpace(r) {
|
||||
return body[start:n], true
|
||||
}
|
||||
}
|
||||
k = n + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// decodeRune is utf8.DecodeRuneInString, kept local so the section reader has
|
||||
// one obvious way to step through the body.
|
||||
func decodeRune(s string) (rune, int) {
|
||||
for i, r := range s {
|
||||
_ = i
|
||||
return r, len(string(r))
|
||||
}
|
||||
return 0, 0
|
||||
}
|
||||
|
||||
// SectionText is the prose under a `## Heading`, with its bullets and blank
|
||||
// lines flattened to one line.
|
||||
//
|
||||
// Used for a workforce level's description, which is a sentence rather than a
|
||||
// list.
|
||||
func SectionText(body, heading string) string {
|
||||
source, ok := sectionSource(body, heading)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
parts := []string{}
|
||||
for _, l := range strings.Split(source, "\n") {
|
||||
l = jsTrim(stripListMarker(l))
|
||||
if l == "" {
|
||||
continue
|
||||
}
|
||||
parts = append(parts, l)
|
||||
}
|
||||
return jsTrim(strings.Join(parts, " "))
|
||||
}
|
||||
|
||||
// stripListMarker removes a leading `-`, `*` or `1.` / `1)` bullet.
|
||||
//
|
||||
// /^\s*(?:[-*]|\d+[.)])\s+/
|
||||
func stripListMarker(l string) string {
|
||||
i := 0
|
||||
for i < len(l) {
|
||||
r, size := decodeRune(l[i:])
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
i += size
|
||||
}
|
||||
marker := i
|
||||
switch {
|
||||
case i < len(l) && (l[i] == '-' || l[i] == '*'):
|
||||
i++
|
||||
default:
|
||||
digits := i
|
||||
for i < len(l) && l[i] >= '0' && l[i] <= '9' {
|
||||
i++
|
||||
}
|
||||
if i == digits || i >= len(l) || (l[i] != '.' && l[i] != ')') {
|
||||
return l
|
||||
}
|
||||
i++
|
||||
}
|
||||
// `\s+` after the marker is required; without it there is no list item.
|
||||
space := i
|
||||
for i < len(l) {
|
||||
r, size := decodeRune(l[i:])
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
i += size
|
||||
}
|
||||
if i == space {
|
||||
return l
|
||||
}
|
||||
_ = marker
|
||||
return l[i:]
|
||||
}
|
||||
210
go-api/internal/definition/jsvalue.go
Normal file
210
go-api/internal/definition/jsvalue.go
Normal file
@@ -0,0 +1,210 @@
|
||||
package definition
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// JavaScript value semantics, reproduced exactly.
|
||||
//
|
||||
// The frontend parser is JavaScript, and the compatibility contract is with
|
||||
// THAT parser, not with an idealised YAML. Three of its behaviours are load
|
||||
// bearing and none of them are Go's defaults:
|
||||
//
|
||||
// - `\s` and `String.prototype.trim` cover a different set of code points
|
||||
// than `unicode.IsSpace`. JS treats U+FEFF as whitespace and U+0085 as
|
||||
// not; Go is the other way round. A definition is trimmed on the way
|
||||
// through the parser at least four times, so the difference is reachable.
|
||||
// - Truthiness decides whether `id:` is used or derived, whether a name
|
||||
// falls back to `Untitled skill`, and what `prompt:` becomes. `0`, `false`
|
||||
// and `""` are falsy; `"0"` and `[]` are not.
|
||||
// - `String(x)` and `Number(x)` have defined results for every type, and the
|
||||
// validator interpolates them into messages an author reads. `[1,2]`
|
||||
// stringifies to `1,2`, an object to `[object Object]`.
|
||||
//
|
||||
// Reimplementing these is not gold-plating: each one is exercised by a case in
|
||||
// the conformance suite because each one is reachable from a definition an
|
||||
// author could write.
|
||||
|
||||
// jsIsSpace reports whether r is whitespace to JavaScript — the union of
|
||||
// WhiteSpace and LineTerminator in the specification.
|
||||
//
|
||||
// Deliberately NOT unicode.IsSpace: that set includes U+0085 (NEL), which JS
|
||||
// does not, and excludes U+FEFF, which JS does.
|
||||
func jsIsSpace(r rune) bool {
|
||||
switch r {
|
||||
case '\t', '\n', '\v', '\f', '\r', ' ',
|
||||
0x00A0, 0x1680, 0x2028, 0x2029, 0x202F, 0x205F, 0x3000, 0xFEFF:
|
||||
return true
|
||||
}
|
||||
return r >= 0x2000 && r <= 0x200A
|
||||
}
|
||||
|
||||
// jsTrim is String.prototype.trim.
|
||||
func jsTrim(s string) string { return strings.TrimFunc(s, jsIsSpace) }
|
||||
|
||||
// jsTrimStart is String.prototype.trimStart.
|
||||
func jsTrimStart(s string) string { return strings.TrimLeftFunc(s, jsIsSpace) }
|
||||
|
||||
// jsTruthy is the `!!x` of a parsed YAML value.
|
||||
//
|
||||
// The parser produces only nil, bool, float64, string, []any and
|
||||
// map[string]any, so those are the only cases that can arise. An empty array
|
||||
// and an empty object are both truthy in JavaScript, which is why they are not
|
||||
// listed alongside the empty string.
|
||||
func jsTruthy(v any) bool {
|
||||
switch x := v.(type) {
|
||||
case nil:
|
||||
return false
|
||||
case bool:
|
||||
return x
|
||||
case float64:
|
||||
return x != 0 && !math.IsNaN(x)
|
||||
case string:
|
||||
return x != ""
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// jsNumberToString is JavaScript's Number → String conversion for the values
|
||||
// this parser can produce.
|
||||
//
|
||||
// `-0` prints as `0`, integers print without a decimal point, and everything
|
||||
// else takes the shortest representation that round-trips — which is what
|
||||
// strconv's 'g' with precision -1 gives, in the range a definition can reach.
|
||||
func jsNumberToString(f float64) string {
|
||||
switch {
|
||||
case math.IsNaN(f):
|
||||
return "NaN"
|
||||
case math.IsInf(f, 1):
|
||||
return "Infinity"
|
||||
case math.IsInf(f, -1):
|
||||
return "-Infinity"
|
||||
case f == 0:
|
||||
return "0" // collapses -0
|
||||
}
|
||||
if f == math.Trunc(f) && math.Abs(f) < 1e21 {
|
||||
return strconv.FormatFloat(f, 'f', -1, 64)
|
||||
}
|
||||
return strconv.FormatFloat(f, 'g', -1, 64)
|
||||
}
|
||||
|
||||
// jsString is the `String(x)` of a parsed YAML value.
|
||||
//
|
||||
// Arrays join on `,` with nil rendering as the empty string, which is
|
||||
// Array.prototype.toString; a mapping renders as `[object Object]`. Both are
|
||||
// reachable: `trigger:` may be written as a list, and the message an author
|
||||
// reads interpolates the result.
|
||||
func jsString(v any) string {
|
||||
switch x := v.(type) {
|
||||
case nil:
|
||||
return "null"
|
||||
case bool:
|
||||
if x {
|
||||
return "true"
|
||||
}
|
||||
return "false"
|
||||
case float64:
|
||||
return jsNumberToString(x)
|
||||
case string:
|
||||
return x
|
||||
case []any:
|
||||
parts := make([]string, len(x))
|
||||
for i, item := range x {
|
||||
if item == nil {
|
||||
parts[i] = ""
|
||||
continue
|
||||
}
|
||||
parts[i] = jsString(item)
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
case map[string]any:
|
||||
return "[object Object]"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// jsTrimmed is the frontend's `trimmed()` helper: String(value ?? empty).trim().
|
||||
//
|
||||
// The nullish coalescing matters: a nil renders as the empty string here,
|
||||
// where a bare String(null) would render as the four characters `null`.
|
||||
func jsTrimmed(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
return jsTrim(jsString(v))
|
||||
}
|
||||
|
||||
// jsNumber is the `Number(x)` of a parsed YAML value, NaN where JavaScript
|
||||
// gives NaN.
|
||||
//
|
||||
// Only reached from `version:`, where the result is checked with
|
||||
// Number.isInteger. The string cases below are the ones a YAML scalar can
|
||||
// still be carrying at that point: a quoted `"3"` stays a string, and so does
|
||||
// anything the numeric patterns in toScalar declined.
|
||||
func jsNumber(v any) float64 {
|
||||
switch x := v.(type) {
|
||||
case nil:
|
||||
return 0
|
||||
case bool:
|
||||
if x {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
case float64:
|
||||
return x
|
||||
case string:
|
||||
return jsNumberFromString(x)
|
||||
case []any:
|
||||
// Number([]) is 0 and Number([3]) is 3, via the same String()
|
||||
// conversion; anything longer stringifies with a comma and fails.
|
||||
if len(x) == 0 {
|
||||
return 0
|
||||
}
|
||||
if len(x) == 1 {
|
||||
return jsNumberFromString(jsString(x[0]))
|
||||
}
|
||||
}
|
||||
return math.NaN()
|
||||
}
|
||||
|
||||
func jsNumberFromString(s string) float64 {
|
||||
s = jsTrim(s)
|
||||
if s == "" {
|
||||
return 0
|
||||
}
|
||||
switch s {
|
||||
case "Infinity", "+Infinity":
|
||||
return math.Inf(1)
|
||||
case "-Infinity":
|
||||
return math.Inf(-1)
|
||||
}
|
||||
// The radix prefixes JavaScript accepts in a numeric string literal. Signs
|
||||
// are not permitted with them, which ParseUint enforces by rejecting the
|
||||
// leading character.
|
||||
if len(s) > 2 && s[0] == '0' {
|
||||
var base int
|
||||
switch s[1] {
|
||||
case 'x', 'X':
|
||||
base = 16
|
||||
case 'o', 'O':
|
||||
base = 8
|
||||
case 'b', 'B':
|
||||
base = 2
|
||||
}
|
||||
if base != 0 {
|
||||
n, err := strconv.ParseUint(s[2:], base, 64)
|
||||
if err != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
return float64(n)
|
||||
}
|
||||
}
|
||||
f, err := strconv.ParseFloat(s, 64)
|
||||
if err != nil {
|
||||
return math.NaN()
|
||||
}
|
||||
return f
|
||||
}
|
||||
349
go-api/internal/definition/skill.go
Normal file
349
go-api/internal/definition/skill.go
Normal file
@@ -0,0 +1,349 @@
|
||||
package definition
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// One Markdown definition → one skill.
|
||||
//
|
||||
// A port of parseSkill and validateSkillSource in src/lib/skills/registry.js,
|
||||
// restricted to the definition contract — see the package documentation in
|
||||
// definition.go for exactly where that boundary is and why the `ui:` and
|
||||
// `owliver:` blocks are on the other side of it.
|
||||
|
||||
// Level is one rung of a workforce ladder, read from the body's own headings.
|
||||
type Level struct {
|
||||
Level string `json:"level"`
|
||||
Label string `json:"label"`
|
||||
Summary string `json:"summary"`
|
||||
}
|
||||
|
||||
// Skill is a definition as the backend reads it.
|
||||
//
|
||||
// The five fields migration 000005 projects into columns — ID, Name,
|
||||
// Description, Status, Pages — are the compatibility contract; the rest is
|
||||
// carried because it is free once the frontmatter is parsed and because the
|
||||
// conformance suite compares it.
|
||||
type Skill struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
|
||||
// Pages as the AUTHOR WROTE THEM, not canonicalised.
|
||||
//
|
||||
// This is not an oversight and must not be "fixed": parseSkill keeps the
|
||||
// declared strings, so a skill written against `Talent Pool` is registered
|
||||
// under `Talent Pool` and resolved through normalizeKey at every use.
|
||||
// Agents are the other way round — see Agent.Pages. Canonicalising here
|
||||
// would make the backend's projection disagree with the editor's.
|
||||
Pages []string `json:"pages"`
|
||||
|
||||
Kind string `json:"kind"`
|
||||
Category string `json:"category"`
|
||||
Actions []string `json:"actions"`
|
||||
Triggers []string `json:"triggers"`
|
||||
DeclaredTriggers bool `json:"declaredTriggers"`
|
||||
Prompt *string `json:"prompt"`
|
||||
SkillID *string `json:"skillId"`
|
||||
Levels []Level `json:"levels"`
|
||||
|
||||
// Body is the Markdown after the frontmatter, trimmed. The definition
|
||||
// itself is NEVER rewritten — see Definition.Markdown.
|
||||
Body string `json:"-"`
|
||||
|
||||
// Deferred names the frontmatter blocks whose semantics this package does
|
||||
// not check and the frontend does. Empty for every definition the backend
|
||||
// can fully validate on its own. See package documentation.
|
||||
Deferred []string `json:"deferred,omitempty"`
|
||||
}
|
||||
|
||||
// AuthoredPath is the origin an authored definition has when the caller names
|
||||
// none. It is a value rather than an absence for one reason: it is the
|
||||
// frontend's own default parameter.
|
||||
//
|
||||
// parseSkill(raw, { path = 'custom', custom = false } = {})
|
||||
// parseAgent(raw, { path = 'custom', custom = false } = {})
|
||||
//
|
||||
// validateSkillSource and validateAgentSource both call their parser with no
|
||||
// path, so every definition the EDITOR checks derives its last-resort id from
|
||||
// the literal string `custom`. That is the same call the backend is making — a
|
||||
// definition submitted to the API is authored, not shipped — so the backend
|
||||
// must derive the same id.
|
||||
//
|
||||
// The difference is reachable and it is not cosmetic. A definition with no
|
||||
// `id:`, no `name:` and a valid `pages:` list gets the id `custom` on the
|
||||
// frontend, passes the id-format check and is ACCEPTED. Deriving no id here
|
||||
// would refuse it with "The frontmatter needs an `id`." — a definition that
|
||||
// validates in the editor and fails on save, which is the exact failure mode
|
||||
// this package exists to prevent. Fixture case: id-omitted-unnamed.
|
||||
const AuthoredPath = "custom"
|
||||
|
||||
// Options carries what the caller knows that the definition does not.
|
||||
type Options struct {
|
||||
// Path is the definition's origin, used only as the last fallback for an
|
||||
// id. Leave it empty for anything authored rather than shipped — which is
|
||||
// what the backend always has — and it becomes AuthoredPath, exactly as the
|
||||
// frontend's default parameter does.
|
||||
Path string
|
||||
}
|
||||
|
||||
// path is the origin an id is derived from, with the frontend's default
|
||||
// applied.
|
||||
func (o Options) path() string {
|
||||
if o.Path == "" {
|
||||
return AuthoredPath
|
||||
}
|
||||
return o.Path
|
||||
}
|
||||
|
||||
// ParseSkill reads a skill definition.
|
||||
//
|
||||
// Returns a *Error when the frontmatter cannot be read. A definition with no
|
||||
// frontmatter at all is not an error here: it parses to a skill carrying the
|
||||
// derived id and no pages, and ValidateSkill is what refuses it — the same
|
||||
// division of labour the frontend has.
|
||||
func ParseSkill(raw string, opts Options) (*Skill, error) {
|
||||
doc, err := ParseFrontmatter(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data := doc.Data
|
||||
|
||||
declaredPages, _ := data["pages"].([]any)
|
||||
|
||||
// `id:` always wins. An explicit id is the address other definitions and
|
||||
// stored preferences refer to, and deriving over the top of one would
|
||||
// silently rename a skill. Slugging the name is what an author means by
|
||||
// leaving it out; the filename is right only for a file, which is why it
|
||||
// is last.
|
||||
id := jsTrim(jsString(data["id"]))
|
||||
if !jsTruthy(data["id"]) {
|
||||
id = slugify(data["name"])
|
||||
if id == "" {
|
||||
id = fileStem(opts.path())
|
||||
}
|
||||
}
|
||||
|
||||
levels := sectionLevels(doc.Body)
|
||||
|
||||
// Two things wear the same format. A definition that names a ladder is a
|
||||
// workforce skill; nothing else distinguishes them, so an author declares
|
||||
// one by writing one rather than by setting a flag.
|
||||
kind := jsTrim(jsString(data["kind"]))
|
||||
if !jsTruthy(data["kind"]) {
|
||||
kind = "assistant"
|
||||
if len(levels) > 0 {
|
||||
kind = "workforce"
|
||||
}
|
||||
}
|
||||
|
||||
skill := &Skill{
|
||||
ID: id,
|
||||
Kind: kind,
|
||||
Levels: levels,
|
||||
Body: doc.Body,
|
||||
Name: "Untitled skill",
|
||||
Pages: stringsOf(declaredPages),
|
||||
}
|
||||
|
||||
if jsTruthy(data["name"]) {
|
||||
skill.Name = jsString(data["name"])
|
||||
}
|
||||
if jsTruthy(data["description"]) {
|
||||
skill.Description = jsString(data["description"])
|
||||
}
|
||||
if s, ok := data["category"].(string); ok {
|
||||
skill.Category = jsTrim(s)
|
||||
}
|
||||
|
||||
// The whole of a skill's lifecycle, and deliberately a coercion rather
|
||||
// than a check: the frontend reads anything that is not `inactive` as
|
||||
// `active`, so `status: bogus` registers as active rather than being
|
||||
// refused. Reproduced, not corrected — see the divergence note in
|
||||
// docs/phase-4d-parser-contract.md.
|
||||
skill.Status = "active"
|
||||
if s, ok := data["status"].(string); ok && s == "inactive" {
|
||||
skill.Status = "inactive"
|
||||
}
|
||||
|
||||
if actions, ok := data["actions"].([]any); ok {
|
||||
skill.Actions = stringsOf(actions)
|
||||
} else {
|
||||
skill.Actions = []string{}
|
||||
}
|
||||
|
||||
// A skill with no declared triggers answers to its own name, so a
|
||||
// definition that omits the field is still reachable by asking for it.
|
||||
// Explicit triggers replace the fallback rather than adding to it.
|
||||
triggers, hasTriggers := data["triggers"].([]any)
|
||||
skill.DeclaredTriggers = hasTriggers && len(triggers) > 0
|
||||
skill.Triggers = []string{}
|
||||
if skill.DeclaredTriggers {
|
||||
for _, t := range triggers {
|
||||
skill.Triggers = append(skill.Triggers, strings.ToLower(jsString(t)))
|
||||
}
|
||||
} else if jsTruthy(data["name"]) {
|
||||
skill.Triggers = append(skill.Triggers, strings.ToLower(jsString(data["name"])))
|
||||
}
|
||||
|
||||
if jsTruthy(data["prompt"]) {
|
||||
p := jsString(data["prompt"])
|
||||
skill.Prompt = &p
|
||||
}
|
||||
|
||||
// The capability in the skill graph a workforce definition governs:
|
||||
// `skill:`, or the id with a `-training` suffix dropped and dashes swapped
|
||||
// for underscores.
|
||||
if kind == "workforce" {
|
||||
base := id
|
||||
if jsTruthy(data["skill"]) {
|
||||
base = jsString(data["skill"])
|
||||
} else {
|
||||
base = strings.TrimSuffix(base, "-training")
|
||||
}
|
||||
s := strings.ReplaceAll(base, "-", "_")
|
||||
skill.SkillID = &s
|
||||
}
|
||||
|
||||
skill.Deferred = deferredBlocks(data)
|
||||
return skill, nil
|
||||
}
|
||||
|
||||
// sectionLevels reads the ladder a workforce definition defines, in order, from
|
||||
// the body's own headings. A rung with no prose is not a rung.
|
||||
func sectionLevels(body string) []Level {
|
||||
out := []Level{}
|
||||
for _, heading := range levelHeadings {
|
||||
summary := SectionText(body, heading)
|
||||
if summary == "" {
|
||||
continue
|
||||
}
|
||||
out = append(out, Level{
|
||||
Level: strings.ToLower(heading),
|
||||
Label: heading,
|
||||
Summary: summary,
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// deferredBlocks names the frontmatter this package does not semantically
|
||||
// check. See the package documentation for why they are deferred rather than
|
||||
// validated or rejected.
|
||||
func deferredBlocks(data map[string]any) []string {
|
||||
out := []string{}
|
||||
for _, key := range []string{"ui", "owliver"} {
|
||||
if _, present := data[key]; present {
|
||||
out = append(out, key)
|
||||
}
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// stringsOf renders a parsed sequence as the strings the frontend would read
|
||||
// out of it. Non-string entries are stringified rather than dropped, because
|
||||
// that is what every consumer of `pages` and `actions` does with them.
|
||||
func stringsOf(list []any) []string {
|
||||
out := make([]string, 0, len(list))
|
||||
for _, v := range list {
|
||||
out = append(out, jsString(v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func fileStem(path string) string {
|
||||
if path == "" {
|
||||
return ""
|
||||
}
|
||||
if i := strings.LastIndexByte(path, '/'); i >= 0 {
|
||||
path = path[i+1:]
|
||||
}
|
||||
return strings.TrimSuffix(path, ".md")
|
||||
}
|
||||
|
||||
// ValidateSkill decides whether a skill definition may be stored.
|
||||
//
|
||||
// Returns nil when it may. The order is the order an author would fix things
|
||||
// in, and every message below is the frontend's message character for
|
||||
// character — an author who sees one in the editor and a different one from
|
||||
// the API is being told about two different problems.
|
||||
//
|
||||
// Two rules are the backend's own and are marked as such: the size bound and
|
||||
// the deferred-block rule. Both are explained in
|
||||
// docs/phase-4d-parser-contract.md.
|
||||
func ValidateSkill(raw string) error {
|
||||
if jsTrim(raw) == "" {
|
||||
return &Rejection{Message: "Paste or upload a Markdown definition."}
|
||||
}
|
||||
|
||||
// The backend's own rule, from migration 000005's markdown_size CHECK. The
|
||||
// editor does not enforce it, so a definition over the bound is one the
|
||||
// frontend accepts and the DATABASE refuses; refusing it here turns a
|
||||
// constraint violation into a message. See the contract document.
|
||||
if n := len([]rune(raw)); n > MaxMarkdownLength {
|
||||
return &Rejection{
|
||||
BackendOnly: true,
|
||||
Message: fmt.Sprintf(
|
||||
"That definition is %d characters. The limit is %d.", n, MaxMarkdownLength),
|
||||
}
|
||||
}
|
||||
|
||||
skill, err := ParseSkill(raw, Options{})
|
||||
if err != nil {
|
||||
// The subset reports the line it failed on, which is far more useful
|
||||
// than "could not be parsed".
|
||||
return &Rejection{Message: jsTrim("That definition could not be parsed. " + err.Error())}
|
||||
}
|
||||
|
||||
if skill.ID == "" {
|
||||
return &Rejection{Message: "The frontmatter needs an `id`."}
|
||||
}
|
||||
if !isDefinitionID(skill.ID) {
|
||||
return &Rejection{Message: "`id` must be lower-case letters, numbers and dashes."}
|
||||
}
|
||||
// Faithful to the frontend, where `name` has already fallen back to
|
||||
// `Untitled skill` and this check can therefore never fire. Kept so the
|
||||
// two validators have the same shape and the same order.
|
||||
if skill.Name == "" {
|
||||
return &Rejection{Message: "The frontmatter needs a `name`."}
|
||||
}
|
||||
// The backend's own rule, and it must be asked BEFORE the generic one
|
||||
// below. A `ui:` block declares the pages it draws on, and parseSkill falls
|
||||
// back to those pages when `pages:` is absent — a fallback this package
|
||||
// cannot compute, because it does not read the `ui:` vocabulary. Rather
|
||||
// than report an empty page list it never really established, say what is
|
||||
// actually missing. No shipped definition relies on the fallback: all
|
||||
// nineteen that carry a `ui:` or `owliver:` block also declare `pages:`.
|
||||
if len(skill.Pages) == 0 && len(skill.Deferred) > 0 {
|
||||
return &Rejection{
|
||||
BackendOnly: true,
|
||||
Message: "A definition with a `ui:` block needs an explicit `pages:` list.",
|
||||
}
|
||||
}
|
||||
if len(skill.Pages) == 0 {
|
||||
return &Rejection{Message: "The frontmatter needs at least one `pages` entry."}
|
||||
}
|
||||
|
||||
unknown := []string{}
|
||||
for _, p := range skill.Pages {
|
||||
if !SurfaceExists(p) {
|
||||
unknown = append(unknown, p)
|
||||
}
|
||||
}
|
||||
if len(unknown) > 0 {
|
||||
plural := ""
|
||||
if len(unknown) > 1 {
|
||||
plural = "s"
|
||||
}
|
||||
return &Rejection{Message: fmt.Sprintf(
|
||||
"Unsupported page%s: %s. Supported pages: %s.",
|
||||
plural, strings.Join(unknown, ", "), strings.Join(SupportedPages, ", "))}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
9633
go-api/internal/definition/testdata/oracle.json
vendored
Normal file
9633
go-api/internal/definition/testdata/oracle.json
vendored
Normal file
File diff suppressed because one or more lines are too long
184
go-api/internal/definition/vocabulary.go
Normal file
184
go-api/internal/definition/vocabulary.go
Normal file
@@ -0,0 +1,184 @@
|
||||
package definition
|
||||
|
||||
import "strings"
|
||||
|
||||
// The closed vocabulary a definition is allowed to name.
|
||||
//
|
||||
// Every table here is a transcription of a frontend table, and the conformance
|
||||
// suite asserts each one against the vocabulary the JavaScript actually
|
||||
// exports (testdata/oracle.json, `vocabulary`) — so a page added to
|
||||
// surfaces.js or an icon added to vocabulary.js fails a test here rather than
|
||||
// silently making the two ends disagree about what is valid.
|
||||
//
|
||||
// Nothing is looked up dynamically and nothing is constructed: a definition
|
||||
// names a key and this file answers whether that key exists.
|
||||
|
||||
// pageSurfaces mirrors SKILL_SURFACES in src/lib/skills/surfaces.js — id first,
|
||||
// then the alternative spellings an author may use for it.
|
||||
var pageSurfaces = []struct {
|
||||
ID string
|
||||
Aliases []string
|
||||
}{
|
||||
{ID: "control-center"},
|
||||
{ID: "positions"},
|
||||
{ID: "create-position", Aliases: []string{"new-position"}},
|
||||
{ID: "candidates"},
|
||||
{ID: "hired-history", Aliases: []string{"hired"}},
|
||||
{ID: "talent-pool"},
|
||||
{ID: "krow-forge", Aliases: []string{"university", "forge"}},
|
||||
{ID: "analytics"},
|
||||
{ID: "activity"},
|
||||
{ID: "workspace-agent-configure"},
|
||||
{ID: "settings"},
|
||||
{ID: "workspace"},
|
||||
{ID: "workspace-agents"},
|
||||
{ID: "workspace-skills"},
|
||||
{ID: "workspace-skill-configure"},
|
||||
{ID: "skill-development"},
|
||||
{ID: "profile"},
|
||||
{ID: "candidates-analysis"},
|
||||
}
|
||||
|
||||
// surfaceByKey resolves every spelling — canonical or alias — to its canonical
|
||||
// id.
|
||||
var surfaceByKey = func() map[string]string {
|
||||
m := map[string]string{}
|
||||
for _, s := range pageSurfaces {
|
||||
m[s.ID] = s.ID
|
||||
for _, a := range s.Aliases {
|
||||
m[a] = s.ID
|
||||
}
|
||||
}
|
||||
return m
|
||||
}()
|
||||
|
||||
// SupportedPages is every canonical page id, in declaration order. The order is
|
||||
// the order the frontend lists them in when it refuses an unsupported page, and
|
||||
// that message is compared byte for byte.
|
||||
var SupportedPages = func() []string {
|
||||
out := make([]string, len(pageSurfaces))
|
||||
for i, s := range pageSurfaces {
|
||||
out[i] = s.ID
|
||||
}
|
||||
return out
|
||||
}()
|
||||
|
||||
// normalizeKey is surfaces.js's own: trimmed, lower-cased, with spaces and
|
||||
// underscores read as dashes. `Talent Pool` and `talent_pool` both reach
|
||||
// `talent-pool`; the canonical keys never widen, only what an author may type
|
||||
// to reach them.
|
||||
func normalizeKey(page any) string {
|
||||
s := strings.ToLower(jsTrim(jsString(page)))
|
||||
if page == nil {
|
||||
s = ""
|
||||
}
|
||||
return strings.Map(func(r rune) rune {
|
||||
if r == ' ' || r == '_' || jsIsSpace(r) {
|
||||
return '-'
|
||||
}
|
||||
return r
|
||||
}, s)
|
||||
}
|
||||
|
||||
// SurfaceExists reports whether a declared page name refers to a real surface.
|
||||
func SurfaceExists(page any) bool {
|
||||
_, ok := surfaceByKey[normalizeKey(page)]
|
||||
return ok
|
||||
}
|
||||
|
||||
// CanonicalPage is the canonical id a declared page name refers to, or "".
|
||||
func CanonicalPage(page any) string { return surfaceByKey[normalizeKey(page)] }
|
||||
|
||||
/* ── Agent vocabulary — src/lib/agents/vocabulary.js ─────────────────────── */
|
||||
|
||||
var (
|
||||
AgentStatuses = []string{"draft", "published", "archived"}
|
||||
ReasoningModes = []string{"fast", "balanced", "deep"}
|
||||
KnowledgeKinds = []string{"note", "link", "skill-reference"}
|
||||
AgentAccess = []string{"all", "specific"}
|
||||
PermissionRole = []string{"manager", "editor", "viewer"}
|
||||
AgentIcons = []string{
|
||||
"owliver", "sparkles", "briefcase", "users", "user-check",
|
||||
"layers", "graduation-cap", "bar-chart", "activity", "shield",
|
||||
}
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultAgentStatus = "draft"
|
||||
DefaultReasoning = "balanced"
|
||||
DefaultAgentIcon = "owliver"
|
||||
DefaultKnowledgeKind = "note"
|
||||
DefaultAgentAccess = "all"
|
||||
DefaultPermission = "viewer"
|
||||
)
|
||||
|
||||
func contains(list []string, want string) bool {
|
||||
for _, v := range list {
|
||||
if v == want {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
/* ── Skill vocabulary ────────────────────────────────────────────────────── */
|
||||
|
||||
// SkillStatuses is the whole of a skill's lifecycle. There is no version and no
|
||||
// publish step: migration 000005 states the same two values.
|
||||
var SkillStatuses = []string{"active", "inactive"}
|
||||
|
||||
// levelHeadings are the rungs a workforce skill may define, and the reason a
|
||||
// definition is read as workforce rather than assistant. Order is the ladder's.
|
||||
var levelHeadings = []string{"Beginner", "Intermediate", "Advanced", "Expert"}
|
||||
|
||||
// slugify is uiConfig.js's, used to derive an id from a name.
|
||||
//
|
||||
// String(value || '').toLowerCase().trim()
|
||||
// .replace(/[^a-z0-9]+/g, '-').replace(/^-|-$/g, '')
|
||||
//
|
||||
// The lower-casing happens BEFORE the class replacement, so an upper-case
|
||||
// letter becomes itself rather than a dash.
|
||||
func slugify(value any) string {
|
||||
if !jsTruthy(value) {
|
||||
return ""
|
||||
}
|
||||
s := jsTrim(strings.ToLower(jsString(value)))
|
||||
|
||||
var b strings.Builder
|
||||
dash := false
|
||||
for _, r := range s {
|
||||
if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') {
|
||||
b.WriteRune(r)
|
||||
dash = false
|
||||
continue
|
||||
}
|
||||
if !dash {
|
||||
b.WriteByte('-')
|
||||
dash = true
|
||||
}
|
||||
}
|
||||
out := b.String()
|
||||
out = strings.TrimPrefix(out, "-")
|
||||
out = strings.TrimSuffix(out, "-")
|
||||
return out
|
||||
}
|
||||
|
||||
// isDefinitionID is the id format the frontend validator enforces and the
|
||||
// definition_id CHECK in migration 000005 restates.
|
||||
func isDefinitionID(id string) bool {
|
||||
if id == "" {
|
||||
return false
|
||||
}
|
||||
for i, r := range id {
|
||||
lower := r >= 'a' && r <= 'z'
|
||||
digit := r >= '0' && r <= '9'
|
||||
if lower || digit {
|
||||
continue
|
||||
}
|
||||
if r == '-' && i > 0 {
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
362
go-api/internal/definition/yaml.go
Normal file
362
go-api/internal/definition/yaml.go
Normal file
@@ -0,0 +1,362 @@
|
||||
package definition
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// The YAML subset a Krow definition is allowed to use.
|
||||
//
|
||||
// A line-for-line port of the frontend's src/lib/skills/yaml.js. That file is
|
||||
// the specification and this is the second implementation of it, so every
|
||||
// decision below is made because the JavaScript makes it — including the ones
|
||||
// a YAML library would make differently.
|
||||
//
|
||||
// Supported, and nothing else:
|
||||
//
|
||||
// - block maps and block sequences, nested to any depth
|
||||
// - scalars: strings, integers, floats, booleans, null
|
||||
// - quoted strings, for values containing `:` or `#`
|
||||
// - `- key: value`, a mapping whose first key sits on the dash
|
||||
// - `#` comments, and blank lines
|
||||
//
|
||||
// Anchors, aliases, merge keys, multi-document files, flow mappings, flow
|
||||
// sequences, block scalars and tags are NOT supported, and are not silently
|
||||
// half-read: an unparseable line is an error carrying the line number, so a
|
||||
// definition either means what it says or is refused with somewhere to look.
|
||||
//
|
||||
// No dependency. A general YAML library would accept a much larger language
|
||||
// than the frontend does, and every construct it accepted and the frontend did
|
||||
// not would be a definition the backend stores and the editor cannot read —
|
||||
// exactly the failure this package exists to prevent. The subset is small
|
||||
// enough to port exactly, so it is ported exactly.
|
||||
//
|
||||
// Nothing here evaluates anything. There is no reflection, no template, no
|
||||
// code path from a definition to execution of any kind: a definition is
|
||||
// configuration, and this is the boundary that keeps it configuration.
|
||||
|
||||
// Error is a definition that could not be read, carrying the line it failed on.
|
||||
//
|
||||
// Line is 1-based and counts lines of FRONTMATTER, not of the file — which is
|
||||
// what the JavaScript reports, because parseYaml is handed the fenced block
|
||||
// rather than the document. Reproduced rather than improved: an author who
|
||||
// sees one message in the editor and another from the API is being told about
|
||||
// two different problems.
|
||||
type Error struct {
|
||||
Line int
|
||||
Message string
|
||||
}
|
||||
|
||||
func (e *Error) Error() string { return e.Message }
|
||||
|
||||
// errIndent and errPair are the two failures the subset has, worded exactly as
|
||||
// the frontend words them.
|
||||
func errIndent(line int) *Error {
|
||||
return &Error{Line: line, Message: fmt.Sprintf("Unexpected indentation on line %d", line)}
|
||||
}
|
||||
|
||||
func errPair(line int, content string) *Error {
|
||||
return &Error{Line: line, Message: fmt.Sprintf("Line %d is not `key: value`: %s", line, content)}
|
||||
}
|
||||
|
||||
// keyPair matches `key: value` and is the only shape a mapping entry may take.
|
||||
// The key alphabet is the frontend's: letters, digits, underscore, dot, dash —
|
||||
// which is why `my key: value` is a refusal rather than a key with a space.
|
||||
// jsSpaceClass is the JavaScript `\s` character class. Go's own `\s` is
|
||||
// ASCII-only, and the difference is reachable: a non-breaking space after the
|
||||
// colon is whitespace to the editor's parser and would be part of the value
|
||||
// here.
|
||||
const jsSpaceClass = `[\t\n\v\f\r \x{00A0}\x{1680}\x{2000}-\x{200A}\x{2028}\x{2029}\x{202F}\x{205F}\x{3000}\x{FEFF}]`
|
||||
|
||||
var keyPair = regexp.MustCompile(`^([A-Za-z0-9_.-]+):` + jsSpaceClass + `*([\s\S]*)$`)
|
||||
|
||||
var (
|
||||
reInt = regexp.MustCompile(`^-?[0-9]+$`)
|
||||
reFloat = regexp.MustCompile(`^-?[0-9]*\.[0-9]+$`)
|
||||
)
|
||||
|
||||
// line is one significant line, reduced to what the parser needs to decide.
|
||||
type line struct {
|
||||
number int // 1-based, within the frontmatter block
|
||||
indent int // leading whitespace, tabs counted as two
|
||||
content string
|
||||
}
|
||||
|
||||
// readLines drops blank lines and whole-line comments, and measures what is
|
||||
// left.
|
||||
//
|
||||
// Indentation is counted in code points with a tab worth two spaces, which is
|
||||
// what the JavaScript does and is why a tab-indented sequence sits at the same
|
||||
// depth as a two-space one.
|
||||
func readLines(source string) []line {
|
||||
out := []line{}
|
||||
for i, text := range strings.Split(source, "\n") {
|
||||
trimmed := jsTrim(text)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
// A whole-line comment: `^\s*#`.
|
||||
if strings.HasPrefix(jsTrimStart(text), "#") {
|
||||
continue
|
||||
}
|
||||
indent := 0
|
||||
for _, r := range text {
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
if r == '\t' {
|
||||
indent += 2
|
||||
continue
|
||||
}
|
||||
indent++
|
||||
}
|
||||
out = append(out, line{number: i + 1, indent: indent, content: trimmed})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// quoted matches a scalar wrapped in one kind of quote, end to end.
|
||||
//
|
||||
// Greedy and anchored at both ends, as in the frontend: `"a" "b"` is therefore
|
||||
// ONE quoted string whose content is `a" "b`, not two. That is a strange
|
||||
// reading, and it is the reading the editor gives, so it is the reading here.
|
||||
func quotedScalar(value string) (string, bool) {
|
||||
if len(value) < 2 {
|
||||
return "", false
|
||||
}
|
||||
q := value[0]
|
||||
if q != '\'' && q != '"' {
|
||||
return "", false
|
||||
}
|
||||
if value[len(value)-1] != q {
|
||||
return "", false
|
||||
}
|
||||
inner := value[1 : len(value)-1]
|
||||
// The only escape the subset has: a doubled quote is one quote.
|
||||
return strings.ReplaceAll(inner, string([]byte{q, q}), string(q)), true
|
||||
}
|
||||
|
||||
// stripTrailingComment removes an unquoted trailing `#` comment.
|
||||
//
|
||||
// `\s+#.*$` applied once, leftmost — so `ops # a # b` loses everything from
|
||||
// the first spaced hash, and `ops#1` loses nothing, because a hash inside a
|
||||
// word is part of the word.
|
||||
func stripTrailingComment(value string) string {
|
||||
runes := []rune(value)
|
||||
for i := 0; i < len(runes); i++ {
|
||||
if !jsIsSpace(runes[i]) {
|
||||
continue
|
||||
}
|
||||
j := i
|
||||
for j < len(runes) && jsIsSpace(runes[j]) {
|
||||
j++
|
||||
}
|
||||
if j < len(runes) && runes[j] == '#' {
|
||||
return string(runes[:i])
|
||||
}
|
||||
i = j - 1
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// toScalar reads one written value: `true`, `false`, `null`, a number, a
|
||||
// quoted string, or the string as written.
|
||||
func toScalar(raw string) any {
|
||||
value := jsTrim(raw)
|
||||
|
||||
switch value {
|
||||
case "", "~", "null":
|
||||
return nil
|
||||
case "true":
|
||||
return true
|
||||
case "false":
|
||||
return false
|
||||
}
|
||||
|
||||
// Quoted: taken literally, which is how a value containing `:` or `#` is
|
||||
// written. No escape processing beyond the doubled quote.
|
||||
if inner, ok := quotedScalar(value); ok {
|
||||
return inner
|
||||
}
|
||||
|
||||
if reInt.MatchString(value) || reFloat.MatchString(value) {
|
||||
if f, err := strconv.ParseFloat(value, 64); err == nil {
|
||||
return f
|
||||
}
|
||||
}
|
||||
|
||||
return jsTrim(stripTrailingComment(value))
|
||||
}
|
||||
|
||||
// cursor is shared down the recursion so a child consumes the lines it owns.
|
||||
type cursor struct{ i int }
|
||||
|
||||
// parseBlock reads one block at indent or deeper.
|
||||
//
|
||||
// Map or sequence depending on what the first line at this level is, which is
|
||||
// how YAML itself decides.
|
||||
func parseBlock(lines []line, c *cursor, indent int) (any, *Error) {
|
||||
if c.i >= len(lines) {
|
||||
return nil, nil
|
||||
}
|
||||
first := lines[c.i]
|
||||
if strings.HasPrefix(first.content, "- ") || first.content == "-" {
|
||||
return parseSequence(lines, c, indent)
|
||||
}
|
||||
return parseMapping(lines, c, indent)
|
||||
}
|
||||
|
||||
// dashPrefix is the `-` and the whitespace after it, as `^-\s*` consumes them.
|
||||
func dashPrefix(content string) int {
|
||||
if !strings.HasPrefix(content, "-") {
|
||||
return 0
|
||||
}
|
||||
n := 1
|
||||
for _, r := range content[1:] {
|
||||
if !jsIsSpace(r) {
|
||||
break
|
||||
}
|
||||
n += len(string(r))
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func parseSequence(lines []line, c *cursor, indent int) (any, *Error) {
|
||||
out := []any{}
|
||||
|
||||
for c.i < len(lines) {
|
||||
cur := lines[c.i]
|
||||
if cur.indent < indent {
|
||||
break
|
||||
}
|
||||
if cur.indent > indent {
|
||||
return nil, errIndent(cur.number)
|
||||
}
|
||||
if !strings.HasPrefix(cur.content, "-") {
|
||||
break
|
||||
}
|
||||
|
||||
cut := dashPrefix(cur.content)
|
||||
rest := cur.content[cut:]
|
||||
c.i++
|
||||
|
||||
if rest == "" {
|
||||
// `-` alone: the item is the indented block beneath it.
|
||||
if c.i < len(lines) && lines[c.i].indent > indent {
|
||||
item, err := parseBlock(lines, c, lines[c.i].indent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, item)
|
||||
continue
|
||||
}
|
||||
out = append(out, nil)
|
||||
continue
|
||||
}
|
||||
|
||||
// `- key: value` opens a mapping whose first key sits on the dash. The
|
||||
// remaining keys are indented to where that key started.
|
||||
if m := keyPair.FindStringSubmatch(rest); m != nil {
|
||||
keyIndent := indent + cut
|
||||
item := map[string]any{}
|
||||
key, value := m[1], m[2]
|
||||
|
||||
if value == "" && c.i < len(lines) && lines[c.i].indent > indent {
|
||||
block, err := parseBlock(lines, c, lines[c.i].indent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
item[key] = block
|
||||
} else {
|
||||
item[key] = toScalar(value)
|
||||
}
|
||||
|
||||
for c.i < len(lines) && lines[c.i].indent == keyIndent &&
|
||||
!strings.HasPrefix(lines[c.i].content, "- ") {
|
||||
more, err := parseMapping(lines, c, keyIndent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if m, ok := more.(map[string]any); ok {
|
||||
for k, v := range m {
|
||||
item[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
out = append(out, item)
|
||||
continue
|
||||
}
|
||||
|
||||
out = append(out, toScalar(rest))
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func parseMapping(lines []line, c *cursor, indent int) (any, *Error) {
|
||||
out := map[string]any{}
|
||||
|
||||
for c.i < len(lines) {
|
||||
cur := lines[c.i]
|
||||
if cur.indent < indent {
|
||||
break
|
||||
}
|
||||
if cur.indent > indent {
|
||||
return nil, errIndent(cur.number)
|
||||
}
|
||||
if strings.HasPrefix(cur.content, "- ") {
|
||||
break
|
||||
}
|
||||
|
||||
m := keyPair.FindStringSubmatch(cur.content)
|
||||
if m == nil {
|
||||
return nil, errPair(cur.number, cur.content)
|
||||
}
|
||||
|
||||
key, value := m[1], m[2]
|
||||
c.i++
|
||||
|
||||
if value != "" {
|
||||
out[key] = toScalar(value)
|
||||
continue
|
||||
}
|
||||
|
||||
// An empty value means the value is the block below — or nothing.
|
||||
if c.i < len(lines) && lines[c.i].indent > indent {
|
||||
block, err := parseBlock(lines, c, lines[c.i].indent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out[key] = block
|
||||
continue
|
||||
}
|
||||
out[key] = nil
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ParseYAML reads one document of the subset as plain data.
|
||||
//
|
||||
// Returns map[string]any, []any, or the empty map for an empty document.
|
||||
// Anything it cannot read is an error rather than a guess, so a malformed
|
||||
// definition is reported to its author instead of being registered in a shape
|
||||
// nobody intended.
|
||||
func ParseYAML(source string) (any, error) {
|
||||
lines := readLines(source)
|
||||
if len(lines) == 0 {
|
||||
return map[string]any{}, nil
|
||||
}
|
||||
|
||||
c := &cursor{}
|
||||
value, err := parseBlock(lines, c, lines[0].indent)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if c.i < len(lines) {
|
||||
return nil, errIndent(lines[c.i].number)
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
748
go-api/internal/domain/definitions_schema_test.go
Normal file
748
go-api/internal/domain/definitions_schema_test.go
Normal file
@@ -0,0 +1,748 @@
|
||||
package domain_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// Phase 4C — the shape of the authored-definition tables.
|
||||
//
|
||||
// These test the MIGRATION, not any Go code: there is no agent or skill package
|
||||
// yet, and there deliberately is not one until Phase 4D. What is under test is
|
||||
// whether the database refuses the things it is supposed to refuse.
|
||||
//
|
||||
// An external test package (`domain_test`) rather than `package domain`,
|
||||
// because testutil imports seeder which imports domain — reachable from an
|
||||
// external test binary, an import cycle from an internal one.
|
||||
|
||||
const (
|
||||
agents = "agent_definitions"
|
||||
skills = "skill_definitions"
|
||||
)
|
||||
|
||||
// fixture is a migrated sandbox with one organization and two users.
|
||||
type fixture struct {
|
||||
pool *pgxpool.Pool
|
||||
ctx context.Context
|
||||
orgID string
|
||||
alice string
|
||||
bob string
|
||||
}
|
||||
|
||||
func newFixture(t *testing.T, label string) *fixture {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, label)
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
f := &fixture{pool: pool, ctx: ctx}
|
||||
if err := pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Defs Org','defs-org') RETURNING id::text`).
|
||||
Scan(&f.orgID); err != nil {
|
||||
t.Fatalf("create organization: %v", err)
|
||||
}
|
||||
f.alice = f.newUser(t, "alice@example.test")
|
||||
f.bob = f.newUser(t, "bob@example.test")
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fixture) newUser(t *testing.T, email string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3) RETURNING id::text`,
|
||||
f.orgID, email, email).Scan(&id); err != nil {
|
||||
t.Fatalf("create user %s: %v", email, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
// row is one candidate definition. Any field may be made deliberately wrong.
|
||||
type row struct {
|
||||
table string
|
||||
defID string
|
||||
orgID string
|
||||
visibility string
|
||||
owner *string
|
||||
createdBy *string
|
||||
markdown string
|
||||
status string
|
||||
version *int
|
||||
}
|
||||
|
||||
func (f *fixture) insert(r row) (string, error) {
|
||||
cols := []string{"definition_id", "org_id", "visibility", "owner_user_id", "created_by", "markdown"}
|
||||
vals := []string{"$1::text", "$2::uuid", "$3::text", "$4::uuid", "$5::uuid", "$6::text"}
|
||||
args := []any{r.defID, r.orgID, r.visibility, r.owner, r.createdBy, r.markdown}
|
||||
|
||||
if r.status != "" {
|
||||
cols, vals = append(cols, "status"), append(vals, "$7::text")
|
||||
args = append(args, r.status)
|
||||
}
|
||||
if r.version != nil {
|
||||
cols = append(cols, "version")
|
||||
vals = append(vals, "$"+itoa(len(args)+1)+"::integer")
|
||||
args = append(args, *r.version)
|
||||
}
|
||||
|
||||
var id string
|
||||
err := f.pool.QueryRow(f.ctx,
|
||||
"INSERT INTO "+r.table+" ("+strings.Join(cols, ", ")+") VALUES ("+
|
||||
strings.Join(vals, ", ")+") RETURNING id::text", args...).Scan(&id)
|
||||
return id, err
|
||||
}
|
||||
|
||||
func itoa(n int) string {
|
||||
if n < 10 {
|
||||
return string(rune('0' + n))
|
||||
}
|
||||
return string(rune('0'+n/10)) + string(rune('0'+n%10))
|
||||
}
|
||||
|
||||
// personal and organization build a valid row of each tier, so a test can
|
||||
// change exactly one thing and see whether the database notices.
|
||||
func (f *fixture) personal(table, defID, owner string) row {
|
||||
return row{table: table, defID: defID, orgID: f.orgID, visibility: "personal",
|
||||
owner: &owner, createdBy: &owner, markdown: "---\nid: " + defID + "\n---\n"}
|
||||
}
|
||||
|
||||
func (f *fixture) organization(table, defID, author string) row {
|
||||
return row{table: table, defID: defID, orgID: f.orgID, visibility: "organization",
|
||||
owner: nil, createdBy: &author, markdown: "---\nid: " + defID + "\n---\n"}
|
||||
}
|
||||
|
||||
func mustInsert(t *testing.T, f *fixture, r row) string {
|
||||
t.Helper()
|
||||
id, err := f.insert(r)
|
||||
if err != nil {
|
||||
t.Fatalf("a valid %s row was refused: %v", r.table, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func refused(t *testing.T, f *fixture, r row, wantConstraint, why string) {
|
||||
t.Helper()
|
||||
_, err := f.insert(r)
|
||||
if err == nil {
|
||||
t.Fatalf("%s: the row was ACCEPTED — %s", r.table, why)
|
||||
}
|
||||
if wantConstraint != "" && !strings.Contains(err.Error(), wantConstraint) {
|
||||
t.Errorf("%s: refused by %v, want the %s constraint", r.table, err, wantConstraint)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 1, 2. The ownership invariant ──────────────────────────────────────── */
|
||||
|
||||
func TestVisibilityRequiresMatchingOwnership(t *testing.T) {
|
||||
f := newFixture(t, "defs_ownership")
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
t.Run(table, func(t *testing.T) {
|
||||
// The two valid shapes.
|
||||
mustInsert(t, f, f.personal(table, "valid-personal", f.alice))
|
||||
mustInsert(t, f, f.organization(table, "valid-org", f.alice))
|
||||
|
||||
// 1. A personal definition with no owner belongs to nobody.
|
||||
bad := f.personal(table, "no-owner", f.alice)
|
||||
bad.owner = nil
|
||||
refused(t, f, bad, "visibility_owner",
|
||||
"a personal definition must have an owner")
|
||||
|
||||
// 2. An organization definition with an owner is two answers to
|
||||
// "whose is this", which is one too many.
|
||||
bad = f.organization(table, "with-owner", f.alice)
|
||||
bad.owner = &f.alice
|
||||
refused(t, f, bad, "visibility_owner",
|
||||
"an organization definition must not have an owner")
|
||||
|
||||
// And an unrecognised tier is not a tier.
|
||||
bad = f.personal(table, "bad-tier", f.alice)
|
||||
bad.visibility = "public"
|
||||
refused(t, f, bad, "visibility_check", "`public` is not a visibility")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 3. definition_id format ────────────────────────────────────────────── */
|
||||
|
||||
// The same rule the frontend validator enforces, restated in the database so a
|
||||
// caller that bypasses the application cannot store an id the registry could
|
||||
// never address.
|
||||
func TestDefinitionIDFormat(t *testing.T) {
|
||||
f := newFixture(t, "defs_idformat")
|
||||
|
||||
valid := []string{"a", "board", "krow-workforce-agent", "x1", "a-1-b", "0abc"}
|
||||
invalid := map[string]string{
|
||||
"leading dash": "-board",
|
||||
"upper case": "Board",
|
||||
"underscore": "my_skill",
|
||||
"space": "my skill",
|
||||
"trailing dot": "board.",
|
||||
"empty": "",
|
||||
"slash": "custom/board",
|
||||
"unicode": "bòard",
|
||||
"sql-ish": "a'; DROP TABLE users; --",
|
||||
"newline": "board\nx",
|
||||
}
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
t.Run(table, func(t *testing.T) {
|
||||
for _, id := range valid {
|
||||
if _, err := f.insert(f.personal(table, id, f.alice)); err != nil {
|
||||
t.Errorf("valid id %q was refused: %v", id, err)
|
||||
}
|
||||
}
|
||||
for name, id := range invalid {
|
||||
bad := f.personal(table, id, f.bob)
|
||||
refused(t, f, bad, "definition_id_format", "id "+name+" ("+id+") is not a valid id")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4, 5, 6, 7. Status vocabularies and version ────────────────────────── */
|
||||
|
||||
func TestAgentStatusAndVersion(t *testing.T) {
|
||||
f := newFixture(t, "defs_agentstatus")
|
||||
|
||||
// 4. The three agent statuses, and nothing else.
|
||||
for _, status := range []string{"draft", "published", "archived"} {
|
||||
r := f.personal(agents, "s-"+status, f.alice)
|
||||
r.status = status
|
||||
mustInsert(t, f, r)
|
||||
}
|
||||
for _, status := range []string{"active", "inactive", "live", "DRAFT", ""} {
|
||||
r := f.personal(agents, "bad-status", f.bob)
|
||||
r.status = status
|
||||
if status == "" {
|
||||
continue // an omitted status takes the default; tested below
|
||||
}
|
||||
refused(t, f, r, "status_check", "`"+status+"` is not an agent status")
|
||||
}
|
||||
|
||||
// The default is draft: creating an agent must never publish it.
|
||||
id := mustInsert(t, f, f.personal(agents, "defaulted", f.bob))
|
||||
var status string
|
||||
var version int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT status, version FROM agent_definitions WHERE id = $1::uuid`, id).
|
||||
Scan(&status, &version); err != nil {
|
||||
t.Fatalf("read back: %v", err)
|
||||
}
|
||||
if status != "draft" {
|
||||
t.Errorf("default status = %q, want draft", status)
|
||||
}
|
||||
if version != 1 {
|
||||
t.Errorf("default version = %d, want 1", version)
|
||||
}
|
||||
|
||||
// 6. A version is a whole number of 1 or more.
|
||||
for _, v := range []int{0, -1, -100} {
|
||||
r := f.personal(agents, "bad-version", f.bob)
|
||||
r.version = &v
|
||||
refused(t, f, r, "version_check", "version must be at least 1")
|
||||
}
|
||||
for _, v := range []int{1, 2, 9999} {
|
||||
r := f.personal(agents, "v-ok", f.alice)
|
||||
r.version = &v
|
||||
r.defID = "v-ok-" + itoa(v%100)
|
||||
if _, err := f.insert(r); err != nil {
|
||||
t.Errorf("version %d was refused: %v", v, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillStatusAndNoVersion(t *testing.T) {
|
||||
f := newFixture(t, "defs_skillstatus")
|
||||
|
||||
// 5. The two skill statuses, and nothing else.
|
||||
for _, status := range []string{"active", "inactive"} {
|
||||
r := f.personal(skills, "s-"+status, f.alice)
|
||||
r.status = status
|
||||
mustInsert(t, f, r)
|
||||
}
|
||||
for _, status := range []string{"draft", "published", "archived", "ACTIVE"} {
|
||||
r := f.personal(skills, "bad-status", f.bob)
|
||||
r.status = status
|
||||
refused(t, f, r, "status_check", "`"+status+"` is not a skill status")
|
||||
}
|
||||
|
||||
// The default is active — a skill is on unless somebody turns it off.
|
||||
id := mustInsert(t, f, f.personal(skills, "defaulted", f.bob))
|
||||
var status string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT status FROM skill_definitions WHERE id = $1::uuid`, id).Scan(&status); err != nil {
|
||||
t.Fatalf("read back: %v", err)
|
||||
}
|
||||
if status != "active" {
|
||||
t.Errorf("default status = %q, want active", status)
|
||||
}
|
||||
|
||||
// 7. Skills have NO version. The frontend has no notion of one, so the
|
||||
// column must not exist — inventing it "for symmetry" would create a field
|
||||
// nothing can set and nothing can mean.
|
||||
var exists int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT count(*)::int FROM information_schema.columns
|
||||
WHERE table_schema='public' AND table_name='skill_definitions' AND column_name='version'`).
|
||||
Scan(&exists); err != nil {
|
||||
t.Fatalf("look for a version column: %v", err)
|
||||
}
|
||||
if exists != 0 {
|
||||
t.Error("skill_definitions has a version column; skills have no version concept")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8. Markdown bound ──────────────────────────────────────────────────── */
|
||||
|
||||
func TestMarkdownSizeBound(t *testing.T) {
|
||||
f := newFixture(t, "defs_markdown")
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
t.Run(table, func(t *testing.T) {
|
||||
// The largest definition shipped with the product is 3,156 bytes,
|
||||
// so anything realistic is far inside the bound.
|
||||
ok := f.personal(table, "big-but-fine", f.alice)
|
||||
ok.markdown = strings.Repeat("x", 65536)
|
||||
mustInsert(t, f, ok)
|
||||
|
||||
over := f.personal(table, "too-big", f.bob)
|
||||
over.markdown = strings.Repeat("x", 65537)
|
||||
refused(t, f, over, "markdown_size", "a definition over the size bound")
|
||||
|
||||
empty := f.personal(table, "empty-md", f.bob)
|
||||
empty.markdown = ""
|
||||
refused(t, f, empty, "markdown_size", "an empty definition cannot parse")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The Markdown is stored byte-for-byte. A definition has to survive a round
|
||||
// trip to a .md file on disk, so anything that rewrote it here — trimming,
|
||||
// newline normalisation, unicode folding — would break that.
|
||||
func TestMarkdownIsStoredVerbatim(t *testing.T) {
|
||||
f := newFixture(t, "defs_verbatim")
|
||||
|
||||
// A BOM, CRLF endings, trailing spaces and a tab — exactly the four things
|
||||
// normalizeDefinition exists to tolerate. The database must not "help" by
|
||||
// removing any of them: normalising is the parser's job, on read.
|
||||
source := "\ufeff---\r\nid: verbatim\r\nname: Verbatim\r\n---\r\n\r\n# Verbatim \r\n\ttabbed\n"
|
||||
r := f.personal(agents, "verbatim", f.alice)
|
||||
r.markdown = source
|
||||
id := mustInsert(t, f, r)
|
||||
|
||||
var stored string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT markdown FROM agent_definitions WHERE id = $1::uuid`, id).Scan(&stored); err != nil {
|
||||
t.Fatalf("read back: %v", err)
|
||||
}
|
||||
if stored != source {
|
||||
t.Errorf("the stored Markdown differs from what was written:\n in %q\n out %q", source, stored)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 9, 10, 11, 12. Foreign keys and deletion ───────────────────────────── */
|
||||
|
||||
func TestForeignKeysAndDeleteBehaviour(t *testing.T) {
|
||||
f := newFixture(t, "defs_fk")
|
||||
missing := "00000000-0000-0000-0000-000000000000"
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
t.Run(table+"/rejects unknown references", func(t *testing.T) {
|
||||
// 9. An organization that does not exist.
|
||||
bad := f.personal(table, "bad-org", f.alice)
|
||||
bad.orgID = missing
|
||||
refused(t, f, bad, "org_id_fkey", "org_id must reference a real organization")
|
||||
|
||||
// 10. An owner that does not exist.
|
||||
bad = f.personal(table, "bad-owner", f.alice)
|
||||
bad.owner = &missing
|
||||
refused(t, f, bad, "owner_user_id_fkey", "owner_user_id must reference a real user")
|
||||
|
||||
// 11. An author that does not exist.
|
||||
bad = f.organization(table, "bad-author", f.alice)
|
||||
bad.createdBy = &missing
|
||||
refused(t, f, bad, "created_by_fkey", "created_by must reference a real user")
|
||||
})
|
||||
}
|
||||
|
||||
// 12. Deletion, three behaviours, each different and each deliberate.
|
||||
t.Run("deleting the owner destroys their personal definitions", func(t *testing.T) {
|
||||
carol := f.newUser(t, "carol@example.test")
|
||||
mustInsert(t, f, f.personal(agents, "carols-agent", carol))
|
||||
mustInsert(t, f, f.personal(skills, "carols-skill", carol))
|
||||
|
||||
if _, err := f.pool.Exec(f.ctx, `DELETE FROM users WHERE id = $1::uuid`, carol); err != nil {
|
||||
t.Fatalf("delete the user: %v", err)
|
||||
}
|
||||
for _, table := range []string{agents, skills} {
|
||||
var n int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
"SELECT count(*)::int FROM "+table+" WHERE definition_id LIKE 'carols-%'").Scan(&n); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Errorf("%s: %d personal definitions survive their deleted owner, want 0", table, n)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("deleting the author keeps the organization's definition", func(t *testing.T) {
|
||||
dave := f.newUser(t, "dave@example.test")
|
||||
id := mustInsert(t, f, f.organization(agents, "daves-shared-agent", dave))
|
||||
|
||||
if _, err := f.pool.Exec(f.ctx, `DELETE FROM users WHERE id = $1::uuid`, dave); err != nil {
|
||||
t.Fatalf("delete the user: %v", err)
|
||||
}
|
||||
var author *string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT created_by::text FROM agent_definitions WHERE id = $1::uuid`, id).Scan(&author); err != nil {
|
||||
t.Fatalf("the shared definition did not survive its author: %v", err)
|
||||
}
|
||||
if author != nil {
|
||||
t.Errorf("created_by = %v, want NULL after the author was deleted", *author)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("deleting the organization destroys both tiers", func(t *testing.T) {
|
||||
g := newFixture(t, "defs_orgcascade")
|
||||
mustInsert(t, g, g.personal(agents, "doomed-personal", g.alice))
|
||||
mustInsert(t, g, g.organization(skills, "doomed-shared", g.alice))
|
||||
|
||||
if _, err := g.pool.Exec(g.ctx, `DELETE FROM organizations WHERE id = $1::uuid`, g.orgID); err != nil {
|
||||
t.Fatalf("delete the organization: %v", err)
|
||||
}
|
||||
for _, table := range []string{agents, skills} {
|
||||
var n int
|
||||
if err := g.pool.QueryRow(g.ctx, "SELECT count(*)::int FROM "+table).Scan(&n); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Errorf("%s: %d rows survive their deleted organization, want 0", table, n)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── 13, 14. Uniqueness, per tier ───────────────────────────────────────── */
|
||||
|
||||
func TestUniquenessPerTier(t *testing.T) {
|
||||
f := newFixture(t, "defs_unique")
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
t.Run(table, func(t *testing.T) {
|
||||
// 13. One personal definition per id per owner.
|
||||
mustInsert(t, f, f.personal(table, "board", f.alice))
|
||||
refused(t, f, f.personal(table, "board", f.alice), "personal_key",
|
||||
"one user cannot hold two personal definitions of the same id")
|
||||
|
||||
// A different user may hold their own, which is the whole point of
|
||||
// personal definitions.
|
||||
mustInsert(t, f, f.personal(table, "board", f.bob))
|
||||
|
||||
// 14. One organization definition per id per organization.
|
||||
mustInsert(t, f, f.organization(table, "board", f.alice))
|
||||
refused(t, f, f.organization(table, "board", f.bob), "org_key",
|
||||
"one organization cannot hold two shared definitions of the same id")
|
||||
|
||||
// Personal and organization definitions of the SAME id coexist:
|
||||
// that is shadow-by-id, and it is the reason definition_id is not
|
||||
// globally unique.
|
||||
var personal, shared int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
"SELECT count(*) FILTER (WHERE visibility='personal'), "+
|
||||
"count(*) FILTER (WHERE visibility='organization') "+
|
||||
"FROM "+table+" WHERE definition_id = 'board'").Scan(&personal, &shared); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if personal != 2 || shared != 1 {
|
||||
t.Errorf("board: %d personal + %d shared, want 2 + 1", personal, shared)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// A second organization may hold its own definition of the same id.
|
||||
t.Run("across organizations", func(t *testing.T) {
|
||||
var otherOrg string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Other','other-defs') RETURNING id::text`).
|
||||
Scan(&otherOrg); err != nil {
|
||||
t.Fatalf("create the second organization: %v", err)
|
||||
}
|
||||
var erin string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid,'erin@example.test','Erin')
|
||||
RETURNING id::text`, otherOrg).Scan(&erin); err != nil {
|
||||
t.Fatalf("create a user in the second organization: %v", err)
|
||||
}
|
||||
r := f.organization(agents, "board", erin)
|
||||
r.orgID = otherOrg
|
||||
mustInsert(t, f, r)
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Schema shape ───────────────────────────────────────────────────────── */
|
||||
|
||||
func TestDefinitionTablesShape(t *testing.T) {
|
||||
f := newFixture(t, "defs_shape")
|
||||
|
||||
shared := map[string]string{
|
||||
"id": "uuid",
|
||||
"definition_id": "text",
|
||||
"org_id": "uuid",
|
||||
"visibility": "text",
|
||||
"owner_user_id": "uuid",
|
||||
"created_by": "text-or-uuid", // placeholder, replaced below
|
||||
"markdown": "text",
|
||||
"status": "text",
|
||||
"name": "text",
|
||||
"description": "text",
|
||||
"pages": "ARRAY",
|
||||
"created_date": "timestamp with time zone",
|
||||
"updated_date": "timestamp with time zone",
|
||||
}
|
||||
shared["created_by"] = "uuid"
|
||||
|
||||
nullable := map[string]bool{"owner_user_id": true, "created_by": true}
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
want := map[string]string{}
|
||||
for k, v := range shared {
|
||||
want[k] = v
|
||||
}
|
||||
if table == agents {
|
||||
want["version"] = "integer"
|
||||
}
|
||||
|
||||
t.Run(table, func(t *testing.T) {
|
||||
rows, err := f.pool.Query(f.ctx,
|
||||
`SELECT column_name, data_type, is_nullable
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema='public' AND table_name=$1`, table)
|
||||
if err != nil {
|
||||
t.Fatalf("read columns: %v", err)
|
||||
}
|
||||
got := map[string]string{}
|
||||
for rows.Next() {
|
||||
var name, kind, isNullable string
|
||||
if err := rows.Scan(&name, &kind, &isNullable); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
got[name] = kind
|
||||
if (isNullable == "YES") != nullable[name] {
|
||||
t.Errorf("%s.%s is_nullable=%s, want nullable=%v", table, name, isNullable, nullable[name])
|
||||
}
|
||||
}
|
||||
rows.Close()
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("read columns: %v", err)
|
||||
}
|
||||
|
||||
for name, kind := range want {
|
||||
if got[name] == "" {
|
||||
t.Errorf("%s.%s is missing", table, name)
|
||||
} else if got[name] != kind {
|
||||
t.Errorf("%s.%s is %s, want %s", table, name, got[name], kind)
|
||||
}
|
||||
}
|
||||
for name := range got {
|
||||
if _, expected := want[name]; !expected {
|
||||
t.Errorf("%s has an unexpected column %q", table, name)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefinitionIndexes(t *testing.T) {
|
||||
f := newFixture(t, "defs_indexes")
|
||||
|
||||
want := map[string][]string{
|
||||
agents: {
|
||||
"agent_definitions_pkey",
|
||||
"agent_definitions_personal_key",
|
||||
"agent_definitions_org_key",
|
||||
"agent_definitions_org_visibility_idx",
|
||||
"agent_definitions_owner_idx",
|
||||
"agent_definitions_published_idx",
|
||||
},
|
||||
skills: {
|
||||
"skill_definitions_pkey",
|
||||
"skill_definitions_personal_key",
|
||||
"skill_definitions_org_key",
|
||||
"skill_definitions_org_visibility_idx",
|
||||
"skill_definitions_owner_idx",
|
||||
"skill_definitions_active_idx",
|
||||
},
|
||||
}
|
||||
|
||||
for table, names := range want {
|
||||
rows, err := f.pool.Query(f.ctx,
|
||||
`SELECT indexname, indexdef FROM pg_indexes WHERE schemaname='public' AND tablename=$1`, table)
|
||||
if err != nil {
|
||||
t.Fatalf("list indexes: %v", err)
|
||||
}
|
||||
got := map[string]string{}
|
||||
for rows.Next() {
|
||||
var name, def string
|
||||
if err := rows.Scan(&name, &def); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
got[name] = def
|
||||
}
|
||||
rows.Close()
|
||||
|
||||
for _, name := range names {
|
||||
if got[name] == "" {
|
||||
t.Errorf("%s: index %s is missing", table, name)
|
||||
}
|
||||
}
|
||||
// The two uniqueness indexes must be partial and unique, or they mean
|
||||
// something other than what they are named.
|
||||
for _, name := range []string{table + "_personal_key", table + "_org_key"} {
|
||||
def := got[name]
|
||||
if !strings.Contains(def, "UNIQUE") {
|
||||
t.Errorf("%s is not UNIQUE: %s", name, def)
|
||||
}
|
||||
if !strings.Contains(def, "WHERE") {
|
||||
t.Errorf("%s is not partial: %s", name, def)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// created_by is deliberately unindexed: attribution only, no listing is
|
||||
// keyed by it, and its SET NULL scan happens only when a user is deleted.
|
||||
for _, table := range []string{agents, skills} {
|
||||
var n int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT count(*)::int FROM pg_indexes
|
||||
WHERE schemaname='public' AND tablename=$1 AND indexdef LIKE '%(created_by)%'`,
|
||||
table).Scan(&n); err != nil {
|
||||
t.Fatalf("look for a created_by index: %v", err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Errorf("%s has an index on created_by; it was deliberately omitted", table)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Reversibility ──────────────────────────────────────────────────────── */
|
||||
|
||||
func TestMigration000005IsReversible(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, "defs_reversible")
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
const up = "000005_agent_skill_definitions.up.sql"
|
||||
const down = "000005_agent_skill_definitions.down.sql"
|
||||
|
||||
exists := func(name string) bool {
|
||||
var reg *string
|
||||
if err := pool.QueryRow(ctx, `SELECT to_regclass('public.' || $1)::text`, name).Scan(®); err != nil {
|
||||
t.Fatalf("to_regclass(%s): %v", name, err)
|
||||
}
|
||||
return reg != nil
|
||||
}
|
||||
|
||||
for _, table := range []string{agents, skills} {
|
||||
if !exists(table) {
|
||||
t.Fatalf("%s does not exist before the rollback", table)
|
||||
}
|
||||
}
|
||||
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, down); err != nil {
|
||||
t.Fatalf("apply %s: %v", down, err)
|
||||
}
|
||||
for _, table := range []string{agents, skills} {
|
||||
if exists(table) {
|
||||
t.Errorf("%s survived the rollback", table)
|
||||
}
|
||||
}
|
||||
|
||||
// The rollback must reach nothing that predates it.
|
||||
for _, table := range []string{"users", "organizations", "sessions", "user_preferences", "job_postings"} {
|
||||
if !exists(table) {
|
||||
t.Fatalf("the rollback dropped %s, which 000005 did not create", table)
|
||||
}
|
||||
}
|
||||
// And no enum type was created, so none can be left behind.
|
||||
var leftover int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_type t JOIN pg_namespace n ON n.oid = t.typnamespace
|
||||
WHERE n.nspname='public' AND t.typtype='e'
|
||||
AND t.typname IN ('definition_visibility','agent_status','skill_status')`).Scan(&leftover); err != nil {
|
||||
t.Fatalf("look for leftover types: %v", err)
|
||||
}
|
||||
if leftover != 0 {
|
||||
t.Errorf("%d enum types left behind by the rollback", leftover)
|
||||
}
|
||||
|
||||
// Re-applying restores exactly what was removed.
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, up); err != nil {
|
||||
t.Fatalf("re-apply %s: %v", up, err)
|
||||
}
|
||||
for _, table := range []string{agents, skills} {
|
||||
if !exists(table) {
|
||||
t.Errorf("%s did not come back", table)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every migration still has a matching down file, and 000005 is the newest.
|
||||
func TestMigrationPairsIncluding000005(t *testing.T) {
|
||||
ups := testutil.MigrationFiles(t, ".up.sql")
|
||||
downs := testutil.MigrationFiles(t, ".down.sql")
|
||||
if len(ups) != len(downs) {
|
||||
t.Fatalf("%d up and %d down migrations", len(ups), len(downs))
|
||||
}
|
||||
for i, up := range ups {
|
||||
want := strings.TrimSuffix(up, ".up.sql") + ".down.sql"
|
||||
if downs[i] != want {
|
||||
t.Errorf("%s has no matching down migration (found %s)", up, downs[i])
|
||||
}
|
||||
}
|
||||
if len(ups) != 5 {
|
||||
t.Errorf("%d migrations, want 5", len(ups))
|
||||
}
|
||||
if ups[4] != "000005_agent_skill_definitions.up.sql" {
|
||||
t.Errorf("the last migration is %s", ups[4])
|
||||
}
|
||||
}
|
||||
|
||||
// 000005 creates exactly two tables and nothing else. The Phase 4B decision was
|
||||
// explicit about which tables must NOT appear; this is that decision, asserted.
|
||||
func TestMigrationAddsExactlyTwoTables(t *testing.T) {
|
||||
f := newFixture(t, "defs_tablecount")
|
||||
|
||||
var n int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT count(*)::int FROM information_schema.tables
|
||||
WHERE table_schema='public' AND table_type='BASE TABLE'`).Scan(&n); err != nil {
|
||||
t.Fatalf("count tables: %v", err)
|
||||
}
|
||||
// 17 from 000001 + sessions from 000004 + the two here. schema_migrations is
|
||||
// golang-migrate's and is absent when the files are applied directly.
|
||||
if n != 20 {
|
||||
t.Errorf("%d base tables after every migration, want 20", n)
|
||||
}
|
||||
|
||||
for _, forbidden := range []string{
|
||||
"definition_versions", "definition_permissions", "agent_skills",
|
||||
"agent_subagents", "agent_knowledge", "conversations",
|
||||
"conversation_messages", "conversation_feedback",
|
||||
} {
|
||||
var reg *string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT to_regclass('public.' || $1)::text`, forbidden).Scan(®); err != nil {
|
||||
t.Fatalf("to_regclass: %v", err)
|
||||
}
|
||||
if reg != nil {
|
||||
t.Errorf("table %s exists; Phase 4B deferred or rejected it", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
76
go-api/internal/domain/errors.go
Normal file
76
go-api/internal/domain/errors.go
Normal file
@@ -0,0 +1,76 @@
|
||||
package domain
|
||||
|
||||
import "fmt"
|
||||
|
||||
// Error is an API-level failure carrying the contract's error code.
|
||||
// See api-contract.md §5.
|
||||
type Error struct {
|
||||
Code string
|
||||
Message string
|
||||
Details map[string]string
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *Error) Error() string { return e.Message }
|
||||
func (e *Error) Unwrap() error { return e.cause }
|
||||
|
||||
// NotFound reproduces store.js's thrown message verbatim: "<Entity> <id> not
|
||||
// found", using the frontend's entity name rather than the table name.
|
||||
func NotFound(entity, id string) *Error {
|
||||
return &Error{Code: "not_found", Message: fmt.Sprintf("%s %s not found", entity, id)}
|
||||
}
|
||||
|
||||
func Invalid(msg string) *Error {
|
||||
return &Error{Code: "invalid_query", Message: msg}
|
||||
}
|
||||
|
||||
func Validation(msg string, details map[string]string) *Error {
|
||||
if details == nil {
|
||||
details = map[string]string{}
|
||||
}
|
||||
return &Error{Code: "validation_failed", Message: msg, Details: details}
|
||||
}
|
||||
|
||||
func Conflict(msg string) *Error {
|
||||
return &Error{Code: "conflict", Message: msg}
|
||||
}
|
||||
|
||||
// Unauthenticated is every "you are not signed in" answer: no cookie, an
|
||||
// unknown token, an expired session, a suspended user, a wrong password, an
|
||||
// email that does not exist.
|
||||
//
|
||||
// One constructor for all of them, deliberately. The distinctions matter in the
|
||||
// server log and must not reach the client: which of those it was tells an
|
||||
// attacker whether an address is registered, whether an account is suspended,
|
||||
// or whether a guessed token was ever real.
|
||||
func Unauthenticated() *Error {
|
||||
return &Error{Code: "unauthorized", Message: "authentication required"}
|
||||
}
|
||||
|
||||
// Forbidden is the answer to an authenticated caller whose role does not permit
|
||||
// the operation.
|
||||
//
|
||||
// Distinct from Unauthenticated: 401 means "I do not know who you are", 403
|
||||
// means "I know exactly who you are and the answer is still no". Conflating
|
||||
// them makes a client retry a login that will not help.
|
||||
//
|
||||
// The message names neither the role the caller has nor the roles that would
|
||||
// have worked. That is not secrecy for its own sake — it is that an endpoint
|
||||
// which answers "employers only" to a talent user is an endpoint that maps the
|
||||
// organization's privilege structure for anyone who asks.
|
||||
//
|
||||
// Note what does NOT come through here: a row belonging to another
|
||||
// organization, or to another person, is not forbidden — it is absent. Those
|
||||
// answer 404 by way of a SQL predicate, so existence never leaks.
|
||||
func Forbidden() *Error {
|
||||
return &Error{Code: "forbidden", Message: "you do not have access to this operation"}
|
||||
}
|
||||
|
||||
// RateLimited is the answer to too many failed sign-in attempts.
|
||||
func RateLimited(msg string) *Error {
|
||||
return &Error{Code: "rate_limited", Message: msg}
|
||||
}
|
||||
|
||||
func Internal(err error) *Error {
|
||||
return &Error{Code: "internal", Message: "internal error", cause: err}
|
||||
}
|
||||
330
go-api/internal/domain/policy.go
Normal file
330
go-api/internal/domain/policy.go
Normal file
@@ -0,0 +1,330 @@
|
||||
package domain
|
||||
|
||||
// Authorization policy: who may perform which operation on which resource, and
|
||||
// which rows they may see when they get there.
|
||||
//
|
||||
// This file is hand-written and `resources_gen.go` is generated, which is the
|
||||
// whole reason they are separate. Regenerating the descriptors from the live
|
||||
// schema must never silently drop an access rule, and a column appearing in the
|
||||
// database must never grant anybody anything by accident.
|
||||
//
|
||||
// Three properties hold here by construction:
|
||||
//
|
||||
// - DENY BY DEFAULT. A resource with no policy permits nothing, to anyone. A
|
||||
// resource added to the schema tomorrow is unreachable until somebody
|
||||
// writes down who may reach it. TestEveryResourceHasAPolicy makes the
|
||||
// omission loud rather than silent.
|
||||
// - ROLE IS users.role, ALWAYS. Never account_type — which the user can
|
||||
// change on themselves through PATCH /me — and never anything read from a
|
||||
// request body, a header or the browser.
|
||||
// - OWNERSHIP IS A SQL PREDICATE, NOT A FILTER. TalentScope describes a WHERE
|
||||
// clause the repository adds beside the organization scope. Rows a talent
|
||||
// user may not see are never fetched, so they cannot leak through a count,
|
||||
// a total or a bug in a later loop.
|
||||
//
|
||||
// Authorization is checked in the handler, before any query runs, and answers
|
||||
// 403. Organization and ownership are predicates, so a row outside them is
|
||||
// simply absent and answers 404 — the caller cannot tell "exists but not yours"
|
||||
// from "does not exist", which is the point.
|
||||
|
||||
// Role is the authorization authority. It mirrors the users_role_check
|
||||
// constraint in migration 000001 and there are deliberately no others.
|
||||
type Role string
|
||||
|
||||
const (
|
||||
RoleAdmin Role = "admin"
|
||||
RoleEmployer Role = "employer"
|
||||
RoleTalent Role = "talent"
|
||||
)
|
||||
|
||||
// ParseRole converts a stored users.role into a Role, reporting whether it is
|
||||
// one this API recognises. An unrecognised value authorizes nothing.
|
||||
func ParseRole(s string) (Role, bool) {
|
||||
switch Role(s) {
|
||||
case RoleAdmin:
|
||||
return RoleAdmin, true
|
||||
case RoleEmployer:
|
||||
return RoleEmployer, true
|
||||
case RoleTalent:
|
||||
return RoleTalent, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
/* ── Row visibility ─────────────────────────────────────────────────────── */
|
||||
|
||||
// ScopeKind is how a resource decides which rows a talent user may see.
|
||||
type ScopeKind int
|
||||
|
||||
const (
|
||||
// ScopeNone: no extra predicate. Every row in the organization is visible.
|
||||
ScopeNone ScopeKind = iota
|
||||
// ScopeUserID: Column = the authenticated user's id.
|
||||
ScopeUserID
|
||||
// ScopeEmail: Column = the authenticated user's email.
|
||||
//
|
||||
// Used where the schema ties a row to a person by email string rather than
|
||||
// by a foreign key — assignments, evidence, shift records, applications,
|
||||
// activity. Those columns have no FK (see the Phase 3D audit, F-05), so the
|
||||
// write path is what makes this trustworthy: a talent caller never supplies
|
||||
// the value, it is derived from the session. See Derived.
|
||||
ScopeEmail
|
||||
// ScopeOwnApplications: the row references a job application belonging to
|
||||
// the authenticated user. Ownership by reference rather than by column —
|
||||
// an AI interview names an application, and the application names a person.
|
||||
ScopeOwnApplications
|
||||
// ScopeActivePostings: visibility rather than ownership. A talent user sees
|
||||
// the postings they could apply to, not the organization's drafts, paused
|
||||
// roles or closed history.
|
||||
ScopeActivePostings
|
||||
)
|
||||
|
||||
// Scope is the predicate applied to a talent caller's rows.
|
||||
type Scope struct {
|
||||
Kind ScopeKind
|
||||
// Column is the column carrying the owner, for ScopeUserID and ScopeEmail.
|
||||
// For ScopeOwnApplications it is the column referencing the application.
|
||||
// For ScopeActivePostings it is the status column.
|
||||
//
|
||||
// Required by every kind except ScopeNone: a scope naming a column the
|
||||
// resource does not have matches no rows at all, which is the safe
|
||||
// direction to fail but is still a bug worth noticing.
|
||||
Column string
|
||||
}
|
||||
|
||||
/* ── Server-owned values ────────────────────────────────────────────────── */
|
||||
|
||||
// DeriveSource names which fact about the caller fills a column.
|
||||
type DeriveSource int
|
||||
|
||||
const (
|
||||
DeriveUserID DeriveSource = iota
|
||||
DeriveEmail
|
||||
DeriveFullName
|
||||
DeriveAccountType
|
||||
)
|
||||
|
||||
// Derived is a column the server fills in on insert from the session.
|
||||
//
|
||||
// Every column named here is also ReadOnly in the descriptors, so a value in a
|
||||
// request body is dropped before it reaches SQL. This is the other half: the
|
||||
// column still has to be filled, and the only acceptable source is the
|
||||
// authenticated identity.
|
||||
type Derived struct {
|
||||
Column string
|
||||
Source DeriveSource
|
||||
// TalentOnly restricts the derivation to talent callers.
|
||||
//
|
||||
// It exists because two different questions wear the same shape. `created_by`
|
||||
// and the user_activity columns record WHO ACTED, so they are the session
|
||||
// user whoever that is. `worker_profiles.user_id`, `job_applications.email`
|
||||
// and `evidence.worker_email` record WHO THE ROW IS ABOUT — and when an
|
||||
// admin creates a candidate's profile or logs an application on their
|
||||
// behalf, the subject is emphatically not the admin. Deriving those
|
||||
// unconditionally would quietly file every candidate's record under the
|
||||
// operator who typed it in.
|
||||
TalentOnly bool
|
||||
}
|
||||
|
||||
/* ── Policy ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Policy is one resource's access rules.
|
||||
//
|
||||
// A nil Policy denies everything. An empty role list for an operation denies
|
||||
// that operation to everyone, which is how an operation the resource does not
|
||||
// support is expressed.
|
||||
type Policy struct {
|
||||
List []Role
|
||||
Get []Role
|
||||
Create []Role
|
||||
Update []Role
|
||||
Delete []Role
|
||||
|
||||
// TalentScope narrows which rows a talent caller may read or write. It is
|
||||
// applied to talent and to nobody else: admin and employer see the whole
|
||||
// organization, which is what an operator console is for.
|
||||
TalentScope Scope
|
||||
|
||||
// Derived fills server-owned columns on insert.
|
||||
Derived []Derived
|
||||
}
|
||||
|
||||
// Allows reports whether a role may perform an operation.
|
||||
//
|
||||
// The zero answer is no: a nil policy, an unknown operation and an unlisted
|
||||
// role all deny.
|
||||
func (p *Policy) Allows(op Op, role Role) bool {
|
||||
if p == nil {
|
||||
return false
|
||||
}
|
||||
for _, r := range p.rolesFor(op) {
|
||||
if r == role {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (p *Policy) rolesFor(op Op) []Role {
|
||||
switch op {
|
||||
case OpList:
|
||||
return p.List
|
||||
case OpGet:
|
||||
return p.Get
|
||||
case OpCreate:
|
||||
return p.Create
|
||||
case OpUpdate:
|
||||
return p.Update
|
||||
case OpDelete:
|
||||
return p.Delete
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ScopeFor returns the row predicate that applies to a role. Only talent is
|
||||
// scoped; every other role sees the organization.
|
||||
func (p *Policy) ScopeFor(role Role) Scope {
|
||||
if p == nil || role != RoleTalent {
|
||||
return Scope{}
|
||||
}
|
||||
return p.TalentScope
|
||||
}
|
||||
|
||||
/* ── The table ──────────────────────────────────────────────────────────── */
|
||||
|
||||
var (
|
||||
everyone = []Role{RoleAdmin, RoleEmployer, RoleTalent}
|
||||
// operators are the roles that run the organization's hiring and workforce:
|
||||
// they see and act on the whole tenant. Talent is not one of them.
|
||||
operators = []Role{RoleAdmin, RoleEmployer}
|
||||
adminOnly = []Role{RoleAdmin}
|
||||
)
|
||||
|
||||
// policies is the authorization contract, keyed by URL path.
|
||||
//
|
||||
// Read this table as the answer to "who may call this, and which rows do they
|
||||
// get". It is the only place those two questions are answered.
|
||||
var policies = map[string]*Policy{
|
||||
|
||||
// Postings are the organization's shop window. Operators author them;
|
||||
// talent sees the ones that are open, and nothing else — not the drafts,
|
||||
// not the paused roles, not the closed history.
|
||||
"job-postings": {
|
||||
List: everyone, Get: everyone,
|
||||
Create: operators, Update: operators,
|
||||
TalentScope: Scope{Kind: ScopeActivePostings, Column: "status"},
|
||||
Derived: []Derived{{Column: "created_by", Source: DeriveUserID}},
|
||||
},
|
||||
|
||||
// An application is written by a person about themselves. Talent may file
|
||||
// one and read their own; only operators may move it through the funnel or
|
||||
// remove it. A talent caller never supplies the email — it is the session's,
|
||||
// which is what makes the read predicate below mean anything.
|
||||
"job-applications": {
|
||||
List: everyone, Create: everyone,
|
||||
Update: operators, Delete: operators,
|
||||
TalentScope: Scope{Kind: ScopeEmail, Column: "email"},
|
||||
Derived: []Derived{{Column: "email", Source: DeriveEmail, TalentOnly: true}},
|
||||
},
|
||||
|
||||
// An interview belongs to an application, and the application belongs to a
|
||||
// person. Talent reads their own and may sit one for an application of
|
||||
// theirs; the insert guard in the repository enforces the second half.
|
||||
"ai-interviews": {
|
||||
List: everyone, Create: everyone,
|
||||
TalentScope: Scope{Kind: ScopeOwnApplications, Column: "application_id"},
|
||||
},
|
||||
|
||||
// The employment record of the organization's workforce. Operators only:
|
||||
// it carries endorsements, review dates and reviewer names, which are the
|
||||
// organization's assessment of a person rather than the person's own data.
|
||||
"staff": {
|
||||
List: operators, Create: operators, Update: operators,
|
||||
},
|
||||
|
||||
// A worker's own profile: contact details, address, salary expectations,
|
||||
// personality assessment. Talent may read and maintain theirs and no other.
|
||||
// user_id is server-owned, so a talent caller cannot claim someone else's
|
||||
// profile by naming them, and cannot hand theirs away.
|
||||
"worker-profiles": {
|
||||
List: everyone, Create: everyone, Update: everyone,
|
||||
TalentScope: Scope{Kind: ScopeUserID, Column: "user_id"},
|
||||
Derived: []Derived{{Column: "user_id", Source: DeriveUserID, TalentOnly: true}},
|
||||
},
|
||||
|
||||
// Who is on which position. Operators allocate; talent reads their own
|
||||
// roster and cannot create one — being assigned to work is not a thing you
|
||||
// do to yourself.
|
||||
"assignments": {
|
||||
List: everyone, Create: operators,
|
||||
TalentScope: Scope{Kind: ScopeEmail, Column: "worker_email"},
|
||||
},
|
||||
|
||||
// Attendance. Read-only for everyone over the API — the seeder owns these
|
||||
// rows — and talent sees only their own shifts.
|
||||
"shift-records": {
|
||||
List: everyone,
|
||||
TalentScope: Scope{Kind: ScopeEmail, Column: "worker_email"},
|
||||
},
|
||||
|
||||
// The training library. Everyone learns from it; only admin authors it.
|
||||
// Employer is excluded from authoring deliberately: courses with a NULL
|
||||
// org_id are the shared platform library, visible to every tenant, so a
|
||||
// write here can reach beyond the writer's own organization.
|
||||
"courses": {
|
||||
List: everyone, Get: everyone,
|
||||
Create: adminOnly, Update: adminOnly,
|
||||
},
|
||||
|
||||
"learning-paths": {List: everyone},
|
||||
|
||||
// Organization taxonomy: the role names positions are filed under.
|
||||
"role-categories": {
|
||||
List: everyone, Create: operators,
|
||||
},
|
||||
|
||||
// Organization taxonomy: the certifications positions may require.
|
||||
// Deleting one changes what every existing posting means, so it is admin's.
|
||||
"certifications": {
|
||||
List: everyone, Create: operators, Delete: adminOnly,
|
||||
},
|
||||
|
||||
// The audit log. Append-only by schema (no update, no delete). Anyone may
|
||||
// write an entry about themselves — and only about themselves: all four
|
||||
// identity columns are server-derived, so an entry cannot be attributed to
|
||||
// someone else. Operators read the organization's log; talent reads theirs.
|
||||
"user-activity": {
|
||||
List: everyone, Create: everyone,
|
||||
TalentScope: Scope{Kind: ScopeEmail, Column: "user_email"},
|
||||
Derived: []Derived{
|
||||
{Column: "user_id", Source: DeriveUserID},
|
||||
{Column: "user_email", Source: DeriveEmail},
|
||||
{Column: "user_name", Source: DeriveFullName},
|
||||
{Column: "account_type", Source: DeriveAccountType},
|
||||
},
|
||||
},
|
||||
|
||||
// Proof of work: a worker submits it, the organization verifies it.
|
||||
// Talent may submit their own and read it back; the verdict is an operator
|
||||
// judgement, so talent cannot PATCH.
|
||||
"evidence": {
|
||||
List: everyone, Create: everyone, Update: operators,
|
||||
TalentScope: Scope{Kind: ScopeEmail, Column: "worker_email"},
|
||||
Derived: []Derived{{Column: "worker_email", Source: DeriveEmail, TalentOnly: true}},
|
||||
},
|
||||
|
||||
// Badges have no endpoints at all (Ops: 0 — the frontend's Badge.list call
|
||||
// has 404ed since Phase 2C). The empty policy is written out rather than
|
||||
// omitted so that the resource is deliberately closed rather than merely
|
||||
// forgotten, and so TestEveryResourceHasAPolicy passes honestly.
|
||||
"badges": {},
|
||||
}
|
||||
|
||||
// init attaches the policies to the descriptors.
|
||||
//
|
||||
// A resource with no entry keeps a nil Policy and therefore permits nothing.
|
||||
func init() {
|
||||
for _, r := range AllResources {
|
||||
r.Policy = policies[r.Path]
|
||||
}
|
||||
}
|
||||
191
go-api/internal/domain/policy_test.go
Normal file
191
go-api/internal/domain/policy_test.go
Normal file
@@ -0,0 +1,191 @@
|
||||
package domain
|
||||
|
||||
import "testing"
|
||||
|
||||
// Invariants of the policy table itself. No database: these catch the mistakes
|
||||
// that would otherwise only show up as a missing 403 in an integration test, or
|
||||
// not at all.
|
||||
|
||||
// Every resource must say who may reach it. A resource added to the schema and
|
||||
// left out of policies.go is unreachable — which is the safe direction, and
|
||||
// still a mistake worth failing on rather than discovering in production.
|
||||
func TestEveryResourceHasAPolicy(t *testing.T) {
|
||||
for _, r := range AllResources {
|
||||
if r.Policy == nil {
|
||||
t.Errorf("resource %q (%s) has no policy: it permits nothing, which is safe but almost certainly unintended",
|
||||
r.Name, r.Path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A nil policy denies everything. This is the property the test above relies on
|
||||
// being true, so it is asserted rather than assumed.
|
||||
func TestNilPolicyDeniesEverything(t *testing.T) {
|
||||
var p *Policy
|
||||
for _, op := range []Op{OpList, OpGet, OpCreate, OpUpdate, OpDelete} {
|
||||
for _, role := range []Role{RoleAdmin, RoleEmployer, RoleTalent} {
|
||||
if p.Allows(op, role) {
|
||||
t.Errorf("a nil policy allowed op %d for %s", op, role)
|
||||
}
|
||||
}
|
||||
}
|
||||
if got := p.ScopeFor(RoleTalent); got.Kind != ScopeNone {
|
||||
t.Error("a nil policy returned a scope")
|
||||
}
|
||||
}
|
||||
|
||||
// A policy must not grant an operation the resource does not expose. Such a
|
||||
// grant is dead — no route is registered — but it reads as permission and would
|
||||
// become real the moment the operation is added.
|
||||
func TestPolicyGrantsNothingWithoutARoute(t *testing.T) {
|
||||
ops := []struct {
|
||||
op Op
|
||||
name string
|
||||
}{
|
||||
{OpList, "List"}, {OpGet, "Get"}, {OpCreate, "Create"},
|
||||
{OpUpdate, "Update"}, {OpDelete, "Delete"},
|
||||
}
|
||||
for _, r := range AllResources {
|
||||
if r.Policy == nil {
|
||||
continue
|
||||
}
|
||||
for _, o := range ops {
|
||||
granted := len(r.Policy.rolesFor(o.op)) > 0
|
||||
if granted && !r.Supports(o.op) {
|
||||
t.Errorf("%s: policy grants %s but the resource has no such route", r.Path, o.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// An unrecognised role authorizes nothing, whatever the policy says.
|
||||
func TestUnknownRoleIsDenied(t *testing.T) {
|
||||
if _, ok := ParseRole("superuser"); ok {
|
||||
t.Fatal("ParseRole accepted a role outside the users_role_check constraint")
|
||||
}
|
||||
if _, ok := ParseRole(""); ok {
|
||||
t.Fatal("ParseRole accepted an empty role")
|
||||
}
|
||||
for _, r := range AllResources {
|
||||
if r.Policy.Allows(OpList, Role("superuser")) {
|
||||
t.Errorf("%s allows an unknown role", r.Path)
|
||||
}
|
||||
}
|
||||
// The three real ones parse.
|
||||
for _, want := range []Role{RoleAdmin, RoleEmployer, RoleTalent} {
|
||||
if got, ok := ParseRole(string(want)); !ok || got != want {
|
||||
t.Errorf("ParseRole(%q) = %q, %v", want, got, ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every column the server derives must also be ReadOnly, or a request body
|
||||
// could still set it on a path the derivation does not cover.
|
||||
func TestDerivedColumnsAreReadOnlyOrTalentScoped(t *testing.T) {
|
||||
for _, r := range AllResources {
|
||||
if r.Policy == nil {
|
||||
continue
|
||||
}
|
||||
for _, d := range r.Policy.Derived {
|
||||
col, ok := r.Column(d.Column)
|
||||
if !ok {
|
||||
t.Errorf("%s: derives %q, which is not a column", r.Path, d.Column)
|
||||
continue
|
||||
}
|
||||
// A TalentOnly derivation intentionally leaves the column writable
|
||||
// for operators — an admin filing a candidate's application must be
|
||||
// able to say whose it is. The unconditional ones must be sealed.
|
||||
if !d.TalentOnly && !col.ReadOnly {
|
||||
t.Errorf("%s.%s is derived unconditionally but is not ReadOnly: a request body could still set it",
|
||||
r.Path, d.Column)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The six columns Phase 3D closed. Named explicitly, so that regenerating the
|
||||
// descriptors without the SERVER_OWNED map in gen_resources.py fails loudly
|
||||
// rather than silently reopening the holes.
|
||||
func TestServerOwnedColumnsAreReadOnly(t *testing.T) {
|
||||
sealed := map[string][]string{
|
||||
"worker-profiles": {"user_id"},
|
||||
"user-activity": {"user_id", "user_email", "user_name", "account_type"},
|
||||
"job-postings": {"created_by"},
|
||||
}
|
||||
for path, cols := range sealed {
|
||||
res, ok := ResourceByPath[path]
|
||||
if !ok {
|
||||
t.Fatalf("resource %s is missing", path)
|
||||
}
|
||||
for _, name := range cols {
|
||||
col, ok := res.Column(name)
|
||||
if !ok {
|
||||
t.Errorf("%s has no column %s", path, name)
|
||||
continue
|
||||
}
|
||||
if !col.ReadOnly {
|
||||
t.Errorf("%s.%s is not ReadOnly — a client could supply it", path, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// And org_id everywhere, which predates Phase 3D and must stay that way.
|
||||
for _, r := range AllResources {
|
||||
if col, ok := r.Column("org_id"); ok && !col.ReadOnly {
|
||||
t.Errorf("%s.org_id is not ReadOnly", r.Path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Talent is the only scoped role. If a scope ever applied to an operator the
|
||||
// admin console would start losing rows, which is a failure mode worth pinning.
|
||||
func TestOnlyTalentIsRowScoped(t *testing.T) {
|
||||
for _, r := range AllResources {
|
||||
for _, role := range []Role{RoleAdmin, RoleEmployer} {
|
||||
if got := r.Policy.ScopeFor(role); got.Kind != ScopeNone {
|
||||
t.Errorf("%s scopes rows for %s: operators see the whole organization", r.Path, role)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every talent scope must name a column the resource actually has.
|
||||
func TestTalentScopesNameRealColumns(t *testing.T) {
|
||||
for _, r := range AllResources {
|
||||
scope := r.Policy.ScopeFor(RoleTalent)
|
||||
if scope.Kind == ScopeNone {
|
||||
continue
|
||||
}
|
||||
if scope.Column == "" {
|
||||
t.Errorf("%s has a talent scope with no column", r.Path)
|
||||
continue
|
||||
}
|
||||
if _, ok := r.Column(scope.Column); !ok {
|
||||
t.Errorf("%s scopes on %q, which is not one of its columns", r.Path, scope.Column)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Talent must not reach an operator resource by having a scope but no grant,
|
||||
// or a grant but no scope where one is required. This pins the shape of the
|
||||
// contract: wherever talent may list a resource that also holds other people's
|
||||
// rows, a scope must narrow it.
|
||||
func TestTalentGrantsHaveScopesWhereRowsAreShared(t *testing.T) {
|
||||
// Resources whose rows are the organization's rather than any one person's:
|
||||
// a talent grant here is deliberate and needs no ownership predicate.
|
||||
shared := map[string]bool{
|
||||
"courses": true, "learning-paths": true,
|
||||
"role-categories": true, "certifications": true,
|
||||
}
|
||||
for _, r := range AllResources {
|
||||
if !r.Policy.Allows(OpList, RoleTalent) {
|
||||
continue
|
||||
}
|
||||
if shared[r.Path] {
|
||||
continue
|
||||
}
|
||||
if r.Policy.ScopeFor(RoleTalent).Kind == ScopeNone {
|
||||
t.Errorf("%s: talent may list it but no ownership scope narrows the rows", r.Path)
|
||||
}
|
||||
}
|
||||
}
|
||||
35
go-api/internal/domain/record.go
Normal file
35
go-api/internal/domain/record.go
Normal file
@@ -0,0 +1,35 @@
|
||||
package domain
|
||||
|
||||
// Record is one row as the API exposes it: the frontend's exact field names
|
||||
// mapped to JSON-ready values.
|
||||
//
|
||||
// A map rather than a struct per resource. The frontend treats every record as
|
||||
// an opaque bag of fields it round-trips unchanged — `store.js` stores whatever
|
||||
// it was handed and returns a clone — and fourteen structs totalling ~350
|
||||
// fields would add a transcription risk without adding a guarantee. Typing
|
||||
// lives in the Column descriptors instead, where validation and SQL both read
|
||||
// it from one place.
|
||||
type Record map[string]any
|
||||
|
||||
// ListParams is a parsed, validated collection query.
|
||||
type ListParams struct {
|
||||
Sort string // column name, without the leading '-'
|
||||
Desc bool
|
||||
Limit int
|
||||
Offset int
|
||||
Filters []Filter
|
||||
}
|
||||
|
||||
// Filter is one equality or membership test. See api-contract.md §6.
|
||||
type Filter struct {
|
||||
Column *Column
|
||||
Values []string // len > 1 means IN
|
||||
}
|
||||
|
||||
// Page is a collection result plus the metadata the envelope reports.
|
||||
type Page struct {
|
||||
Records []Record
|
||||
Total int
|
||||
Limit int
|
||||
Offset int
|
||||
}
|
||||
164
go-api/internal/domain/resource.go
Normal file
164
go-api/internal/domain/resource.go
Normal file
@@ -0,0 +1,164 @@
|
||||
// Package domain describes the API's resources: their columns, types and the
|
||||
// operations each one supports.
|
||||
//
|
||||
// The descriptors here are the single place the contract in
|
||||
// `docs/api-contract.md` is encoded. The repository, service and HTTP layers
|
||||
// are all driven from them, so a semantic is implemented once and applies
|
||||
// identically to every resource — which is the point. Fourteen hand-written
|
||||
// repositories would be fourteen chances to get NULLS LAST wrong.
|
||||
package domain
|
||||
|
||||
import "fmt"
|
||||
|
||||
// Kind is a column's value type, as the API presents it.
|
||||
type Kind int
|
||||
|
||||
const (
|
||||
KindString Kind = iota
|
||||
KindInt
|
||||
KindFloat
|
||||
KindBool
|
||||
KindTimestamp
|
||||
KindDate
|
||||
KindTextArray
|
||||
KindJSON
|
||||
KindUUID
|
||||
KindEnum
|
||||
)
|
||||
|
||||
// Op is a supported operation, as a bit set.
|
||||
type Op uint8
|
||||
|
||||
const (
|
||||
OpList Op = 1 << iota
|
||||
OpGet
|
||||
OpCreate
|
||||
OpUpdate
|
||||
OpDelete
|
||||
)
|
||||
|
||||
// Column is one database column and how the API treats it.
|
||||
type Column struct {
|
||||
Name string
|
||||
Kind Kind
|
||||
PGType string // the type every parameter is explicitly cast to
|
||||
NotNull bool // database-level NOT NULL
|
||||
ReadOnly bool // server-owned: ignored if present in a request body
|
||||
Required bool // must be supplied, non-blank, on create
|
||||
Enum []string // permitted values when Kind == KindEnum
|
||||
}
|
||||
|
||||
// Resource is one API resource and its backing table.
|
||||
type Resource struct {
|
||||
Name string // frontend entity name, used verbatim in error messages
|
||||
Path string // URL segment
|
||||
Table string
|
||||
Columns []Column
|
||||
DefaultSort string
|
||||
DefaultLimit int
|
||||
Ops Op
|
||||
// OrgNullable marks a table where a NULL org_id means "shared across every
|
||||
// organization" — the platform course library. Reads match org OR NULL.
|
||||
OrgNullable bool
|
||||
|
||||
// Policy is who may do what, and which rows they see. Attached from
|
||||
// policy.go, which is hand-written; nil means the resource permits nothing.
|
||||
Policy *Policy
|
||||
|
||||
byName map[string]*Column
|
||||
}
|
||||
|
||||
// Supports reports whether the resource exposes an operation.
|
||||
func (r *Resource) Supports(op Op) bool { return r.Ops&op != 0 }
|
||||
|
||||
// Column looks a column up by name.
|
||||
func (r *Resource) Column(name string) (*Column, bool) {
|
||||
if r.byName == nil {
|
||||
r.byName = make(map[string]*Column, len(r.Columns))
|
||||
for i := range r.Columns {
|
||||
r.byName[r.Columns[i].Name] = &r.Columns[i]
|
||||
}
|
||||
}
|
||||
c, ok := r.byName[name]
|
||||
return c, ok
|
||||
}
|
||||
|
||||
// Filterable reports whether a column may appear as a query filter.
|
||||
//
|
||||
// Arrays and JSON are excluded deliberately. `store.js` compares with `===`,
|
||||
// so a filter against an array column matches nothing today; supporting
|
||||
// containment here would be a silent behaviour change, not a fix.
|
||||
// See api-contract.md §6.
|
||||
func (r *Resource) Filterable(name string) bool {
|
||||
c, ok := r.Column(name)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
switch c.Kind {
|
||||
case KindTextArray, KindJSON:
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// Sortable reports whether a column may be sorted on. Any real column may be.
|
||||
func (r *Resource) Sortable(name string) bool {
|
||||
_, ok := r.Column(name)
|
||||
return ok
|
||||
}
|
||||
|
||||
// SelectExpr is the SQL that reads one column back in its API representation.
|
||||
//
|
||||
// The casts are not cosmetic. pgx hands back a [16]byte for uuid and a
|
||||
// pgtype.Numeric for numeric, neither of which JSON-encodes as the frontend
|
||||
// expects, and both are cheaper to fix in the projection than in Go.
|
||||
func (c Column) SelectExpr() string {
|
||||
switch c.Kind {
|
||||
case KindUUID:
|
||||
return fmt.Sprintf("%s::text AS %s", c.Name, c.Name)
|
||||
case KindFloat:
|
||||
return fmt.Sprintf("%s::float8 AS %s", c.Name, c.Name)
|
||||
case KindDate:
|
||||
return fmt.Sprintf("to_char(%s, 'YYYY-MM-DD') AS %s", c.Name, c.Name)
|
||||
case KindTimestamp:
|
||||
// Reproduces the millisecond ISO-8601 form the seed data uses, so a
|
||||
// record read back over HTTP is byte-identical to what the frontend
|
||||
// has always seen from localStorage.
|
||||
return fmt.Sprintf(
|
||||
`to_char(%s AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS %s`, c.Name, c.Name)
|
||||
case KindString:
|
||||
if c.PGType == "citext" {
|
||||
// pgx has no codec registered for citext, so without this the value
|
||||
// comes back as an unmapped type rather than a string.
|
||||
return fmt.Sprintf("%s::text AS %s", c.Name, c.Name)
|
||||
}
|
||||
return c.Name
|
||||
case KindInt:
|
||||
if c.PGType == "bigint" {
|
||||
// user_activity.id is an identity bigint. Every id the frontend
|
||||
// handles is an opaque string, so it stays one here too.
|
||||
return fmt.Sprintf("%s::text AS %s", c.Name, c.Name)
|
||||
}
|
||||
return c.Name
|
||||
default:
|
||||
return c.Name
|
||||
}
|
||||
}
|
||||
|
||||
// ResourceByPath indexes AllResources by URL segment.
|
||||
var ResourceByPath = func() map[string]*Resource {
|
||||
m := make(map[string]*Resource, len(AllResources))
|
||||
for _, r := range AllResources {
|
||||
m[r.Path] = r
|
||||
}
|
||||
return m
|
||||
}()
|
||||
|
||||
// ResourceByTable indexes AllResources by table name.
|
||||
var ResourceByTable = func() map[string]*Resource {
|
||||
m := make(map[string]*Resource, len(AllResources))
|
||||
for _, r := range AllResources {
|
||||
m[r.Table] = r
|
||||
}
|
||||
return m
|
||||
}()
|
||||
396
go-api/internal/domain/resources_gen.go
Normal file
396
go-api/internal/domain/resources_gen.go
Normal file
@@ -0,0 +1,396 @@
|
||||
// Code generated by scripts/gen_resources.py. DO NOT EDIT BY HAND.
|
||||
// Regenerate with: make gen-resources
|
||||
//
|
||||
// Column names, types, enum values and nullability are read out of
|
||||
// information_schema so they cannot drift from the migrations. The
|
||||
// per-resource metadata (path, default sort, default limit, supported
|
||||
// operations, required fields) comes from docs/api-contract.md.
|
||||
|
||||
package domain
|
||||
|
||||
// AllResources is every resource the API serves.
|
||||
var AllResources = []*Resource{
|
||||
{
|
||||
Name: "JobPosting", Path: "job-postings", Table: "job_postings",
|
||||
DefaultSort: "-created_date", DefaultLimit: 100,
|
||||
Ops: OpList | OpGet | OpCreate | OpUpdate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "created_by", Kind: KindUUID, PGType: "uuid", ReadOnly: true},
|
||||
{Name: "company", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "title", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "role_category", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "description", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "responsibilities", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "qualifications", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "nice_to_haves", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "custom_requirements", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "physical_requirements", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "leadership_expectations", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "attendance_expectations", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "min_experience_years", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "english_required", Kind: KindEnum, PGType: "english_level", NotNull: true, Enum: []string{"basic", "conversational", "fluent", "native"}},
|
||||
{Name: "certifications_required", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "skill_requirements", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "pay_range_min", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "pay_range_max", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "location", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "status", Kind: KindEnum, PGType: "posting_status", NotNull: true, Enum: []string{"draft", "active", "paused", "closed"}},
|
||||
{Name: "ai_generated", Kind: KindBool, PGType: "boolean", NotNull: true},
|
||||
{Name: "headcount", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "start_date", Kind: KindDate, PGType: "date"},
|
||||
{Name: "duration_months", Kind: KindFloat, PGType: "numeric"},
|
||||
{Name: "priority", Kind: KindEnum, PGType: "posting_priority", NotNull: true, Enum: []string{"urgent", "high", "normal"}},
|
||||
{Name: "vetting_criteria", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "JobApplication", Path: "job-applications", Table: "job_applications",
|
||||
DefaultSort: "-ai_score", DefaultLimit: 200,
|
||||
Ops: OpList | OpCreate | OpUpdate | OpDelete,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "job_posting_id", Kind: KindUUID, PGType: "uuid", NotNull: true, Required: true},
|
||||
{Name: "worker_profile_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "job_title", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "applicant_name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "email", Kind: KindString, PGType: "citext", NotNull: true, Required: true},
|
||||
{Name: "phone", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "years_experience", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "english_level", Kind: KindEnum, PGType: "english_level", NotNull: true, Enum: []string{"basic", "conversational", "fluent", "native"}},
|
||||
{Name: "certifications", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "availability", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "skills", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "companies_worked", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "client_rating", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "professional_summary", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "cover_letter", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "selfie_url", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "status", Kind: KindEnum, PGType: "application_status", NotNull: true, Enum: []string{"applied", "ai_screened", "shortlisted", "interview", "hired", "rejected", "assigned"}},
|
||||
{Name: "ai_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "ai_match_label", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "ai_summary", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "ai_strengths", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "ai_gaps", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "ai_recommendation", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "score_breakdown", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "screened_at", Kind: KindTimestamp, PGType: "timestamptz"},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "interview_id", Kind: KindUUID, PGType: "uuid"},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "AIInterview", Path: "ai-interviews", Table: "ai_interviews",
|
||||
DefaultSort: "-created_date", DefaultLimit: 100,
|
||||
Ops: OpList | OpCreate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "application_id", Kind: KindUUID, PGType: "uuid", NotNull: true, Required: true},
|
||||
{Name: "job_posting_id", Kind: KindUUID, PGType: "uuid", NotNull: true, Required: true},
|
||||
{Name: "job_title", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "candidate_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "messages", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "overall_interview_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "verdict", Kind: KindEnum, PGType: "interview_verdict", NotNull: true, Enum: []string{"hire", "maybe", "no"}},
|
||||
{Name: "hire_recommendation", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "integrity_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "ai_flags", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "category_scores", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "strengths", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "concerns", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "best_fit_roles", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "summary", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "reasoning", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "Staff", Path: "staff", Table: "staff",
|
||||
DefaultSort: "-created_date", DefaultLimit: 100,
|
||||
Ops: OpList | OpCreate | OpUpdate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "application_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "job_posting_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_profile_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "email", Kind: KindString, PGType: "citext", NotNull: true, Required: true},
|
||||
{Name: "phone", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "role", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "profile_tier", Kind: KindEnum, PGType: "profile_tier", NotNull: true, Enum: []string{"Beginner", "Cross-Trained", "Skilled"}},
|
||||
{Name: "hire_date", Kind: KindDate, PGType: "date", NotNull: true, Required: true},
|
||||
{Name: "ai_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "status", Kind: KindEnum, PGType: "staff_status", NotNull: true, Enum: []string{"onboarding", "active", "inactive"}},
|
||||
{Name: "client_rating", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "endorsement_text", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "endorsed_skills", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "review_date", Kind: KindDate, PGType: "date"},
|
||||
{Name: "reviewer_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "WorkerProfile", Path: "worker-profiles", Table: "worker_profiles",
|
||||
DefaultSort: "-krow_score", DefaultLimit: 500,
|
||||
Ops: OpList | OpCreate | OpUpdate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "user_id", Kind: KindUUID, PGType: "uuid", ReadOnly: true},
|
||||
{Name: "full_name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "email", Kind: KindString, PGType: "citext", NotNull: true, Required: true},
|
||||
{Name: "phone", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "address", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "selfie_url", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "languages", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "availability", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "transportation", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "certifications", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "experience", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "experience_years", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "current_position", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "desired_position", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "career_goals", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "skills", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "industries", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "personality", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "strengths", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "weaknesses", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "communication_style", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "salary_expectations", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "leadership_potential", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "ai_interview_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "krow_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "reliability_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "profile_completion", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "xp", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "completed_courses", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "earned_badges", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "capabilities", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "shifts_completed", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "attendance_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "performance_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "client_rating", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "supervisor_rating", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "status", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "score_breakdown", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "Course", Path: "courses", Table: "courses",
|
||||
DefaultSort: "-created_date", DefaultLimit: 200,
|
||||
Ops: OpList | OpGet | OpCreate | OpUpdate,
|
||||
OrgNullable: true,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", ReadOnly: true},
|
||||
{Name: "title", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "description", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "category", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "difficulty", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "xp", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "estimated_minutes", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "badge_reward", Kind: KindString, PGType: "text"},
|
||||
{Name: "proof_skill", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "skill_id", Kind: KindString, PGType: "text"},
|
||||
{Name: "target_level", Kind: KindEnum, PGType: "skill_level", Enum: []string{"beginner", "intermediate", "advanced", "expert"}},
|
||||
{Name: "required_level", Kind: KindEnum, PGType: "skill_level", Enum: []string{"beginner", "intermediate", "advanced", "expert"}},
|
||||
{Name: "completion_criteria", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "verification_criteria", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
{Name: "challenge", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "unlock_requirements", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "quiz", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "pass_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "status", Kind: KindEnum, PGType: "course_status", NotNull: true, Enum: []string{"active", "inactive"}},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "training_outline", Kind: KindTextArray, PGType: "text[]", NotNull: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "LearningPath", Path: "learning-paths", Table: "learning_paths",
|
||||
DefaultSort: "-created_date", DefaultLimit: 100,
|
||||
Ops: OpList,
|
||||
OrgNullable: true,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", ReadOnly: true},
|
||||
{Name: "name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "target_role", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "description", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "difficulty", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "steps", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "RoleCategory", Path: "role-categories", Table: "role_categories",
|
||||
DefaultSort: "-created_date", DefaultLimit: 100,
|
||||
Ops: OpList | OpCreate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "Certification", Path: "certifications", Table: "certifications",
|
||||
DefaultSort: "-created_date", DefaultLimit: 200,
|
||||
Ops: OpList | OpCreate | OpDelete,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "UserActivity", Path: "user-activity", Table: "user_activity",
|
||||
DefaultSort: "-created_date", DefaultLimit: 500,
|
||||
Ops: OpList | OpCreate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindInt, PGType: "bigint", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "event_type", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "user_id", Kind: KindUUID, PGType: "uuid", ReadOnly: true},
|
||||
{Name: "user_email", Kind: KindString, PGType: "citext", NotNull: true, ReadOnly: true},
|
||||
{Name: "user_name", Kind: KindString, PGType: "text", NotNull: true, ReadOnly: true},
|
||||
{Name: "account_type", Kind: KindString, PGType: "text", NotNull: true, ReadOnly: true},
|
||||
{Name: "details", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "position_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "application_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "candidate_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "interview_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_email", Kind: KindString, PGType: "citext"},
|
||||
{Name: "metadata", Kind: KindJSON, PGType: "jsonb"},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "Evidence", Path: "evidence", Table: "evidence",
|
||||
DefaultSort: "-created_date", DefaultLimit: 200,
|
||||
Ops: OpList | OpCreate | OpUpdate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "course_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_profile_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "course_title", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "skill", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "worker_email", Kind: KindString, PGType: "citext", NotNull: true, Required: true},
|
||||
{Name: "worker_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "type", Kind: KindEnum, PGType: "challenge_type", NotNull: true, Required: true, Enum: []string{"roleplay", "video", "photo_identify"}},
|
||||
{Name: "media_url", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "transcript", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "ai_verdict", Kind: KindEnum, PGType: "evidence_verdict", NotNull: true, Enum: []string{"verified", "needs_work", "failed"}},
|
||||
{Name: "ai_score", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "ai_rubric", Kind: KindJSON, PGType: "jsonb", NotNull: true},
|
||||
{Name: "ai_feedback", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "supervisor_verified", Kind: KindBool, PGType: "boolean", NotNull: true},
|
||||
{Name: "supervisor_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "verified_date", Kind: KindTimestamp, PGType: "timestamptz"},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "Assignment", Path: "assignments", Table: "assignments",
|
||||
DefaultSort: "-created_date", DefaultLimit: 500,
|
||||
Ops: OpList | OpCreate,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "job_posting_id", Kind: KindUUID, PGType: "uuid", NotNull: true, Required: true},
|
||||
{Name: "application_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_profile_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_email", Kind: KindString, PGType: "citext", NotNull: true, Required: true},
|
||||
{Name: "worker_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "starts_at", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, Required: true},
|
||||
{Name: "ends_at", Kind: KindTimestamp, PGType: "timestamptz"},
|
||||
{Name: "status", Kind: KindEnum, PGType: "assignment_status", NotNull: true, Enum: []string{"active", "completed", "cancelled"}},
|
||||
{Name: "source", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "match_score", Kind: KindInt, PGType: "int"},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "ShiftRecord", Path: "shift-records", Table: "shift_records",
|
||||
DefaultSort: "-created_date", DefaultLimit: 500,
|
||||
Ops: OpList,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "staff_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "assignment_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "job_posting_id", Kind: KindUUID, PGType: "uuid"},
|
||||
{Name: "worker_name", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "worker_email", Kind: KindString, PGType: "citext", NotNull: true},
|
||||
{Name: "role", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "role_category", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "shift_date", Kind: KindDate, PGType: "date", NotNull: true},
|
||||
{Name: "scheduled_start", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true},
|
||||
{Name: "scheduled_end", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true},
|
||||
{Name: "scheduled_hours", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "actual_start", Kind: KindTimestamp, PGType: "timestamptz"},
|
||||
{Name: "actual_end", Kind: KindTimestamp, PGType: "timestamptz"},
|
||||
{Name: "actual_hours", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "overtime_hours", Kind: KindFloat, PGType: "numeric", NotNull: true},
|
||||
{Name: "minutes_late", Kind: KindInt, PGType: "int", NotNull: true},
|
||||
{Name: "status", Kind: KindEnum, PGType: "shift_status", NotNull: true, Enum: []string{"present", "late", "absent", "no_show"}},
|
||||
{Name: "notes", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
// Badge serves NO endpoint: useBadges has zero consumers and every
|
||||
// badge the UI renders comes from worker_profiles.earned_badges. The
|
||||
// descriptor exists so the seeder can write the table. api-contract.md §2.
|
||||
{
|
||||
Name: "Badge", Path: "badges", Table: "badges",
|
||||
DefaultSort: "-created_date", DefaultLimit: 200,
|
||||
Ops: 0,
|
||||
Columns: []Column{
|
||||
{Name: "id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "legacy_id", Kind: KindString, PGType: "text", ReadOnly: true},
|
||||
{Name: "org_id", Kind: KindUUID, PGType: "uuid", NotNull: true, ReadOnly: true},
|
||||
{Name: "name", Kind: KindString, PGType: "text", NotNull: true, Required: true},
|
||||
{Name: "description", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "image_url", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "level", Kind: KindEnum, PGType: "badge_level", NotNull: true, Enum: []string{"bronze", "silver", "gold", "platinum"}},
|
||||
{Name: "requirements", Kind: KindString, PGType: "text", NotNull: true},
|
||||
{Name: "expiration_months", Kind: KindInt, PGType: "int"},
|
||||
{Name: "verification_status", Kind: KindEnum, PGType: "badge_verification", NotNull: true, Enum: []string{"pending", "verified", "expired"}},
|
||||
{Name: "created_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
{Name: "updated_date", Kind: KindTimestamp, PGType: "timestamptz", NotNull: true, ReadOnly: true},
|
||||
},
|
||||
},
|
||||
}
|
||||
220
go-api/internal/httpserver/api.go
Normal file
220
go-api/internal/httpserver/api.go
Normal file
@@ -0,0 +1,220 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/service"
|
||||
)
|
||||
|
||||
// maxBodyBytes bounds a request body. The largest thing the frontend sends is
|
||||
// an AI interview transcript; 4 MB is far above it and far below trouble.
|
||||
const maxBodyBytes = 4 << 20
|
||||
|
||||
// routeResources registers exactly the endpoints each resource supports.
|
||||
//
|
||||
// Only the declared operations are registered, so an unsupported one — DELETE
|
||||
// on a job posting, say — is answered by the mux with 405 rather than by a
|
||||
// handler that has to know it should refuse. The database having a table is
|
||||
// never a reason for an endpoint to exist. See api-contract.md §2.
|
||||
func (s *Server) routeResources(mux *http.ServeMux) int {
|
||||
count := 0
|
||||
for _, svc := range s.api.All() {
|
||||
res := svc.Resource()
|
||||
base := "/api/v1/" + res.Path
|
||||
item := base + "/{id}"
|
||||
|
||||
if res.Supports(domain.OpList) {
|
||||
mux.HandleFunc("GET "+base, s.handleList(svc))
|
||||
count++
|
||||
}
|
||||
if res.Supports(domain.OpCreate) {
|
||||
mux.HandleFunc("POST "+base, s.handleCreate(svc))
|
||||
count++
|
||||
}
|
||||
if res.Supports(domain.OpGet) {
|
||||
mux.HandleFunc("GET "+item, s.handleGet(svc))
|
||||
count++
|
||||
}
|
||||
if res.Supports(domain.OpUpdate) {
|
||||
mux.HandleFunc("PATCH "+item, s.handleUpdate(svc))
|
||||
count++
|
||||
}
|
||||
if res.Supports(domain.OpDelete) {
|
||||
mux.HandleFunc("DELETE "+item, s.handleDelete(svc))
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
// authorize is the role gate. It runs before any query.
|
||||
//
|
||||
// It answers 403 and nothing else — never 404, and never a message naming the
|
||||
// role required. Which rows the caller may then see is a separate question,
|
||||
// answered in SQL by the repository, and its refusal is a 404 so that existence
|
||||
// does not leak. Keeping the two apart is what makes "403 means your role, 404
|
||||
// means not yours or not there" a rule a client can rely on.
|
||||
//
|
||||
// The role comes from the session-resolved identity. A role the API does not
|
||||
// recognise authorizes nothing.
|
||||
func (s *Server) authorize(w http.ResponseWriter, r *http.Request,
|
||||
svc *service.Service, op domain.Op) (authctx.Identity, bool) {
|
||||
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
// Unreachable: the middleware refuses an unauthenticated request before
|
||||
// the router sees it. A missing identity here is a wiring bug, not a
|
||||
// client error.
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return authctx.Identity{}, false
|
||||
}
|
||||
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || !svc.Resource().Policy.Allows(op, role) {
|
||||
s.log.Warn("authorization refused",
|
||||
"user_id", ident.UserID, "role", ident.Role,
|
||||
"resource", svc.Resource().Path, "method", r.Method, "path", r.URL.Path)
|
||||
writeError(w, s.log, domain.Forbidden())
|
||||
return authctx.Identity{}, false
|
||||
}
|
||||
return ident, true
|
||||
}
|
||||
|
||||
func (s *Server) handleList(svc *service.Service) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
ident, ok := s.authorize(w, r, svc, domain.OpList)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
params, err := svc.ParseList(r.URL.Query())
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
page, err := svc.List(r.Context(), ident, params)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writePage(w, page)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleGet(svc *service.Service) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
ident, ok := s.authorize(w, r, svc, domain.OpGet)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
rec, err := svc.Get(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleCreate(svc *service.Service) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
ident, ok := s.authorize(w, r, svc, domain.OpCreate)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
rec, err := svc.Create(r.Context(), ident, body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusCreated, rec)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleUpdate(svc *service.Service) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
ident, ok := s.authorize(w, r, svc, domain.OpUpdate)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
rec, err := svc.Update(r.Context(), ident, r.PathValue("id"), body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleDelete(svc *service.Service) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
ident, ok := s.authorize(w, r, svc, domain.OpDelete)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
rec, err := svc.Delete(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
}
|
||||
|
||||
// decodeBody reads a JSON object body.
|
||||
//
|
||||
// DisallowUnknownFields is not used — the target is a map, so every field is
|
||||
// "known" here. Unknown *columns* are rejected in the service, where the
|
||||
// resource's schema is available to say which those are.
|
||||
// decodeInto reads a JSON body into a typed struct.
|
||||
//
|
||||
// Beside decodeBody rather than replacing it: the resource handlers genuinely
|
||||
// want the open map, because a PATCH body is "whichever fields the caller sent"
|
||||
// and a struct cannot distinguish an absent field from a zero one. The auth
|
||||
// endpoints have a fixed, closed shape, and a struct says so.
|
||||
func decodeInto(r *http.Request, dst any) error {
|
||||
defer func() { _ = r.Body.Close() }()
|
||||
raw, err := io.ReadAll(http.MaxBytesReader(nil, r.Body, maxBodyBytes))
|
||||
if err != nil {
|
||||
return domain.Invalid("request body could not be read")
|
||||
}
|
||||
if len(raw) == 0 {
|
||||
return domain.Invalid("request body must be a JSON object")
|
||||
}
|
||||
if err := json.Unmarshal(raw, dst); err != nil {
|
||||
return domain.Invalid("request body must be a JSON object")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func decodeBody(r *http.Request) (domain.Record, error) {
|
||||
defer func() { _ = r.Body.Close() }()
|
||||
raw, err := io.ReadAll(http.MaxBytesReader(nil, r.Body, maxBodyBytes))
|
||||
if err != nil {
|
||||
return nil, domain.Invalid("request body could not be read")
|
||||
}
|
||||
if len(raw) == 0 {
|
||||
return domain.Record{}, nil
|
||||
}
|
||||
var body domain.Record
|
||||
if err := json.Unmarshal(raw, &body); err != nil {
|
||||
return nil, domain.Invalid("request body must be a JSON object")
|
||||
}
|
||||
if body == nil {
|
||||
return domain.Record{}, nil
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
1114
go-api/internal/httpserver/api_test.go
Normal file
1114
go-api/internal/httpserver/api_test.go
Normal file
File diff suppressed because it is too large
Load Diff
374
go-api/internal/httpserver/auth.go
Normal file
374
go-api/internal/httpserver/auth.go
Normal file
@@ -0,0 +1,374 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/orgctx"
|
||||
)
|
||||
|
||||
// The authentication surface: sign in, sign out, and the middleware that turns
|
||||
// a cookie into an identity.
|
||||
//
|
||||
// The shape of the whole thing is one sentence: the browser holds an opaque
|
||||
// random string it cannot read, the database holds SHA-256 of that string, and
|
||||
// every protected request is a lookup from one to the other. No claim travels
|
||||
// in the request. There is no token in a JSON body, no user id in a query
|
||||
// string, no organization in a header — those are all things a client can
|
||||
// write, and a client writing its own identity is the bug this replaces.
|
||||
|
||||
// sessionCookieName is the cookie the browser holds.
|
||||
//
|
||||
// The "__Host-" prefix would be stronger — browsers enforce Secure, Path=/ and
|
||||
// no Domain on it — but it also *requires* Secure, which cannot be set over
|
||||
// plain HTTP on localhost. A cookie name that only works in production is worse
|
||||
// than a plain one that works everywhere, so the hardening is done by the
|
||||
// attributes below instead, where it can be conditional.
|
||||
const sessionCookieName = "krow_session"
|
||||
|
||||
/* ── Cookie ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// secureCookies reports whether Secure may be set.
|
||||
//
|
||||
// Secure means "only ever send this over HTTPS". Setting it in development
|
||||
// would mean the browser silently declines to send the cookie back to
|
||||
// http://localhost, and the symptom is an endless loop of successful logins
|
||||
// that never authenticate anything.
|
||||
func (s *Server) secureCookies() bool { return s.cfg.AppEnv != "development" }
|
||||
|
||||
// setSessionCookie writes the raw token to the browser.
|
||||
//
|
||||
// This is the only place the raw token is written to a response, and it goes
|
||||
// into a Set-Cookie header rather than a body: HttpOnly means no script on the
|
||||
// page can read it, which is what makes an XSS bug stop short of session theft.
|
||||
//
|
||||
// maxAge matches the session's own lifetime so the browser drops the cookie at
|
||||
// roughly the moment the server would refuse it. The server is still the
|
||||
// authority — a cookie the browser keeps too long is simply rejected — but a
|
||||
// cookie that expires with its session keeps the two honest.
|
||||
func (s *Server) setSessionCookie(w http.ResponseWriter, token string, lifetime time.Duration) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: sessionCookieName,
|
||||
Value: token,
|
||||
Path: "/",
|
||||
// HttpOnly: script cannot read it.
|
||||
HttpOnly: true,
|
||||
// Lax, not Strict and not None. Strict would drop the cookie on any
|
||||
// cross-site navigation, so following a link into the app would land on
|
||||
// a login page despite a live session. None would require Secure and
|
||||
// would send the cookie on cross-site POSTs, which is the CSRF hole Lax
|
||||
// exists to close.
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
Secure: s.secureCookies(),
|
||||
MaxAge: int(lifetime.Seconds()),
|
||||
})
|
||||
}
|
||||
|
||||
// clearSessionCookie expires the cookie in the browser.
|
||||
//
|
||||
// The attributes must match the ones it was set with — a cookie is identified
|
||||
// by name, domain and path, so clearing it with a different Path leaves the
|
||||
// original in place and the browser keeps sending a token the server has
|
||||
// already deleted.
|
||||
func (s *Server) clearSessionCookie(w http.ResponseWriter) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: sessionCookieName,
|
||||
Value: "",
|
||||
Path: "/",
|
||||
HttpOnly: true,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
Secure: s.secureCookies(),
|
||||
MaxAge: -1,
|
||||
})
|
||||
}
|
||||
|
||||
// sessionToken reads the raw token out of the request, if there is one.
|
||||
func sessionToken(r *http.Request) string {
|
||||
c, err := r.Cookie(sessionCookieName)
|
||||
if err != nil || c == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(c.Value)
|
||||
}
|
||||
|
||||
/* ── Routes ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
func (s *Server) routeAuth(mux *http.ServeMux) int {
|
||||
mux.HandleFunc("POST /api/v1/auth/login", s.handleLogin)
|
||||
mux.HandleFunc("POST /api/v1/auth/logout", s.handleLogout)
|
||||
return 2
|
||||
}
|
||||
|
||||
// loginRequest is the body of POST /api/v1/auth/login.
|
||||
type loginRequest struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
RememberMe bool `json:"remember_me"`
|
||||
}
|
||||
|
||||
// handleLogin verifies a password and issues a session.
|
||||
//
|
||||
// The order of operations is deliberate:
|
||||
//
|
||||
// 1. Parse and validate the *shape* of the request. A missing field is a
|
||||
// malformed request, not a failed login, and saying so reveals nothing.
|
||||
// 2. Check the rate limit, before any expensive work. Refusing early is the
|
||||
// point — an attacker must not be able to make the server hash for them.
|
||||
// 3. Verify the credentials, which takes the same measurable time whether the
|
||||
// email exists or not (see auth.Credentials).
|
||||
// 4. Issue the session and set the cookie.
|
||||
//
|
||||
// Every failure in step 3 produces one identical response. The reason goes to
|
||||
// the log, at warn, with the email — which is already in the request — and
|
||||
// never the password.
|
||||
func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
var req loginRequest
|
||||
if err := decodeInto(r, &req); err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
email := strings.TrimSpace(req.Email)
|
||||
details := map[string]string{}
|
||||
if email == "" {
|
||||
details["email"] = "an email address is required"
|
||||
}
|
||||
if req.Password == "" {
|
||||
details["password"] = "a password is required"
|
||||
}
|
||||
if len(details) > 0 {
|
||||
writeError(w, s.log, domain.Validation("email and password are required", details))
|
||||
return
|
||||
}
|
||||
|
||||
// Two budgets, both consulted, both counted. The per-email budget stops one
|
||||
// account being ground down from many addresses; the per-address budget,
|
||||
// which is wider, stops one host working through many accounts. They are
|
||||
// separate limiters because they are deliberately different sizes — see the
|
||||
// note on Server.
|
||||
addr := clientAddr(r)
|
||||
emailKey := strings.ToLower(email)
|
||||
for _, check := range []struct {
|
||||
limiter *attemptLimiter
|
||||
key string
|
||||
scope string
|
||||
}{
|
||||
{s.loginByEmail, emailKey, "email"},
|
||||
{s.loginByAddr, addr, "address"},
|
||||
} {
|
||||
if ok, retryAfter := check.limiter.Allow(check.key); !ok {
|
||||
w.Header().Set("Retry-After", retryAfterSeconds(retryAfter))
|
||||
s.log.Warn("login rate limited", "scope", check.scope,
|
||||
"email", email, "addr", addr,
|
||||
"retry_after_seconds", retryAfterSeconds(retryAfter))
|
||||
writeError(w, s.log, domain.RateLimited(
|
||||
"too many sign-in attempts; wait a few minutes and try again"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
user, reason, err := s.credentials.Verify(r.Context(), email, req.Password)
|
||||
if errors.Is(err, auth.ErrInvalidCredentials) {
|
||||
s.loginByEmail.Fail(emailKey)
|
||||
s.loginByAddr.Fail(addr)
|
||||
// The reason is the whole value of this line and must never leave it.
|
||||
s.log.Warn("login failed", "email", email, "addr", addr, "reason", string(reason))
|
||||
writeError(w, s.log, domain.Unauthenticated())
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
// The database is down, or a stored hash is unreadable. The caller's
|
||||
// credentials were never judged, so this is a 500 and not a 401.
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
token, sess, err := s.sessions.Issue(r.Context(), user.ID, req.RememberMe)
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
// A correct password clears the email's penalty, so two typos followed by a
|
||||
// success leave nothing behind. The address counter is left alone: one
|
||||
// correct login should not wipe the budget for every other account being
|
||||
// tried from the same host.
|
||||
s.loginByEmail.Reset(emailKey)
|
||||
|
||||
s.setSessionCookie(w, token, time.Until(sess.ExpiresAt))
|
||||
|
||||
// Best effort, deliberately after the session exists: a failure to stamp
|
||||
// last_login_at is a lost diagnostic, not a reason to refuse a sign-in that
|
||||
// has already succeeded.
|
||||
if err := s.users.MarkLoggedIn(r.Context(), user.ID, s.now()); err != nil {
|
||||
s.log.Warn("could not record last_login_at", "user_id", user.ID, "error", err)
|
||||
}
|
||||
|
||||
s.log.Info("login", "user_id", user.ID, "email", user.Email,
|
||||
"remember_me", req.RememberMe, "session_id", sess.ID,
|
||||
"expires_at", sess.ExpiresAt, "absolute_expires_at", sess.AbsoluteExpiresAt)
|
||||
|
||||
// The body is the user, in exactly the shape GET /me returns, so the
|
||||
// frontend can render the signed-in state without a second round trip.
|
||||
//
|
||||
// The token is NOT here and must never be. It went out in a Set-Cookie
|
||||
// header the page cannot read; putting it in the body would hand it to
|
||||
// every script on the page and undo HttpOnly entirely.
|
||||
record, err := s.userRecord(r.Context(), s.db.Pool, user.ID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, record)
|
||||
}
|
||||
|
||||
// handleLogout revokes the session behind the cookie and clears the cookie.
|
||||
//
|
||||
// Idempotent by construction: no cookie, an unknown token and a live session
|
||||
// all end the same way — the cookie is cleared and the answer is 200. Logging
|
||||
// out is a request to not be signed in, and the caller is not signed in
|
||||
// afterwards in every one of those cases.
|
||||
//
|
||||
// It is deliberately public. Requiring a valid session to log out means a user
|
||||
// whose session has already expired gets a 401 from the one action that would
|
||||
// have tidied up their stale cookie.
|
||||
func (s *Server) handleLogout(w http.ResponseWriter, r *http.Request) {
|
||||
if token := sessionToken(r); token != "" {
|
||||
if err := s.sessions.Revoke(r.Context(), token); err != nil {
|
||||
// Revoke already treats "no such session" as success, so this is a
|
||||
// real failure — the database, most likely. Clearing the cookie is
|
||||
// still the right thing to do, and reporting a 500 for a logout
|
||||
// would leave the caller signed in with no way to fix it.
|
||||
s.log.Error("could not revoke session on logout", "error", err)
|
||||
}
|
||||
}
|
||||
s.clearSessionCookie(w)
|
||||
writeJSON(w, http.StatusOK, envelope{Data: map[string]any{"status": "signed_out"}})
|
||||
}
|
||||
|
||||
/* ── Middleware ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// publicPaths are the only endpoints reachable without a session.
|
||||
//
|
||||
// An allowlist rather than a list of protected prefixes, so the failure mode of
|
||||
// forgetting to update it is a route that refuses everyone — not one that
|
||||
// serves everyone. A new endpoint is private until someone deliberately says
|
||||
// otherwise, which is the direction a mistake should fall in.
|
||||
var publicPaths = map[string]bool{
|
||||
"/health": true,
|
||||
"/api/v1/auth/login": true,
|
||||
"/api/v1/auth/logout": true,
|
||||
}
|
||||
|
||||
// authenticate resolves the session cookie into an identity, or refuses.
|
||||
//
|
||||
// This replaces devOrgMiddleware, which put a fixed organization on every
|
||||
// request with no credential behind it. The seam is the same one that comment
|
||||
// promised: everything downstream still reads the organization from
|
||||
// orgctx, and not one service or repository changed.
|
||||
//
|
||||
// What the request cannot influence: nothing here reads the body, the query
|
||||
// string or any header other than Cookie. The user id, the organization and the
|
||||
// role are all read from the sessions and users tables, keyed by a token the
|
||||
// client cannot forge without already holding it.
|
||||
func (s *Server) authenticate(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if publicPaths[r.URL.Path] {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
token := sessionToken(r)
|
||||
if token == "" {
|
||||
writeError(w, s.log, domain.Unauthenticated())
|
||||
return
|
||||
}
|
||||
|
||||
sess, err := s.sessions.Authenticate(r.Context(), token)
|
||||
if err != nil {
|
||||
// Not found and expired are logged apart and answered identically.
|
||||
// Clearing the cookie stops the browser re-sending a token that
|
||||
// will never work again.
|
||||
s.log.Debug("session rejected", "reason", sessionRejection(err), "path", r.URL.Path)
|
||||
if errors.Is(err, auth.ErrSessionNotFound) || errors.Is(err, auth.ErrSessionExpired) ||
|
||||
errors.Is(err, auth.ErrEmptyToken) {
|
||||
s.clearSessionCookie(w)
|
||||
writeError(w, s.log, domain.Unauthenticated())
|
||||
return
|
||||
}
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
// The user is re-read on every request rather than cached in the
|
||||
// session row, so suspending an account takes effect on the account's
|
||||
// next request instead of whenever its session happens to lapse.
|
||||
user, err := s.users.FindByID(r.Context(), sess.UserID)
|
||||
if err != nil {
|
||||
if errors.Is(err, auth.ErrUserNotFound) {
|
||||
// The FK cascades, so this should be unreachable. If it happens
|
||||
// the session is orphaned and worth destroying.
|
||||
s.log.Warn("session references a missing user", "session_id", sess.ID)
|
||||
_ = s.sessions.RevokeID(r.Context(), sess.ID)
|
||||
s.clearSessionCookie(w)
|
||||
writeError(w, s.log, domain.Unauthenticated())
|
||||
return
|
||||
}
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
if !user.IsActive() {
|
||||
// Suspension revokes on contact. Leaving the session alive would
|
||||
// mean a suspended account keeps a working cookie for up to thirty
|
||||
// days, refused one request at a time.
|
||||
s.log.Warn("session for an inactive user revoked",
|
||||
"user_id", user.ID, "status", user.Status)
|
||||
_ = s.sessions.RevokeID(r.Context(), sess.ID)
|
||||
s.clearSessionCookie(w)
|
||||
writeError(w, s.log, domain.Unauthenticated())
|
||||
return
|
||||
}
|
||||
|
||||
id := authctx.Identity{
|
||||
UserID: user.ID, OrgID: user.OrgID, Email: user.Email,
|
||||
FullName: user.FullName, Role: user.Role, AccountType: user.AccountType,
|
||||
Status: user.Status, SessionID: sess.ID, ExpiresAt: sess.ExpiresAt,
|
||||
}
|
||||
ctx := authctx.With(r.Context(), id)
|
||||
// The organization comes from the user's row, never from the request.
|
||||
// Every service and repository already takes it as a parameter, so this
|
||||
// one line is the whole of the tenancy change.
|
||||
ctx = orgctx.With(ctx, user.OrgID)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
|
||||
// sessionRejection names why a session was refused, for the log only.
|
||||
func sessionRejection(err error) string {
|
||||
switch {
|
||||
case errors.Is(err, auth.ErrSessionNotFound):
|
||||
return "not_found"
|
||||
case errors.Is(err, auth.ErrSessionExpired):
|
||||
return "expired"
|
||||
case errors.Is(err, auth.ErrEmptyToken):
|
||||
return "empty_token"
|
||||
default:
|
||||
return "error"
|
||||
}
|
||||
}
|
||||
|
||||
// retryAfterSeconds renders a duration for the Retry-After header, rounded up
|
||||
// and never below one second — "Retry-After: 0" invites an immediate retry.
|
||||
func retryAfterSeconds(d time.Duration) string {
|
||||
secs := int(d.Round(time.Second) / time.Second)
|
||||
if secs < 1 {
|
||||
secs = 1
|
||||
}
|
||||
return strconv.Itoa(secs)
|
||||
}
|
||||
926
go-api/internal/httpserver/auth_test.go
Normal file
926
go-api/internal/httpserver/auth_test.go
Normal file
@@ -0,0 +1,926 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// shortSessions keeps the expiry arithmetic in these tests small. The real
|
||||
// lifetimes are asserted against the production policy in
|
||||
// TestSessionLifetimes, which is the test that would catch a change to them.
|
||||
var shortSessions = auth.Policy{
|
||||
IdleLifetime: time.Hour,
|
||||
AbsoluteLifetime: 3 * time.Hour,
|
||||
RememberIdleLifetime: 24 * time.Hour,
|
||||
RememberAbsoluteLifetime: 72 * time.Hour,
|
||||
}
|
||||
|
||||
// clockedAPI is newAPI with a clock the test drives, so expiry can be reached
|
||||
// without sleeping.
|
||||
func clockedAPI(t *testing.T, opts ...httpserver.Option) (*api, *time.Time) {
|
||||
t.Helper()
|
||||
now := time.Date(2026, 8, 22, 9, 0, 0, 0, time.UTC)
|
||||
all := append([]httpserver.Option{
|
||||
httpserver.WithClock(func() time.Time { return now }),
|
||||
httpserver.WithSessionPolicy(shortSessions),
|
||||
}, opts...)
|
||||
return newAPI(t, all...), &now
|
||||
}
|
||||
|
||||
/* ── 1-5. Login and its failure modes ───────────────────────────────────── */
|
||||
|
||||
// 1. A correct email and password sign in and return the user.
|
||||
func TestLoginSucceeds(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
result := signIn(t, a.handler, a.email, harnessPassword, false)
|
||||
if result.code != http.StatusOK {
|
||||
t.Fatalf("login = %d, want 200 (%v)", result.code, result.body)
|
||||
}
|
||||
|
||||
data, ok := result.body["data"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("login body has no data object: %v", result.body)
|
||||
}
|
||||
if data["id"] != a.userID {
|
||||
t.Errorf("login returned user %v, want %v", data["id"], a.userID)
|
||||
}
|
||||
if data["email"] != a.email {
|
||||
t.Errorf("login returned email %v, want %v", data["email"], a.email)
|
||||
}
|
||||
// The frontend renders the signed-in state straight from this body, so the
|
||||
// embedded preferences must be there as they are on GET /me.
|
||||
if _, ok := data["preferences"].(map[string]any); !ok {
|
||||
t.Error("login response has no embedded preferences")
|
||||
}
|
||||
|
||||
// The email is case-insensitive: users type their address how they like.
|
||||
upper := signIn(t, a.handler, strings.ToUpper(a.email), harnessPassword, false)
|
||||
if upper.code != http.StatusOK {
|
||||
t.Errorf("login with an upper-case email = %d, want 200", upper.code)
|
||||
}
|
||||
|
||||
// last_login_at is stamped.
|
||||
var lastLogin *time.Time
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT last_login_at FROM users WHERE id = $1::uuid`, a.userID).Scan(&lastLogin); err != nil {
|
||||
t.Fatalf("read last_login_at: %v", err)
|
||||
}
|
||||
if lastLogin == nil {
|
||||
t.Error("a successful login did not record last_login_at")
|
||||
}
|
||||
}
|
||||
|
||||
// 2, 3, 4, 5. Every credential failure is externally identical.
|
||||
//
|
||||
// This is one test rather than four because the property under test is the
|
||||
// sameness: a wrong password, an unknown address, an account with no password
|
||||
// and a suspended account must be indistinguishable from outside. Asserting
|
||||
// each in isolation would not catch the one thing that matters.
|
||||
func TestLoginFailuresAreIndistinguishable(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
// A second account, suspended, with a valid password set.
|
||||
suspended := newUser(t, a.h.Pool, a.orgID, "suspended@example.test", "employer")
|
||||
setStatus(t, a.h.Pool, suspended, "suspended")
|
||||
|
||||
// A third with no password at all — the state every seeded user starts in.
|
||||
passwordless := "nopassword@example.test"
|
||||
if _, err := a.h.Pool.Exec(context.Background(),
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, 'No Password')`,
|
||||
a.orgID, passwordless); err != nil {
|
||||
t.Fatalf("create the passwordless user: %v", err)
|
||||
}
|
||||
|
||||
cases := map[string]struct{ email, password string }{
|
||||
"wrong password": {a.email, "definitely-not-the-password"},
|
||||
"unknown email": {"nobody@example.invalid", harnessPassword},
|
||||
"suspended user": {"suspended@example.test", harnessPassword},
|
||||
"no password set": {passwordless, harnessPassword},
|
||||
"empty-ish password": {a.email, "x"},
|
||||
}
|
||||
|
||||
var bodies []string
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
result := signIn(t, a.handler, tc.email, tc.password, false)
|
||||
if result.code != http.StatusUnauthorized {
|
||||
t.Fatalf("login = %d, want 401", result.code)
|
||||
}
|
||||
if result.cookie != nil && result.cookie.Value != "" {
|
||||
t.Error("a failed login set a session cookie")
|
||||
}
|
||||
bodies = append(bodies, result.raw.Body.String())
|
||||
|
||||
// The message must not name the reason.
|
||||
body := strings.ToLower(result.raw.Body.String())
|
||||
for _, leak := range []string{
|
||||
"password", "suspend", "inactive", "not found", "no such",
|
||||
"unknown", "exist", "email address is not",
|
||||
} {
|
||||
if strings.Contains(body, leak) {
|
||||
t.Errorf("the login error mentions %q, which distinguishes the failure:\n%s",
|
||||
leak, result.raw.Body.String())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// 5. Byte-for-byte identical, not merely "all 401".
|
||||
for i := 1; i < len(bodies); i++ {
|
||||
if bodies[i] != bodies[0] {
|
||||
t.Errorf("login failures differ:\n%s\nvs\n%s", bodies[0], bodies[i])
|
||||
}
|
||||
}
|
||||
|
||||
// A suspended user with the RIGHT password is still refused. Worth its own
|
||||
// assertion: this is the check that must come after the password test, or
|
||||
// the timing of the refusal leaks that the account exists.
|
||||
if got := signIn(t, a.handler, "suspended@example.test", harnessPassword, false); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("a suspended user with a correct password got %d, want 401", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// A malformed request is a 422 about the request, not a 401 about an account.
|
||||
// Saying "you did not send a password" reveals nothing about any user.
|
||||
func TestLoginRejectsMalformedRequests(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
for name, payload := range map[string]any{
|
||||
"no email": map[string]any{"password": harnessPassword},
|
||||
"no password": map[string]any{"email": a.email},
|
||||
"both blank": map[string]any{"email": " ", "password": ""},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
r := a.doAnon("POST", "/api/v1/auth/login", payload)
|
||||
if r.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("login = %d, want 422", r.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6-9. The session cookie ────────────────────────────────────────────── */
|
||||
|
||||
// 6, 7. The cookie is created with the right attributes, and the token appears
|
||||
// nowhere a script could read it.
|
||||
func TestLoginSetsHardenedCookieAndNeverReturnsTheToken(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
result := signIn(t, a.handler, a.email, harnessPassword, false)
|
||||
|
||||
c := result.cookie
|
||||
if c == nil {
|
||||
t.Fatal("login set no session cookie")
|
||||
}
|
||||
if c.Value == "" {
|
||||
t.Fatal("the session cookie is empty")
|
||||
}
|
||||
if !c.HttpOnly {
|
||||
t.Error("the session cookie is not HttpOnly: a script on the page could read it")
|
||||
}
|
||||
if c.SameSite != http.SameSiteLaxMode {
|
||||
t.Errorf("SameSite = %v, want Lax", c.SameSite)
|
||||
}
|
||||
if c.Path != "/" {
|
||||
t.Errorf("Path = %q, want /", c.Path)
|
||||
}
|
||||
// APP_ENV=development in this harness, and localhost is plain HTTP: a
|
||||
// Secure cookie would never be sent back. TestCookieIsSecureOutsideDevelopment
|
||||
// covers the other half.
|
||||
if c.Secure {
|
||||
t.Error("the cookie is Secure in development; the browser would never return it over HTTP")
|
||||
}
|
||||
|
||||
// 7. The raw token is in the Set-Cookie header and nowhere else.
|
||||
if strings.Contains(result.raw.Body.String(), c.Value) {
|
||||
t.Error("the login response body contains the session token")
|
||||
}
|
||||
// And the same for every other response the API gives while signed in.
|
||||
me := a.doWith(c, "GET", "/api/v1/me")
|
||||
encoded, _ := json.Marshal(me.body)
|
||||
if strings.Contains(string(encoded), c.Value) {
|
||||
t.Error("GET /me leaks the session token")
|
||||
}
|
||||
|
||||
// 11 (verification list). The database holds the hash, never the token.
|
||||
var rawRows, hashRows int
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`, c.Value).Scan(&rawRows); err != nil {
|
||||
t.Fatalf("scan for a stored raw token: %v", err)
|
||||
}
|
||||
if rawRows != 0 {
|
||||
t.Error("the raw session token is stored in PostgreSQL")
|
||||
}
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`,
|
||||
auth.HashToken(c.Value)).Scan(&hashRows); err != nil {
|
||||
t.Fatalf("scan for the stored hash: %v", err)
|
||||
}
|
||||
if hashRows != 1 {
|
||||
t.Errorf("%d session rows hold the token's hash, want 1", hashRows)
|
||||
}
|
||||
}
|
||||
|
||||
// Outside development the cookie must be Secure, or it can be read off the wire.
|
||||
func TestCookieIsSecureOutsideDevelopment(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
srv := newServer(t, h, nil)
|
||||
userID, email := seededUser(t, h.Pool)
|
||||
setPassword(t, h.Pool, userID)
|
||||
|
||||
// A production-shaped server over the same database.
|
||||
prod := newServerWithEnv(t, h, "production")
|
||||
if got := signIn(t, prod, email, harnessPassword, false); got.cookie == nil || !got.cookie.Secure {
|
||||
t.Errorf("the cookie is not Secure when APP_ENV=production: %+v", got.cookie)
|
||||
}
|
||||
// The development server, for contrast, on the same database.
|
||||
if got := signIn(t, srv.Handler(), email, harnessPassword, false); got.cookie == nil || got.cookie.Secure {
|
||||
t.Error("the cookie is Secure in development")
|
||||
}
|
||||
}
|
||||
|
||||
// 8, 9. Normal and Remember Me sessions get the lifetimes the decision names,
|
||||
// in the cookie and in the row.
|
||||
func TestSessionLifetimes(t *testing.T) {
|
||||
a := newAPI(t) // the production policy: 12h / 24h and 30d / 90d
|
||||
|
||||
for name, tc := range map[string]struct {
|
||||
remember bool
|
||||
wantIdle time.Duration
|
||||
wantAbsolute time.Duration
|
||||
}{
|
||||
"normal": {false, 12 * time.Hour, 24 * time.Hour},
|
||||
"remember me": {true, 30 * 24 * time.Hour, 90 * 24 * time.Hour},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
before := time.Now()
|
||||
result := signIn(t, a.handler, a.email, harnessPassword, tc.remember)
|
||||
if result.code != http.StatusOK || result.cookie == nil {
|
||||
t.Fatalf("login = %d", result.code)
|
||||
}
|
||||
|
||||
// The cookie's own lifetime matches the session's.
|
||||
wantMaxAge := int(tc.wantIdle.Seconds())
|
||||
if drift := result.cookie.MaxAge - wantMaxAge; drift > 5 || drift < -5 {
|
||||
t.Errorf("cookie Max-Age = %d, want about %d", result.cookie.MaxAge, wantMaxAge)
|
||||
}
|
||||
|
||||
// And so does the row, which is the authority.
|
||||
var expires, absolute time.Time
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT expires_at, absolute_expires_at FROM sessions WHERE token_hash = $1::text`,
|
||||
auth.HashToken(result.cookie.Value)).Scan(&expires, &absolute); err != nil {
|
||||
t.Fatalf("read the session row: %v", err)
|
||||
}
|
||||
assertAbout(t, "expires_at", expires.Sub(before), tc.wantIdle)
|
||||
assertAbout(t, "absolute_expires_at", absolute.Sub(before), tc.wantAbsolute)
|
||||
if absolute.Before(expires) {
|
||||
t.Error("the absolute deadline is before the sliding one")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertAbout(t *testing.T, name string, got, want time.Duration) {
|
||||
t.Helper()
|
||||
if drift := got - want; drift > time.Minute || drift < -time.Minute {
|
||||
t.Errorf("%s is %v from now, want about %v", name, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 10-11. Logout ──────────────────────────────────────────────────────── */
|
||||
|
||||
// 10, 11. Logging out revokes the session, clears the cookie, and is idempotent.
|
||||
func TestLogout(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
// Signed in, the protected endpoint works.
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me before logout = %d, want 200", got.code)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest("POST", "/api/v1/auth/logout", nil)
|
||||
req.AddCookie(a.cookie)
|
||||
rec := httptest.NewRecorder()
|
||||
a.handler.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("logout = %d, want 200 (%s)", rec.Code, rec.Body.String())
|
||||
}
|
||||
// The cookie is expired in the browser.
|
||||
var cleared *http.Cookie
|
||||
for _, c := range rec.Result().Cookies() {
|
||||
if c.Name == sessionCookie {
|
||||
cleared = c
|
||||
}
|
||||
}
|
||||
if cleared == nil {
|
||||
t.Fatal("logout did not clear the session cookie")
|
||||
}
|
||||
if cleared.MaxAge >= 0 || cleared.Value != "" {
|
||||
t.Errorf("the cleared cookie is %+v, want an empty value and a negative Max-Age", cleared)
|
||||
}
|
||||
// The attributes must match the ones it was set with, or the browser keeps
|
||||
// the original alongside this one.
|
||||
if cleared.Path != "/" || !cleared.HttpOnly {
|
||||
t.Errorf("the cleared cookie has different attributes: %+v", cleared)
|
||||
}
|
||||
|
||||
// The row is gone, not merely expired.
|
||||
var rows int
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`,
|
||||
auth.HashToken(a.cookie.Value)).Scan(&rows); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
if rows != 0 {
|
||||
t.Error("logout left the session row in the database")
|
||||
}
|
||||
|
||||
// The old cookie no longer authenticates anything.
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("GET /me after logout = %d, want 401", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// 11. Every shape of logout succeeds: twice over, with a stale cookie, and with
|
||||
// no cookie at all. A user asking not to be signed in is not signed in
|
||||
// afterwards in all three cases, so all three are successes.
|
||||
func TestLogoutIsIdempotent(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
stale := a.cookie
|
||||
|
||||
for i, attempt := range []string{"first", "second", "third"} {
|
||||
req := httptest.NewRequest("POST", "/api/v1/auth/logout", nil)
|
||||
req.AddCookie(stale)
|
||||
rec := httptest.NewRecorder()
|
||||
a.handler.ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("%s logout (attempt %d) = %d, want 200", attempt, i+1, rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// With no cookie at all — a signed-out browser clicking sign out.
|
||||
if got := a.doAnon("POST", "/api/v1/auth/logout", nil); got.code != http.StatusOK {
|
||||
t.Errorf("logout with no cookie = %d, want 200", got.code)
|
||||
}
|
||||
// And with a token that was never real.
|
||||
junk := &http.Cookie{Name: sessionCookie, Value: "not-a-real-token"}
|
||||
if got := a.doWith(junk, "POST", "/api/v1/auth/logout"); got.code != http.StatusOK {
|
||||
t.Errorf("logout with a junk cookie = %d, want 200", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 12-15, 20. The middleware ──────────────────────────────────────────── */
|
||||
|
||||
// 20. Without a session, protected endpoints refuse.
|
||||
func TestUnauthenticatedRequestsAreRefused(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
for _, path := range []string{
|
||||
"/api/v1/me",
|
||||
"/api/v1/me/preferences",
|
||||
"/api/v1/job-postings",
|
||||
"/api/v1/job-applications",
|
||||
"/api/v1/worker-profiles",
|
||||
"/api/v1/courses",
|
||||
"/api/v1/staff",
|
||||
} {
|
||||
got := a.doAnon("GET", path, nil)
|
||||
if got.code != http.StatusUnauthorized {
|
||||
t.Errorf("GET %s without a session = %d, want 401", path, got.code)
|
||||
}
|
||||
if body, _ := got.body["error"].(map[string]any); body == nil || body["code"] != "unauthorized" {
|
||||
t.Errorf("GET %s: error code = %v, want unauthorized", path, got.body)
|
||||
}
|
||||
}
|
||||
|
||||
// Writes too, not only reads.
|
||||
if got := a.doAnon("POST", "/api/v1/job-postings", map[string]any{"title": "x"}); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("POST without a session = %d, want 401", got.code)
|
||||
}
|
||||
if got := a.doAnon("PATCH", "/api/v1/me", map[string]any{"full_name": "x"}); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("PATCH /me without a session = %d, want 401", got.code)
|
||||
}
|
||||
|
||||
// The public three stay public.
|
||||
if got := a.doAnon("GET", "/health", nil); got.code != http.StatusOK {
|
||||
t.Errorf("GET /health without a session = %d, want 200", got.code)
|
||||
}
|
||||
if got := a.doAnon("POST", "/api/v1/auth/logout", nil); got.code != http.StatusOK {
|
||||
t.Errorf("POST /auth/logout without a session = %d, want 200", got.code)
|
||||
}
|
||||
if got := a.doAnon("POST", "/api/v1/auth/login", map[string]any{
|
||||
"email": a.email, "password": harnessPassword,
|
||||
}); got.code != http.StatusOK {
|
||||
t.Errorf("POST /auth/login without a session = %d, want 200", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// 12. A token that does not name a session is refused, whatever it looks like.
|
||||
func TestInvalidSessionsAreRejected(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
unknown, err := auth.GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
for name, value := range map[string]string{
|
||||
"well-formed but unknown": unknown,
|
||||
"junk": "not-a-token-at-all",
|
||||
"empty": "",
|
||||
"the stored hash": auth.HashToken(a.cookie.Value),
|
||||
"the token, altered": a.cookie.Value[:len(a.cookie.Value)-1] + "X",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
got := a.doWith(&http.Cookie{Name: sessionCookie, Value: value}, "GET", "/api/v1/me")
|
||||
if got.code != http.StatusUnauthorized {
|
||||
t.Errorf("GET /me with a %s token = %d, want 401", name, got.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Presenting the *hash* must not work. It is the value in the database, so
|
||||
// a lookup that forgot to hash the cookie would accept it — and a leaked
|
||||
// database dump would then be a set of working credentials.
|
||||
got := a.doWith(&http.Cookie{Name: sessionCookie, Value: auth.HashToken(a.cookie.Value)}, "GET", "/api/v1/me")
|
||||
if got.code == http.StatusOK {
|
||||
t.Fatal("the stored token hash authenticated as a token")
|
||||
}
|
||||
}
|
||||
|
||||
// 13. An expired session is refused, and the row is cleaned up as it is found.
|
||||
func TestExpiredSessionIsRejected(t *testing.T) {
|
||||
a, now := clockedAPI(t)
|
||||
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me while live = %d, want 200", got.code)
|
||||
}
|
||||
|
||||
// Past the idle deadline without using it.
|
||||
*now = now.Add(shortSessions.IdleLifetime + time.Minute)
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("GET /me after expiry = %d, want 401", got.code)
|
||||
}
|
||||
|
||||
var rows int
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`,
|
||||
auth.HashToken(a.cookie.Value)).Scan(&rows); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
if rows != 0 {
|
||||
t.Error("an expired session was refused but left in the database")
|
||||
}
|
||||
}
|
||||
|
||||
// A session used steadily slides forward and keeps working — but never past its
|
||||
// absolute ceiling.
|
||||
func TestSessionSlidesButNotForever(t *testing.T) {
|
||||
a, now := clockedAPI(t)
|
||||
start := *now
|
||||
|
||||
for _, at := range []time.Duration{50 * time.Minute, 105 * time.Minute, 160 * time.Minute} {
|
||||
*now = start.Add(at)
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me at +%v = %d, want 200 — the session should have slid", at, got.code)
|
||||
}
|
||||
}
|
||||
|
||||
*now = start.Add(shortSessions.AbsoluteLifetime)
|
||||
if got := a.do("GET", "/api/v1/me", nil); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("GET /me at the absolute ceiling = %d, want 401", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// 14. The identity is the session's user, not the first or the oldest one.
|
||||
func TestAuthenticatedRequestResolvesTheSessionUser(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
other := newUser(t, a.h.Pool, a.orgID, "second-user@example.test", "employer")
|
||||
result := signIn(t, a.handler, "second-user@example.test", harnessPassword, false)
|
||||
if result.code != http.StatusOK {
|
||||
t.Fatalf("second user login = %d", result.code)
|
||||
}
|
||||
|
||||
got := a.doWith(result.cookie, "GET", "/api/v1/me")
|
||||
if got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me = %d", got.code)
|
||||
}
|
||||
data := got.body["data"].(map[string]any)
|
||||
if data["id"] != other {
|
||||
t.Errorf("GET /me returned %v, want the second user %v", data["id"], other)
|
||||
}
|
||||
if data["email"] != "second-user@example.test" {
|
||||
t.Errorf("GET /me returned email %v", data["email"])
|
||||
}
|
||||
|
||||
// The seeded user's own cookie still resolves to the seeded user: two
|
||||
// sessions, two identities, no crosstalk.
|
||||
first := a.do("GET", "/api/v1/me", nil)
|
||||
if first.body["data"].(map[string]any)["id"] != a.userID {
|
||||
t.Error("the first session no longer resolves to its own user")
|
||||
}
|
||||
}
|
||||
|
||||
// 15. Suspending an account takes effect on its next request, and takes the
|
||||
// session with it.
|
||||
func TestSuspendedUserIsRejectedMidSession(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
victim := newUser(t, a.h.Pool, a.orgID, "about-to-be-suspended@example.test", "employer")
|
||||
result := signIn(t, a.handler, "about-to-be-suspended@example.test", harnessPassword, false)
|
||||
if result.code != http.StatusOK {
|
||||
t.Fatalf("login = %d", result.code)
|
||||
}
|
||||
if got := a.doWith(result.cookie, "GET", "/api/v1/me"); got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me while active = %d, want 200", got.code)
|
||||
}
|
||||
|
||||
setStatus(t, a.h.Pool, victim, "suspended")
|
||||
|
||||
if got := a.doWith(result.cookie, "GET", "/api/v1/me"); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("GET /me after suspension = %d, want 401", got.code)
|
||||
}
|
||||
// The session is destroyed rather than refused one request at a time: a
|
||||
// suspended account must not keep a working cookie for thirty days.
|
||||
var rows int
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*)::int FROM sessions WHERE user_id = $1::uuid`, victim).Scan(&rows); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
if rows != 0 {
|
||||
t.Errorf("%d sessions survive for a suspended user, want 0", rows)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 16-19. /me ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// 16. GET /me is the session's user.
|
||||
func TestMeReturnsTheSessionUser(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
got := a.do("GET", "/api/v1/me", nil)
|
||||
if got.code != http.StatusOK {
|
||||
t.Fatalf("GET /me = %d", got.code)
|
||||
}
|
||||
data := got.body["data"].(map[string]any)
|
||||
if data["id"] != a.userID {
|
||||
t.Errorf("GET /me id = %v, want %v", data["id"], a.userID)
|
||||
}
|
||||
// The frontend contract: these fields must still be here.
|
||||
for _, field := range []string{"id", "email", "full_name", "role", "account_type", "status", "preferences"} {
|
||||
if _, ok := data[field]; !ok {
|
||||
t.Errorf("GET /me no longer returns %q", field)
|
||||
}
|
||||
}
|
||||
// And these must not be.
|
||||
for _, field := range []string{"password_hash", "org_id"} {
|
||||
if _, ok := data[field]; ok {
|
||||
t.Errorf("GET /me exposes %q", field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 17, 18, 19. A user cannot edit what the server owns about them.
|
||||
func TestMeCannotEditServerOwnedFields(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
before := readUserRow(t, a)
|
||||
|
||||
// Everything at once, plus each on its own below, because a handler could
|
||||
// plausibly filter one and not another.
|
||||
got := a.do("PATCH", "/api/v1/me", map[string]any{
|
||||
"role": "admin",
|
||||
"org_id": "00000000-0000-0000-0000-000000000000",
|
||||
"password_hash": "$argon2id$v=19$m=65536,t=3,p=4$YWFhYWFhYWFhYWFhYWFhYQ$" + strings.Repeat("A", 43),
|
||||
"id": "00000000-0000-0000-0000-000000000000",
|
||||
"status": "suspended",
|
||||
"email": "attacker@example.invalid",
|
||||
"full_name": "A Legitimate Rename",
|
||||
})
|
||||
if got.code != http.StatusOK {
|
||||
t.Fatalf("PATCH /me = %d (%v)", got.code, got.body)
|
||||
}
|
||||
|
||||
after := readUserRow(t, a)
|
||||
if after.role != before.role {
|
||||
t.Errorf("role changed from %q to %q — privilege escalation", before.role, after.role)
|
||||
}
|
||||
if after.orgID != before.orgID {
|
||||
t.Errorf("org_id changed from %q to %q — tenancy escape", before.orgID, after.orgID)
|
||||
}
|
||||
if after.passwordHash != before.passwordHash {
|
||||
t.Error("password_hash was overwritten through PATCH /me")
|
||||
}
|
||||
if after.id != before.id {
|
||||
t.Errorf("id changed from %q to %q", before.id, after.id)
|
||||
}
|
||||
if after.status != before.status {
|
||||
t.Errorf("status changed from %q to %q", before.status, after.status)
|
||||
}
|
||||
if after.email != before.email {
|
||||
t.Errorf("email changed from %q to %q", before.email, after.email)
|
||||
}
|
||||
// The one legitimate field in that payload did land, so the endpoint is
|
||||
// filtering rather than refusing everything.
|
||||
if after.fullName != "A Legitimate Rename" {
|
||||
t.Errorf("full_name = %q, want the rename to have applied", after.fullName)
|
||||
}
|
||||
// The response reports the truth rather than echoing the request.
|
||||
if data := got.body["data"].(map[string]any); data["role"] != before.role {
|
||||
t.Errorf("the response reports role %v, want the unchanged %q", data["role"], before.role)
|
||||
}
|
||||
|
||||
// One at a time.
|
||||
for field, value := range map[string]any{
|
||||
"role": "admin",
|
||||
"org_id": "00000000-0000-0000-0000-000000000000",
|
||||
"password_hash": "anything",
|
||||
"status": "suspended",
|
||||
} {
|
||||
t.Run(field, func(t *testing.T) {
|
||||
r := a.do("PATCH", "/api/v1/me", map[string]any{field: value})
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("PATCH /me {%s} = %d", field, r.code)
|
||||
}
|
||||
now := readUserRow(t, a)
|
||||
if now.role != before.role || now.orgID != before.orgID ||
|
||||
now.passwordHash != before.passwordHash || now.status != before.status {
|
||||
t.Errorf("PATCH /me {%s: %v} changed a server-owned field", field, value)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// account_type stays editable: it is a display attribute, not authorization,
|
||||
// and Layout.jsx writes it when the viewer switches surface.
|
||||
if r := a.do("PATCH", "/api/v1/me", map[string]any{"account_type": "talent"}); r.code != http.StatusOK {
|
||||
t.Fatalf("PATCH /me {account_type} = %d", r.code)
|
||||
}
|
||||
if readUserRow(t, a).accountType != "talent" {
|
||||
t.Error("account_type is no longer self-editable; Layout.jsx's role switch depends on it")
|
||||
}
|
||||
}
|
||||
|
||||
type userRow struct {
|
||||
id, orgID, email, fullName, role, accountType, status, passwordHash string
|
||||
}
|
||||
|
||||
func readUserRow(t *testing.T, a *api) userRow {
|
||||
t.Helper()
|
||||
var u userRow
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT id::text, org_id::text, email::text, full_name, role, account_type, status,
|
||||
COALESCE(password_hash, '') FROM users WHERE id = $1::uuid`, a.userID).
|
||||
Scan(&u.id, &u.orgID, &u.email, &u.fullName, &u.role, &u.accountType, &u.status, &u.passwordHash); err != nil {
|
||||
t.Fatalf("read the user row: %v", err)
|
||||
}
|
||||
return u
|
||||
}
|
||||
|
||||
/* ── Identity comes from the session, never from the request ────────────── */
|
||||
|
||||
// The security property the whole phase exists for: nothing a client writes can
|
||||
// change who it is.
|
||||
func TestIdentityCannotBeSuppliedByTheRequest(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
victim := newUser(t, a.h.Pool, a.orgID, "victim@example.test", "admin")
|
||||
|
||||
// A body naming another user.
|
||||
got := a.do("PATCH", "/api/v1/me", map[string]any{
|
||||
"id": victim, "user_id": victim, "full_name": "Renamed By An Impostor",
|
||||
})
|
||||
if got.code != http.StatusOK {
|
||||
t.Fatalf("PATCH /me = %d", got.code)
|
||||
}
|
||||
if got.body["data"].(map[string]any)["id"] != a.userID {
|
||||
t.Error("a user id in the body changed whose record was returned")
|
||||
}
|
||||
var victimName string
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT full_name FROM users WHERE id = $1::uuid`, victim).Scan(&victimName); err != nil {
|
||||
t.Fatalf("read the victim: %v", err)
|
||||
}
|
||||
if victimName == "Renamed By An Impostor" {
|
||||
t.Fatal("a user id in the request body redirected the write to another user")
|
||||
}
|
||||
|
||||
// A query string naming another user, and another organization.
|
||||
for _, q := range []string{
|
||||
"?user_id=" + victim,
|
||||
"?org_id=00000000-0000-0000-0000-000000000000",
|
||||
"?id=" + victim,
|
||||
} {
|
||||
r := a.do("GET", "/api/v1/me"+q, nil)
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("GET /me%s = %d", q, r.code)
|
||||
}
|
||||
if r.body["data"].(map[string]any)["id"] != a.userID {
|
||||
t.Errorf("GET /me%s resolved to a different user", q)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The organization is read from the user's row, so a session in one tenant sees
|
||||
// no data from another. This is what replaced the fixed development org.
|
||||
func TestTenancyFollowsTheSession(t *testing.T) {
|
||||
a := newAPI(t)
|
||||
|
||||
// The seeded organization has data.
|
||||
mine := a.do("GET", "/api/v1/job-postings?limit=100", nil)
|
||||
if mine.code != http.StatusOK {
|
||||
t.Fatalf("GET /job-postings = %d", mine.code)
|
||||
}
|
||||
if len(mine.records(t)) == 0 {
|
||||
t.Fatal("the seeded organization has no job postings; the test proves nothing")
|
||||
}
|
||||
|
||||
// A second organization, with a user of its own and no data.
|
||||
var otherOrg string
|
||||
if err := a.h.Pool.QueryRow(context.Background(),
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Other Tenant', 'other-tenant') RETURNING id::text`).
|
||||
Scan(&otherOrg); err != nil {
|
||||
t.Fatalf("create the second organization: %v", err)
|
||||
}
|
||||
newUser(t, a.h.Pool, otherOrg, "outsider@example.test", "admin")
|
||||
|
||||
result := signIn(t, a.handler, "outsider@example.test", harnessPassword, false)
|
||||
if result.code != http.StatusOK {
|
||||
t.Fatalf("outsider login = %d", result.code)
|
||||
}
|
||||
theirs := a.doWith(result.cookie, "GET", "/api/v1/job-postings?limit=100")
|
||||
if theirs.code != http.StatusOK {
|
||||
t.Fatalf("GET /job-postings as the outsider = %d", theirs.code)
|
||||
}
|
||||
if n := len(theirs.records(t)); n != 0 {
|
||||
t.Errorf("a user in another organization sees %d job postings, want 0", n)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 21. Sweeping ───────────────────────────────────────────────────────── */
|
||||
|
||||
// 21. The sweep collects sessions nobody comes back for, and leaves live ones.
|
||||
func TestSessionSweep(t *testing.T) {
|
||||
a, now := clockedAPI(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// a.cookie is a normal session (1h idle here). Add a Remember Me one.
|
||||
long := signIn(t, a.handler, a.email, harnessPassword, true)
|
||||
if long.code != http.StatusOK {
|
||||
t.Fatalf("remember-me login = %d", long.code)
|
||||
}
|
||||
if n := countSessions(t, a); n != 2 {
|
||||
t.Fatalf("%d sessions before the sweep, want 2", n)
|
||||
}
|
||||
|
||||
// Nothing is due yet.
|
||||
if deleted, err := a.srv.Sessions().Sweep(ctx); err != nil || deleted != 0 {
|
||||
t.Errorf("early sweep deleted %d (err %v), want 0", deleted, err)
|
||||
}
|
||||
|
||||
// Past the short session's deadline, not the long one's.
|
||||
*now = now.Add(shortSessions.IdleLifetime + time.Minute)
|
||||
deleted, err := a.srv.Sessions().Sweep(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Sweep: %v", err)
|
||||
}
|
||||
if deleted != 1 {
|
||||
t.Errorf("sweep deleted %d sessions, want 1", deleted)
|
||||
}
|
||||
if n := countSessions(t, a); n != 1 {
|
||||
t.Errorf("%d sessions remain, want 1", n)
|
||||
}
|
||||
// The survivor still works.
|
||||
if got := a.doWith(long.cookie, "GET", "/api/v1/me"); got.code != http.StatusOK {
|
||||
t.Errorf("the swept database rejected a live session: %d", got.code)
|
||||
}
|
||||
|
||||
// Past everything.
|
||||
*now = now.Add(shortSessions.RememberAbsoluteLifetime)
|
||||
if _, err := a.srv.Sessions().Sweep(ctx); err != nil {
|
||||
t.Fatalf("Sweep: %v", err)
|
||||
}
|
||||
if n := countSessions(t, a); n != 0 {
|
||||
t.Errorf("%d sessions remain after everything expired, want 0", n)
|
||||
}
|
||||
}
|
||||
|
||||
func countSessions(t *testing.T, a *api) int {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := a.h.Pool.QueryRow(context.Background(), `SELECT count(*)::int FROM sessions`).Scan(&n); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
/* ── 22. Rate limiting ──────────────────────────────────────────────────── */
|
||||
|
||||
// 22. Repeated failures are refused, before any password is checked.
|
||||
func TestLoginRateLimit(t *testing.T) {
|
||||
a, now := clockedAPI(t, httpserver.WithLoginRateLimit(3, 50, 10*time.Minute))
|
||||
start := *now
|
||||
|
||||
// Three failures use the budget.
|
||||
for i := 1; i <= 3; i++ {
|
||||
if got := signIn(t, a.handler, a.email, "wrong-password", false); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("failure %d = %d, want 401", i, got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// The fourth is refused as rate limited, not as a bad password.
|
||||
blocked := signIn(t, a.handler, a.email, "wrong-password", false)
|
||||
if blocked.code != http.StatusTooManyRequests {
|
||||
t.Fatalf("the fourth attempt = %d, want 429", blocked.code)
|
||||
}
|
||||
if got := blocked.raw.Header().Get("Retry-After"); got == "" || got == "0" {
|
||||
t.Errorf("Retry-After = %q, want a positive number of seconds", got)
|
||||
}
|
||||
if body, _ := blocked.body["error"].(map[string]any); body == nil || body["code"] != "rate_limited" {
|
||||
t.Errorf("error code = %v, want rate_limited", blocked.body)
|
||||
}
|
||||
|
||||
// The CORRECT password is refused too. The limit is checked before the
|
||||
// credentials, which is the point: an attacker must not be able to make the
|
||||
// server hash for them, and must not learn from a 401-vs-429 difference
|
||||
// whether their guess was right.
|
||||
correct := signIn(t, a.handler, a.email, harnessPassword, false)
|
||||
if correct.code != http.StatusTooManyRequests {
|
||||
t.Errorf("a correct password during a block = %d, want 429", correct.code)
|
||||
}
|
||||
if correct.cookie != nil {
|
||||
t.Error("a rate-limited login still issued a session")
|
||||
}
|
||||
|
||||
// The window passes and the budget returns.
|
||||
*now = start.Add(11 * time.Minute)
|
||||
if got := signIn(t, a.handler, a.email, harnessPassword, false); got.code != http.StatusOK {
|
||||
t.Fatalf("login after the window = %d, want 200", got.code)
|
||||
}
|
||||
|
||||
// A success clears the email's counter, so two typos then a success leaves
|
||||
// nothing behind.
|
||||
for i := 0; i < 2; i++ {
|
||||
if got := signIn(t, a.handler, a.email, "wrong-password", false); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("typo %d = %d, want 401", i+1, got.code)
|
||||
}
|
||||
}
|
||||
if got := signIn(t, a.handler, a.email, harnessPassword, false); got.code != http.StatusOK {
|
||||
t.Fatalf("login after two typos = %d, want 200", got.code)
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
if got := signIn(t, a.handler, a.email, "wrong-password", false); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("after the reset, failure %d = %d, want 401 — the counter did not clear", i+1, got.code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The limit is per email as well as per address, so one account cannot be
|
||||
// ground down by an attacker who has plenty of addresses — and one address
|
||||
// cannot work through plenty of accounts.
|
||||
func TestLoginRateLimitIsPerEmailAndPerAddress(t *testing.T) {
|
||||
// Three per email, five per address: tight enough to reach both bounds in a
|
||||
// handful of attempts, and shaped like the production pair, where the
|
||||
// address budget is the wider one.
|
||||
a, _ := clockedAPI(t, httpserver.WithLoginRateLimit(3, 5, 10*time.Minute))
|
||||
newUser(t, a.h.Pool, a.orgID, "unrelated@example.test", "employer")
|
||||
|
||||
// Exhaust the seeded account's own budget.
|
||||
for i := 0; i < 3; i++ {
|
||||
if got := signIn(t, a.handler, a.email, "wrong-password", false); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("failure %d = %d, want 401", i+1, got.code)
|
||||
}
|
||||
}
|
||||
if got := signIn(t, a.handler, a.email, harnessPassword, false); got.code != http.StatusTooManyRequests {
|
||||
t.Fatalf("the blocked email = %d, want 429", got.code)
|
||||
}
|
||||
|
||||
// A different account from the same address still works: the per-email
|
||||
// budget is per email, so one account being attacked does not lock out
|
||||
// everyone else behind the same NAT.
|
||||
if got := signIn(t, a.handler, "unrelated@example.test", harnessPassword, false); got.code != http.StatusOK {
|
||||
t.Fatalf("a second email from the same address = %d, want 200", got.code)
|
||||
}
|
||||
|
||||
// But the address budget is real. Two more failures from here reach five,
|
||||
// and then nothing from this address gets through, whichever account it
|
||||
// names.
|
||||
for i := 0; i < 2; i++ {
|
||||
if got := signIn(t, a.handler, "unrelated@example.test", "wrong-password", false); got.code != http.StatusUnauthorized {
|
||||
t.Fatalf("address failure %d = %d, want 401", i+1, got.code)
|
||||
}
|
||||
}
|
||||
if got := signIn(t, a.handler, "unrelated@example.test", harnessPassword, false); got.code != http.StatusTooManyRequests {
|
||||
t.Errorf("the address budget was not enforced: %d, want 429", got.code)
|
||||
}
|
||||
}
|
||||
218
go-api/internal/httpserver/authfixture_test.go
Normal file
218
go-api/internal/httpserver/authfixture_test.go
Normal file
@@ -0,0 +1,218 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The shared authentication fixture.
|
||||
//
|
||||
// Every endpoint in this package except /health and the two auth routes now
|
||||
// requires a session, so the harness signs in before it hands a test anything.
|
||||
// That is what keeps the thirty-odd pre-existing tests in api_test.go working
|
||||
// unchanged: they still call a.do("GET", "/api/v1/…"), and the cookie rides
|
||||
// along underneath.
|
||||
//
|
||||
// The alternative — inserting a session row directly — would test the
|
||||
// middleware against a session no login ever produced. Signing in through the
|
||||
// real handler means the fixture itself exercises the flow it depends on.
|
||||
|
||||
// harnessPassword is the password every test account is given. It is a literal
|
||||
// in a test file for a database that is created and dropped by the same
|
||||
// process; it is not a credential for anything that outlives the run.
|
||||
const harnessPassword = "harness-password-not-a-real-secret"
|
||||
|
||||
// harnessHash is argon2id at production cost — about a tenth of a second — so
|
||||
// it is computed once for the whole package rather than once per test.
|
||||
var harnessHash = sync.OnceValues(func() (string, error) {
|
||||
return auth.HashPassword(harnessPassword)
|
||||
})
|
||||
|
||||
// setPassword gives a user a known password.
|
||||
func setPassword(t *testing.T, pool *pgxpool.Pool, userID string) {
|
||||
t.Helper()
|
||||
hash, err := harnessHash()
|
||||
if err != nil {
|
||||
t.Fatalf("hash the harness password: %v", err)
|
||||
}
|
||||
if _, err := pool.Exec(context.Background(),
|
||||
`UPDATE users SET password_hash = $2::text WHERE id = $1::uuid`, userID, hash); err != nil {
|
||||
t.Fatalf("set the harness password: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// setStatus flips a user between 'active' and 'suspended'.
|
||||
func setStatus(t *testing.T, pool *pgxpool.Pool, userID, status string) {
|
||||
t.Helper()
|
||||
if _, err := pool.Exec(context.Background(),
|
||||
`UPDATE users SET status = $2::text WHERE id = $1::uuid`, userID, status); err != nil {
|
||||
t.Fatalf("set status %s: %v", status, err)
|
||||
}
|
||||
}
|
||||
|
||||
// seededUser is the demo user the fixture loads into the test database.
|
||||
func seededUser(t *testing.T, pool *pgxpool.Pool) (id, email string) {
|
||||
t.Helper()
|
||||
if err := pool.QueryRow(context.Background(),
|
||||
`SELECT id::text, email::text FROM users ORDER BY created_date, id LIMIT 1`).
|
||||
Scan(&id, &email); err != nil {
|
||||
t.Fatalf("read the seeded user: %v", err)
|
||||
}
|
||||
return id, email
|
||||
}
|
||||
|
||||
// newUser adds a user to an organization, with the harness password set.
|
||||
func newUser(t *testing.T, pool *pgxpool.Pool, orgID, email, role string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := pool.QueryRow(context.Background(),
|
||||
`INSERT INTO users (org_id, email, full_name, role) VALUES ($1::uuid, $2::citext, $3, $4)
|
||||
RETURNING id::text`, orgID, email, "Test User", role).Scan(&id); err != nil {
|
||||
t.Fatalf("create user %s: %v", email, err)
|
||||
}
|
||||
setPassword(t, pool, id)
|
||||
return id
|
||||
}
|
||||
|
||||
// loginResult is what signIn observed: the response, and the cookie if one was
|
||||
// set. Tests assert on both.
|
||||
type loginResult struct {
|
||||
code int
|
||||
body map[string]any
|
||||
cookie *http.Cookie
|
||||
raw *httptest.ResponseRecorder
|
||||
}
|
||||
|
||||
// signIn posts credentials to the real login handler.
|
||||
func signIn(t *testing.T, handler http.Handler, email, password string, remember bool) loginResult {
|
||||
t.Helper()
|
||||
payload, err := json.Marshal(map[string]any{
|
||||
"email": email, "password": password, "remember_me": remember,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("encode the login payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest("POST", "/api/v1/auth/login", strings.NewReader(string(payload)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
out := loginResult{code: rec.Code, raw: rec}
|
||||
if rec.Body.Len() > 0 {
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &out.body)
|
||||
}
|
||||
for _, c := range rec.Result().Cookies() {
|
||||
if c.Name == sessionCookie {
|
||||
out.cookie = c
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// sessionCookie is the name the server uses. Duplicated here rather than
|
||||
// exported from the package: a test that asserts the cookie name should fail
|
||||
// when the name changes, not silently follow it.
|
||||
const sessionCookie = "krow_session"
|
||||
|
||||
// newServer builds a server over a fresh migrated, seeded database.
|
||||
func newServer(t *testing.T, h *testutil.Harness, origins []string, opts ...httpserver.Option) *httpserver.Server {
|
||||
t.Helper()
|
||||
cfg := &config.Config{
|
||||
AppEnv: "development",
|
||||
HTTP: config.HTTPConfig{
|
||||
Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second,
|
||||
CORSOrigins: origins,
|
||||
},
|
||||
DB: config.DBConfig{Schema: "public"},
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"}, log, opts...)
|
||||
if err != nil {
|
||||
t.Fatalf("build the server: %v", err)
|
||||
}
|
||||
return srv
|
||||
}
|
||||
|
||||
// withSession attaches a cookie to every request passing through, so a test
|
||||
// about something else — CORS, say — is not also a test about signing in.
|
||||
func withSession(handler http.Handler, cookie *http.Cookie) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if cookie != nil {
|
||||
r.AddCookie(cookie)
|
||||
}
|
||||
handler.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// newServerWithEnv builds a server for a given APP_ENV, so the cookie's Secure
|
||||
// flag can be observed on both sides of the development boundary.
|
||||
func newServerWithEnv(t *testing.T, h *testutil.Harness, appEnv string) http.Handler {
|
||||
t.Helper()
|
||||
cfg := &config.Config{
|
||||
AppEnv: appEnv,
|
||||
HTTP: config.HTTPConfig{Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second},
|
||||
DB: config.DBConfig{Schema: "public"},
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"}, log)
|
||||
if err != nil {
|
||||
t.Fatalf("build the %s server: %v", appEnv, err)
|
||||
}
|
||||
return srv.Handler()
|
||||
}
|
||||
|
||||
/* ── Role fixtures (Phase 3D) ───────────────────────────────────────────── */
|
||||
|
||||
// actor is one signed-in user of a known role.
|
||||
type actor struct {
|
||||
name string // for test output only
|
||||
id string
|
||||
email string
|
||||
role string
|
||||
cookie *http.Cookie
|
||||
}
|
||||
|
||||
// signInAs creates a user with the given role and signs them in.
|
||||
func signInAs(t *testing.T, handler http.Handler, pool *pgxpool.Pool, orgID, name, email, role string) actor {
|
||||
t.Helper()
|
||||
id := newUserWithRole(t, pool, orgID, email, role)
|
||||
result := signIn(t, handler, email, harnessPassword, false)
|
||||
if result.code != http.StatusOK || result.cookie == nil {
|
||||
t.Fatalf("could not sign in %s (%s): status %d", name, role, result.code)
|
||||
}
|
||||
return actor{name: name, id: id, email: email, role: role, cookie: result.cookie}
|
||||
}
|
||||
|
||||
// newUserWithRole inserts a user with an explicit role and the harness password.
|
||||
//
|
||||
// Written straight to the database rather than through the API on purpose:
|
||||
// users.role is server-owned and there is deliberately no endpoint that sets
|
||||
// it, which is the property Phase 3D depends on.
|
||||
func newUserWithRole(t *testing.T, pool *pgxpool.Pool, orgID, email, role string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := pool.QueryRow(context.Background(),
|
||||
`INSERT INTO users (org_id, email, full_name, role, account_type)
|
||||
VALUES ($1::uuid, $2::citext, $3, $4::text, 'employer') RETURNING id::text`,
|
||||
orgID, email, "Test "+role, role).Scan(&id); err != nil {
|
||||
t.Fatalf("create %s user %s: %v", role, email, err)
|
||||
}
|
||||
setPassword(t, pool, id)
|
||||
return id
|
||||
}
|
||||
104
go-api/internal/httpserver/cors.go
Normal file
104
go-api/internal/httpserver/cors.go
Normal file
@@ -0,0 +1,104 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Cross-origin access, for local development.
|
||||
//
|
||||
// In Phase 2D the frontend fetches this API directly from the Vite dev server,
|
||||
// which is a different origin (http://localhost:5173 → http://127.0.0.1:8080).
|
||||
// Without these headers the browser makes the request and then refuses to let
|
||||
// the page read the response, which surfaces in the app as an opaque "Failed to
|
||||
// fetch" with a perfectly healthy 200 in the server log.
|
||||
//
|
||||
// This is a transport concern only. No endpoint, request shape, response shape
|
||||
// or status code in docs/api-contract.md changes because of it.
|
||||
|
||||
// corsMaxAge is how long a browser may cache a preflight result. Ten minutes
|
||||
// keeps preflight off the hot path without making an allowlist change take an
|
||||
// awkwardly long time to be noticed in development.
|
||||
const corsMaxAge = 600
|
||||
|
||||
// allowedCORSMethods is every method the router actually registers, plus
|
||||
// OPTIONS for the preflight itself. It is a fixed list rather than something
|
||||
// derived per path: the browser asks about one method at a time and only needs
|
||||
// to know it is permitted in general.
|
||||
var allowedCORSMethods = []string{
|
||||
http.MethodGet, http.MethodPost, http.MethodPatch,
|
||||
http.MethodDelete, http.MethodOptions,
|
||||
}
|
||||
|
||||
// cors answers preflights and marks cross-origin responses as readable.
|
||||
//
|
||||
// Origins are matched exactly against the allowlist and echoed back one at a
|
||||
// time — never "*" — so adding credentials later does not require rewriting
|
||||
// this. A request whose Origin is not on the list is served normally, with no
|
||||
// CORS headers: the API does not refuse it, the browser simply will not hand
|
||||
// the response to the page. That distinction matters, because curl, the health
|
||||
// checker and any server-to-server caller send no Origin at all and must not be
|
||||
// affected by this middleware.
|
||||
//
|
||||
// With an empty allowlist the middleware is not installed at all (see New), so
|
||||
// the same-origin deployment pays nothing for it.
|
||||
func cors(origins []string) func(http.Handler) http.Handler {
|
||||
allowed := make(map[string]bool, len(origins))
|
||||
for _, o := range origins {
|
||||
allowed[o] = true
|
||||
}
|
||||
methods := strings.Join(allowedCORSMethods, ", ")
|
||||
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
origin := r.Header.Get("Origin")
|
||||
|
||||
// Vary on Origin whether or not this particular origin matched: the
|
||||
// response differs by Origin, so a cache that ignored it could hand
|
||||
// one origin's headers to another.
|
||||
w.Header().Add("Vary", "Origin")
|
||||
|
||||
if origin == "" || !allowed[origin] {
|
||||
if isPreflight(r) {
|
||||
// A preflight is never a real request. Answering it with
|
||||
// the router's 404 for "OPTIONS /api/v1/…" would be
|
||||
// misleading; 403 says plainly that the origin was refused.
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Access-Control-Allow-Origin", origin)
|
||||
|
||||
if isPreflight(r) {
|
||||
w.Header().Add("Vary", "Access-Control-Request-Method")
|
||||
w.Header().Add("Vary", "Access-Control-Request-Headers")
|
||||
w.Header().Set("Access-Control-Allow-Methods", methods)
|
||||
// Echo the requested headers rather than listing them. The
|
||||
// frontend sends only Content-Type today; echoing means a
|
||||
// future header does not need a change here to be allowed from
|
||||
// an origin that is already trusted.
|
||||
if h := r.Header.Get("Access-Control-Request-Headers"); h != "" {
|
||||
w.Header().Set("Access-Control-Allow-Headers", h)
|
||||
} else {
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type")
|
||||
}
|
||||
w.Header().Set("Access-Control-Max-Age", strconv.Itoa(corsMaxAge))
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// isPreflight identifies the browser's OPTIONS probe. A bare OPTIONS with no
|
||||
// Access-Control-Request-Method is not a preflight and is left to the router.
|
||||
func isPreflight(r *http.Request) bool {
|
||||
return r.Method == http.MethodOptions &&
|
||||
r.Header.Get("Access-Control-Request-Method") != ""
|
||||
}
|
||||
143
go-api/internal/httpserver/cors_test.go
Normal file
143
go-api/internal/httpserver/cors_test.go
Normal file
@@ -0,0 +1,143 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
const devOrigin = "http://localhost:5173"
|
||||
|
||||
// corsAPI is newAPI with an explicit CORS allowlist. It is separate because
|
||||
// every other test in this package asserts the same-origin behaviour, where the
|
||||
// middleware is not installed at all.
|
||||
func corsAPI(t *testing.T, origins ...string) http.Handler {
|
||||
t.Helper()
|
||||
h := testutil.New(t)
|
||||
srv := newServer(t, h, origins)
|
||||
handler := srv.Handler()
|
||||
|
||||
// The API routes below now require a session. Signing in once and attaching
|
||||
// the cookie to every request keeps these tests about CORS: without it they
|
||||
// would assert 401 and prove nothing about the headers.
|
||||
userID, email := seededUser(t, h.Pool)
|
||||
setPassword(t, h.Pool, userID)
|
||||
result := signIn(t, handler, email, harnessPassword, false)
|
||||
if result.code != http.StatusOK || result.cookie == nil {
|
||||
t.Fatalf("the CORS harness could not sign in: status %d", result.code)
|
||||
}
|
||||
return withSession(handler, result.cookie)
|
||||
}
|
||||
|
||||
func send(handler http.Handler, method, path string, headers map[string]string) *httptest.ResponseRecorder {
|
||||
req := httptest.NewRequest(method, path, nil)
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// An allowed origin gets its own origin echoed back, never "*".
|
||||
func TestCORSAllowsConfiguredOrigin(t *testing.T) {
|
||||
handler := corsAPI(t, devOrigin)
|
||||
|
||||
rec := send(handler, "GET", "/api/v1/job-postings", map[string]string{"Origin": devOrigin})
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d", rec.Code)
|
||||
}
|
||||
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != devOrigin {
|
||||
t.Fatalf("Access-Control-Allow-Origin = %q, want %q", got, devOrigin)
|
||||
}
|
||||
if rec.Header().Get("Vary") == "" {
|
||||
t.Fatal("a response that varies by Origin must say so")
|
||||
}
|
||||
}
|
||||
|
||||
// The preflight the browser sends before a PATCH must succeed without reaching
|
||||
// the router, and must name the methods the frontend uses.
|
||||
func TestCORSPreflight(t *testing.T) {
|
||||
handler := corsAPI(t, devOrigin)
|
||||
|
||||
rec := send(handler, "OPTIONS", "/api/v1/job-applications/some-id", map[string]string{
|
||||
"Origin": devOrigin,
|
||||
"Access-Control-Request-Method": "PATCH",
|
||||
"Access-Control-Request-Headers": "content-type",
|
||||
})
|
||||
if rec.Code != http.StatusNoContent {
|
||||
t.Fatalf("preflight: expected 204, got %d (%s)", rec.Code, rec.Body.String())
|
||||
}
|
||||
allow := rec.Header().Get("Access-Control-Allow-Methods")
|
||||
for _, m := range []string{"GET", "POST", "PATCH", "DELETE"} {
|
||||
if !contains(allow, m) {
|
||||
t.Fatalf("Access-Control-Allow-Methods = %q, missing %s", allow, m)
|
||||
}
|
||||
}
|
||||
if got := rec.Header().Get("Access-Control-Allow-Headers"); got != "content-type" {
|
||||
t.Fatalf("Access-Control-Allow-Headers = %q, want the requested header echoed", got)
|
||||
}
|
||||
if rec.Header().Get("Access-Control-Max-Age") == "" {
|
||||
t.Fatal("preflight result should be cacheable")
|
||||
}
|
||||
}
|
||||
|
||||
// An origin that is not on the list gets no CORS headers, so the browser will
|
||||
// not hand the response to the page.
|
||||
func TestCORSRefusesUnknownOrigin(t *testing.T) {
|
||||
handler := corsAPI(t, devOrigin)
|
||||
|
||||
rec := send(handler, "GET", "/api/v1/job-postings", map[string]string{
|
||||
"Origin": "http://evil.example",
|
||||
})
|
||||
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
|
||||
t.Fatalf("an unlisted origin was allowed: %q", got)
|
||||
}
|
||||
|
||||
pre := send(handler, "OPTIONS", "/api/v1/job-postings", map[string]string{
|
||||
"Origin": "http://evil.example",
|
||||
"Access-Control-Request-Method": "GET",
|
||||
})
|
||||
if pre.Code != http.StatusForbidden {
|
||||
t.Fatalf("preflight from an unlisted origin: expected 403, got %d", pre.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// A caller with no Origin — curl, a health checker, anything server-to-server —
|
||||
// is untouched by the middleware.
|
||||
func TestCORSIgnoresRequestsWithoutOrigin(t *testing.T) {
|
||||
handler := corsAPI(t, devOrigin)
|
||||
|
||||
rec := send(handler, "GET", "/health", nil)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d", rec.Code)
|
||||
}
|
||||
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
|
||||
t.Fatalf("a request with no Origin got CORS headers: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// With no allowlist the middleware is not installed, which is the posture for
|
||||
// any deployment serving the frontend from the API's own origin.
|
||||
func TestCORSOffByDefault(t *testing.T) {
|
||||
handler := corsAPI(t) // no origins
|
||||
|
||||
rec := send(handler, "GET", "/api/v1/job-postings", map[string]string{"Origin": devOrigin})
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200, got %d", rec.Code)
|
||||
}
|
||||
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
|
||||
t.Fatalf("CORS answered with no allowlist configured: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func contains(haystack, needle string) bool {
|
||||
for i := 0; i+len(needle) <= len(haystack); i++ {
|
||||
if haystack[i:i+len(needle)] == needle {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
214
go-api/internal/httpserver/definitions.go
Normal file
214
go-api/internal/httpserver/definitions.go
Normal file
@@ -0,0 +1,214 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
)
|
||||
|
||||
func (s *Server) routeDefinitions(mux *http.ServeMux) int {
|
||||
mux.HandleFunc("GET /api/v1/agent-definitions", s.handleAgentDefinitionsList)
|
||||
mux.HandleFunc("POST /api/v1/agent-definitions", s.handleAgentDefinitionsCreate)
|
||||
mux.HandleFunc("GET /api/v1/agent-definitions/{id}", s.handleAgentDefinitionsGet)
|
||||
mux.HandleFunc("PATCH /api/v1/agent-definitions/{id}", s.handleAgentDefinitionsUpdate)
|
||||
mux.HandleFunc("DELETE /api/v1/agent-definitions/{id}", s.handleAgentDefinitionsDelete)
|
||||
|
||||
mux.HandleFunc("GET /api/v1/skill-definitions", s.handleSkillDefinitionsList)
|
||||
mux.HandleFunc("POST /api/v1/skill-definitions", s.handleSkillDefinitionsCreate)
|
||||
mux.HandleFunc("GET /api/v1/skill-definitions/{id}", s.handleSkillDefinitionsGet)
|
||||
mux.HandleFunc("PATCH /api/v1/skill-definitions/{id}", s.handleSkillDefinitionsUpdate)
|
||||
mux.HandleFunc("DELETE /api/v1/skill-definitions/{id}", s.handleSkillDefinitionsDelete)
|
||||
|
||||
return 10
|
||||
}
|
||||
|
||||
/* ── Agents ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
func (s *Server) handleAgentDefinitionsList(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
params, err := s.definitions.ParseListParams(r.URL.Query())
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
page, err := s.definitions.ListAgents(r.Context(), ident, params)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writePage(w, page)
|
||||
}
|
||||
|
||||
func (s *Server) handleAgentDefinitionsGet(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.GetAgent(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleAgentDefinitionsCreate(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.CreateAgent(r.Context(), ident, body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusCreated, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleAgentDefinitionsUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.UpdateAgent(r.Context(), ident, r.PathValue("id"), body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleAgentDefinitionsDelete(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.DeleteAgent(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
|
||||
/* ── Skills ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
func (s *Server) handleSkillDefinitionsList(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
params, err := s.definitions.ParseListParams(r.URL.Query())
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
page, err := s.definitions.ListSkills(r.Context(), ident, params)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writePage(w, page)
|
||||
}
|
||||
|
||||
func (s *Server) handleSkillDefinitionsGet(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.GetSkill(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleSkillDefinitionsCreate(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.CreateSkill(r.Context(), ident, body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusCreated, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleSkillDefinitionsUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.UpdateSkill(r.Context(), ident, r.PathValue("id"), body)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
|
||||
func (s *Server) handleSkillDefinitionsDelete(w http.ResponseWriter, r *http.Request) {
|
||||
ident, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
rec, err := s.definitions.DeleteSkill(r.Context(), ident, r.PathValue("id"))
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, rec)
|
||||
}
|
||||
848
go-api/internal/httpserver/definitions_api_test.go
Normal file
848
go-api/internal/httpserver/definitions_api_test.go
Normal file
@@ -0,0 +1,848 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Phase 4E — Backend CRUD APIs for authored Agent and Skill definitions.
|
||||
|
||||
const validAgentMD = `---
|
||||
id: test-agent
|
||||
name: Test Agent
|
||||
description: An authored agent for testing
|
||||
status: draft
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
|
||||
## Instructions
|
||||
Execute testing tasks carefully.
|
||||
`
|
||||
|
||||
const validSkillMD = `---
|
||||
id: test-skill
|
||||
name: Test Skill
|
||||
description: An authored skill for testing
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
|
||||
# Test Skill
|
||||
Skill body instructions.
|
||||
`
|
||||
|
||||
/* ── 1. Agent Create Tests ────────────────────────────────────────────────── */
|
||||
|
||||
func TestAgentCreate(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// 1. Valid personal agent -> 201
|
||||
res := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": validAgentMD,
|
||||
"visibility": "personal",
|
||||
})
|
||||
if res.code != http.StatusCreated {
|
||||
t.Fatalf("create personal agent: got status %d (%v)", res.code, res.body)
|
||||
}
|
||||
rec := res.record(t)
|
||||
if rec["definition_id"] != "test-agent" {
|
||||
t.Errorf("definition_id = %v, want test-agent", rec["definition_id"])
|
||||
}
|
||||
if rec["name"] != "Test Agent" {
|
||||
t.Errorf("name = %v, want Test Agent", rec["name"])
|
||||
}
|
||||
if rec["status"] != "draft" {
|
||||
t.Errorf("status = %v, want draft", rec["status"])
|
||||
}
|
||||
if fmt.Sprint(rec["version"]) != "1" {
|
||||
t.Errorf("version = %v, want 1", rec["version"])
|
||||
}
|
||||
if rec["visibility"] != "personal" {
|
||||
t.Errorf("visibility = %v, want personal", rec["visibility"])
|
||||
}
|
||||
|
||||
// 3. Personal fields derived from authenticated identity
|
||||
if rec["owner_user_id"] != r.talA.id {
|
||||
t.Errorf("owner_user_id = %v, want %s", rec["owner_user_id"], r.talA.id)
|
||||
}
|
||||
if rec["created_by"] != r.talA.id {
|
||||
t.Errorf("created_by = %v, want %s", rec["created_by"], r.talA.id)
|
||||
}
|
||||
if rec["org_id"] != r.orgID {
|
||||
t.Errorf("org_id = %v, want %s", rec["org_id"], r.orgID)
|
||||
}
|
||||
|
||||
// 2. Valid organization agent -> 201 (by admin)
|
||||
orgAgentMD := `---
|
||||
id: shared-agent
|
||||
name: Shared Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
## Instructions
|
||||
Shared instructions.
|
||||
`
|
||||
resOrg := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": orgAgentMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if resOrg.code != http.StatusCreated {
|
||||
t.Fatalf("create org agent: got status %d (%v)", resOrg.code, resOrg.body)
|
||||
}
|
||||
orgRec := resOrg.record(t)
|
||||
// 4. Organization fields derived from authenticated identity
|
||||
if orgRec["visibility"] != "organization" {
|
||||
t.Errorf("visibility = %v, want organization", orgRec["visibility"])
|
||||
}
|
||||
if orgRec["owner_user_id"] != nil {
|
||||
t.Errorf("owner_user_id = %v, want nil for organization tier", orgRec["owner_user_id"])
|
||||
}
|
||||
if orgRec["created_by"] != r.admin.id {
|
||||
t.Errorf("created_by = %v, want %s", orgRec["created_by"], r.admin.id)
|
||||
}
|
||||
|
||||
// 5, 6, 7. Client-supplied org_id, owner_user_id, created_by cannot override session
|
||||
manipulatedMD := `---
|
||||
id: spoof-agent
|
||||
name: Spoof Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resSpoof := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": manipulatedMD,
|
||||
"visibility": "personal",
|
||||
"org_id": r.otherOrgID,
|
||||
"owner_user_id": r.talB.id,
|
||||
"created_by": r.admin.id,
|
||||
})
|
||||
if resSpoof.code != http.StatusCreated {
|
||||
t.Fatalf("create spoofed agent: status %d", resSpoof.code)
|
||||
}
|
||||
spoofRec := resSpoof.record(t)
|
||||
if spoofRec["org_id"] != r.orgID {
|
||||
t.Errorf("org_id spoofed: got %v, want %s", spoofRec["org_id"], r.orgID)
|
||||
}
|
||||
if spoofRec["owner_user_id"] != r.talA.id {
|
||||
t.Errorf("owner_user_id spoofed: got %v, want %s", spoofRec["owner_user_id"], r.talA.id)
|
||||
}
|
||||
if spoofRec["created_by"] != r.talA.id {
|
||||
t.Errorf("created_by spoofed: got %v, want %s", spoofRec["created_by"], r.talA.id)
|
||||
}
|
||||
|
||||
// 8. Invalid Markdown -> 422
|
||||
resEmpty := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": "",
|
||||
})
|
||||
if resEmpty.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("empty markdown: got %d, want 422", resEmpty.code)
|
||||
}
|
||||
|
||||
// 9. Invalid definition_id -> 422
|
||||
badIDMD := `---
|
||||
id: Bad_ID!
|
||||
name: Bad ID Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resBadID := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": badIDMD,
|
||||
})
|
||||
if resBadID.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("bad definition_id: got %d, want 422 (%v)", resBadID.code, resBadID.body)
|
||||
}
|
||||
|
||||
// 10. Missing name -> 422
|
||||
noNameMD := `---
|
||||
id: no-name-agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resNoName := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": noNameMD,
|
||||
})
|
||||
if resNoName.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("missing name: got %d, want 422 (%v)", resNoName.code, resNoName.body)
|
||||
}
|
||||
|
||||
// 11. Version > MaxVersion -> 422
|
||||
hugeVersionMD := `---
|
||||
id: huge-v
|
||||
name: Huge Version
|
||||
version: 999999999999999
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resHugeV := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": hugeVersionMD,
|
||||
})
|
||||
if resHugeV.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("huge version: got %d, want 422 (%v)", resHugeV.code, resHugeV.body)
|
||||
}
|
||||
|
||||
// 12. Duplicate personal definition -> 409
|
||||
resDupPersonal := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": validAgentMD,
|
||||
"visibility": "personal",
|
||||
})
|
||||
if resDupPersonal.code != http.StatusConflict {
|
||||
t.Errorf("duplicate personal agent: got %d, want 409 (%v)", resDupPersonal.code, resDupPersonal.body)
|
||||
}
|
||||
|
||||
// 13. Duplicate organization definition -> 409
|
||||
resDupOrg := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": orgAgentMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if resDupOrg.code != http.StatusConflict {
|
||||
t.Errorf("duplicate org agent: got %d, want 409 (%v)", resDupOrg.code, resDupOrg.body)
|
||||
}
|
||||
|
||||
// Shadow-by-id: personal agent with SAME id as organization agent succeeds!
|
||||
resShadow := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": orgAgentMD,
|
||||
"visibility": "personal",
|
||||
})
|
||||
if resShadow.code != http.StatusCreated {
|
||||
t.Errorf("shadow personal agent: got %d, want 201 (%v)", resShadow.code, resShadow.body)
|
||||
}
|
||||
|
||||
// Talent cannot create organization definition -> 403
|
||||
talOrgMD := `---
|
||||
id: tal-org
|
||||
name: Tal Org
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resTalOrg := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": talOrgMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if resTalOrg.code != http.StatusForbidden {
|
||||
t.Errorf("talent create org agent: got %d, want 403 (%v)", resTalOrg.code, resTalOrg.body)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 2. Skill Create Tests ────────────────────────────────────────────────── */
|
||||
|
||||
func TestSkillCreate(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// 14. Valid personal skill -> 201
|
||||
res := r.as(r.talA, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": validSkillMD,
|
||||
"visibility": "personal",
|
||||
})
|
||||
if res.code != http.StatusCreated {
|
||||
t.Fatalf("create personal skill: got %d (%v)", res.code, res.body)
|
||||
}
|
||||
rec := res.record(t)
|
||||
if rec["definition_id"] != "test-skill" {
|
||||
t.Errorf("definition_id = %v, want test-skill", rec["definition_id"])
|
||||
}
|
||||
if rec["name"] != "Test Skill" {
|
||||
t.Errorf("name = %v, want Test Skill", rec["name"])
|
||||
}
|
||||
if rec["status"] != "active" {
|
||||
t.Errorf("status = %v, want active", rec["status"])
|
||||
}
|
||||
if rec["visibility"] != "personal" {
|
||||
t.Errorf("visibility = %v, want personal", rec["visibility"])
|
||||
}
|
||||
if rec["owner_user_id"] != r.talA.id {
|
||||
t.Errorf("owner_user_id = %v, want %s", rec["owner_user_id"], r.talA.id)
|
||||
}
|
||||
// 21. Skills do NOT have a version column
|
||||
if _, hasVersion := rec["version"]; hasVersion {
|
||||
t.Errorf("skill record has version field; skills must not have a version")
|
||||
}
|
||||
|
||||
// 15. Valid organization skill -> 201
|
||||
orgSkillMD := `---
|
||||
id: org-skill
|
||||
name: Org Skill
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
# Org Skill
|
||||
`
|
||||
resOrg := r.as(r.admin, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": orgSkillMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if resOrg.code != http.StatusCreated {
|
||||
t.Fatalf("create org skill: got %d (%v)", resOrg.code, resOrg.body)
|
||||
}
|
||||
|
||||
// 16. Invalid Markdown -> 422
|
||||
resEmpty := r.as(r.admin, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": "",
|
||||
})
|
||||
if resEmpty.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("empty skill markdown: got %d, want 422", resEmpty.code)
|
||||
}
|
||||
|
||||
// 17. Invalid definition_id -> 422
|
||||
badIDMD := `---
|
||||
id: BAD_SKILL
|
||||
name: Bad Skill
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
resBadID := r.as(r.admin, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": badIDMD,
|
||||
})
|
||||
if resBadID.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("bad skill id: got %d, want 422", resBadID.code)
|
||||
}
|
||||
|
||||
// 18. Invalid page -> 422
|
||||
badPageMD := `---
|
||||
id: bad-page-skill
|
||||
name: Bad Page Skill
|
||||
pages:
|
||||
- totally_unknown_page_xyz
|
||||
---
|
||||
`
|
||||
resBadPage := r.as(r.admin, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": badPageMD,
|
||||
})
|
||||
if resBadPage.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("bad skill page: got %d, want 422 (%v)", resBadPage.code, resBadPage.body)
|
||||
}
|
||||
|
||||
// 19. Duplicate personal skill -> 409
|
||||
resDupPers := r.as(r.talA, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": validSkillMD,
|
||||
"visibility": "personal",
|
||||
})
|
||||
if resDupPers.code != http.StatusConflict {
|
||||
t.Errorf("duplicate personal skill: got %d, want 409 (%v)", resDupPers.code, resDupPers.body)
|
||||
}
|
||||
|
||||
// 20. Duplicate organization skill -> 409
|
||||
resDupOrg := r.as(r.admin, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": orgSkillMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if resDupOrg.code != http.StatusConflict {
|
||||
t.Errorf("duplicate org skill: got %d, want 409 (%v)", resDupOrg.code, resDupOrg.body)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 3. List Tests ────────────────────────────────────────────────────────── */
|
||||
|
||||
func TestDefinitionsList(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// Create:
|
||||
// - 1 org agent (admin)
|
||||
// - 1 personal agent for talA
|
||||
// - 1 personal agent for talB
|
||||
// - 1 org agent for outsider (in other org)
|
||||
r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: org-agent-1
|
||||
name: Org Agent 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "organization",
|
||||
})
|
||||
r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: tala-agent
|
||||
name: TalA Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
r.as(r.talB, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: talb-agent
|
||||
name: TalB Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
r.as(r.outsider, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: outsider-agent
|
||||
name: Outsider Agent
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "organization",
|
||||
})
|
||||
|
||||
// 22, 23, 24, 25. Scoping assertions
|
||||
talAList := r.as(r.talA, "GET", "/api/v1/agent-definitions", nil)
|
||||
if talAList.code != http.StatusOK {
|
||||
t.Fatalf("talA list: %d", talAList.code)
|
||||
}
|
||||
talARecs := talAList.records(t)
|
||||
talAIDMap := map[string]bool{}
|
||||
for _, rec := range talARecs {
|
||||
talAIDMap[rec["definition_id"].(string)] = true
|
||||
}
|
||||
|
||||
if !talAIDMap["org-agent-1"] {
|
||||
t.Errorf("talA should see org-agent-1")
|
||||
}
|
||||
if !talAIDMap["tala-agent"] {
|
||||
t.Errorf("talA should see tala-agent")
|
||||
}
|
||||
if talAIDMap["talb-agent"] {
|
||||
t.Errorf("talA must NOT see talB's personal agent")
|
||||
}
|
||||
if talAIDMap["outsider-agent"] {
|
||||
t.Errorf("talA must NOT see outsider organization's agent")
|
||||
}
|
||||
|
||||
// 26. Visibility filter
|
||||
onlyPersonal := r.as(r.talA, "GET", "/api/v1/agent-definitions?visibility=personal", nil).records(t)
|
||||
for _, rec := range onlyPersonal {
|
||||
if rec["visibility"] != "personal" {
|
||||
t.Errorf("expected only personal visibility, got %v", rec["visibility"])
|
||||
}
|
||||
}
|
||||
onlyOrg := r.as(r.talA, "GET", "/api/v1/agent-definitions?visibility=organization", nil).records(t)
|
||||
for _, rec := range onlyOrg {
|
||||
if rec["visibility"] != "organization" {
|
||||
t.Errorf("expected only organization visibility, got %v", rec["visibility"])
|
||||
}
|
||||
}
|
||||
badVis := r.as(r.talA, "GET", "/api/v1/agent-definitions?visibility=invalid_vis", nil)
|
||||
if badVis.code != http.StatusBadRequest {
|
||||
t.Errorf("bad visibility filter: got %d, want 400", badVis.code)
|
||||
}
|
||||
|
||||
// 27. Status filter
|
||||
filteredStatus := r.as(r.talA, "GET", "/api/v1/agent-definitions?status=draft", nil).records(t)
|
||||
for _, rec := range filteredStatus {
|
||||
if rec["status"] != "draft" {
|
||||
t.Errorf("expected draft status, got %v", rec["status"])
|
||||
}
|
||||
}
|
||||
|
||||
// 28. Definition ID filter
|
||||
defIDList := r.as(r.talA, "GET", "/api/v1/agent-definitions?definition_id=tala-agent", nil).records(t)
|
||||
if len(defIDList) != 1 || defIDList[0]["definition_id"] != "tala-agent" {
|
||||
t.Errorf("definition_id filter failed: got %v", defIDList)
|
||||
}
|
||||
|
||||
// 29. Pagination
|
||||
page1 := r.as(r.talA, "GET", "/api/v1/agent-definitions?limit=1&offset=0", nil)
|
||||
meta1 := page1.meta(t)
|
||||
if fmt.Sprint(meta1["limit"]) != "1" || fmt.Sprint(meta1["offset"]) != "0" {
|
||||
t.Errorf("pagination meta: %v", meta1)
|
||||
}
|
||||
|
||||
// 30. Stable sorting
|
||||
sortedAsc := r.as(r.talA, "GET", "/api/v1/agent-definitions?sort=definition_id", nil).records(t)
|
||||
if len(sortedAsc) >= 2 {
|
||||
id0 := sortedAsc[0]["definition_id"].(string)
|
||||
id1 := sortedAsc[1]["definition_id"].(string)
|
||||
if id0 > id1 {
|
||||
t.Errorf("ascending sort failed: %s > %s", id0, id1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4. Get By ID Tests ───────────────────────────────────────────────────── */
|
||||
|
||||
func TestDefinitionsGet(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// Create personal agent for talA
|
||||
createA := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: get-pers-a
|
||||
name: Get Pers A
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
idPersA := createA.record(t)["id"].(string)
|
||||
|
||||
// Create org agent
|
||||
createOrg := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: get-org
|
||||
name: Get Org
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "organization",
|
||||
})
|
||||
idOrg := createOrg.record(t)["id"].(string)
|
||||
|
||||
// Create personal agent for outsider
|
||||
createOutsider := r.as(r.outsider, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: get-pers-outsider
|
||||
name: Get Pers Outsider
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
idOutsider := createOutsider.record(t)["id"].(string)
|
||||
|
||||
// 31. Own personal definition -> 200
|
||||
getPersA := r.as(r.talA, "GET", "/api/v1/agent-definitions/"+idPersA, nil)
|
||||
if getPersA.code != http.StatusOK {
|
||||
t.Errorf("get own personal agent: %d", getPersA.code)
|
||||
}
|
||||
|
||||
// 32. Same-org organization definition -> 200
|
||||
getOrgByTal := r.as(r.talA, "GET", "/api/v1/agent-definitions/"+idOrg, nil)
|
||||
if getOrgByTal.code != http.StatusOK {
|
||||
t.Errorf("get same-org agent: %d", getOrgByTal.code)
|
||||
}
|
||||
|
||||
// 33. Other user's personal definition -> 404 (inaccessible)
|
||||
getPersByTalB := r.as(r.talB, "GET", "/api/v1/agent-definitions/"+idPersA, nil)
|
||||
if getPersByTalB.code != http.StatusNotFound {
|
||||
t.Errorf("get other user personal agent: got %d, want 404", getPersByTalB.code)
|
||||
}
|
||||
|
||||
// 34. Other organization's definition -> 404 (inaccessible)
|
||||
getOutsiderByTalA := r.as(r.talA, "GET", "/api/v1/agent-definitions/"+idOutsider, nil)
|
||||
if getOutsiderByTalA.code != http.StatusNotFound {
|
||||
t.Errorf("get outsider definition: got %d, want 404", getOutsiderByTalA.code)
|
||||
}
|
||||
|
||||
// Malformed UUID -> 404
|
||||
getMalformed := r.as(r.talA, "GET", "/api/v1/agent-definitions/not-a-uuid", nil)
|
||||
if getMalformed.code != http.StatusNotFound {
|
||||
t.Errorf("get malformed uuid: got %d, want 404", getMalformed.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 5. Patch Tests ───────────────────────────────────────────────────────── */
|
||||
|
||||
func TestDefinitionsPatch(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// Create personal agent for talA
|
||||
createA := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: patch-agent
|
||||
name: Initial Name
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
Initial Body`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
idA := createA.record(t)["id"].(string)
|
||||
origCreated := createA.record(t)["created_date"].(string)
|
||||
origUpdated := createA.record(t)["updated_date"].(string)
|
||||
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
|
||||
// 35, 36, 37, 38. Markdown update -> 200, projections updated, markdown verbatim, updated_date changed
|
||||
updatedMD := `---
|
||||
id: patch-agent
|
||||
name: Updated Name
|
||||
version: 2
|
||||
status: published
|
||||
pages:
|
||||
- candidates
|
||||
- positions
|
||||
---
|
||||
|
||||
# Updated Body
|
||||
Verbatim content with trailing spaces
|
||||
`
|
||||
patchRes := r.as(r.talA, "PATCH", "/api/v1/agent-definitions/"+idA, map[string]any{
|
||||
"markdown": updatedMD,
|
||||
})
|
||||
if patchRes.code != http.StatusOK {
|
||||
t.Fatalf("patch agent: %d (%v)", patchRes.code, patchRes.body)
|
||||
}
|
||||
patchedRec := patchRes.record(t)
|
||||
if patchedRec["name"] != "Updated Name" {
|
||||
t.Errorf("name = %v, want Updated Name", patchedRec["name"])
|
||||
}
|
||||
if patchedRec["status"] != "published" {
|
||||
t.Errorf("status = %v, want published", patchedRec["status"])
|
||||
}
|
||||
if fmt.Sprint(patchedRec["version"]) != "2" {
|
||||
t.Errorf("version = %v, want 2", patchedRec["version"])
|
||||
}
|
||||
if patchedRec["markdown"] != updatedMD {
|
||||
t.Errorf("markdown not verbatim:\n got: %q\nwant: %q", patchedRec["markdown"], updatedMD)
|
||||
}
|
||||
if patchedRec["created_date"] != origCreated {
|
||||
t.Errorf("created_date changed on patch")
|
||||
}
|
||||
if patchedRec["updated_date"] == origUpdated {
|
||||
t.Errorf("updated_date did not advance")
|
||||
}
|
||||
|
||||
// 39. Invalid Markdown update -> 422
|
||||
badPatch := r.as(r.talA, "PATCH", "/api/v1/agent-definitions/"+idA, map[string]any{
|
||||
"markdown": "--- invalid yaml --",
|
||||
})
|
||||
if badPatch.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("invalid patch md: got %d, want 422", badPatch.code)
|
||||
}
|
||||
|
||||
// 40. Server-owned fields cannot be modified
|
||||
spoofPatch := r.as(r.talA, "PATCH", "/api/v1/agent-definitions/"+idA, map[string]any{
|
||||
"owner_user_id": r.talB.id,
|
||||
"org_id": r.otherOrgID,
|
||||
"created_by": r.admin.id,
|
||||
})
|
||||
if spoofPatch.code != http.StatusOK {
|
||||
t.Errorf("spoof patch status: %d", spoofPatch.code)
|
||||
}
|
||||
reread := r.as(r.talA, "GET", "/api/v1/agent-definitions/"+idA, nil).record(t)
|
||||
if reread["owner_user_id"] != r.talA.id {
|
||||
t.Errorf("owner_user_id altered on patch: %v", reread["owner_user_id"])
|
||||
}
|
||||
if reread["org_id"] != r.orgID {
|
||||
t.Errorf("org_id altered on patch: %v", reread["org_id"])
|
||||
}
|
||||
|
||||
// 41. Unauthorized update: talB cannot patch talA's definition -> 404
|
||||
resTalBPatch := r.as(r.talB, "PATCH", "/api/v1/agent-definitions/"+idA, map[string]any{
|
||||
"status": "archived",
|
||||
})
|
||||
if resTalBPatch.code != http.StatusNotFound {
|
||||
t.Errorf("talB patch talA: got %d, want 404", resTalBPatch.code)
|
||||
}
|
||||
|
||||
// Create org agent
|
||||
createOrg := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: org-for-patch
|
||||
name: Org Patch
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "organization",
|
||||
})
|
||||
idOrg := createOrg.record(t)["id"].(string)
|
||||
|
||||
// Talent cannot patch org definition -> 403
|
||||
talOrgPatch := r.as(r.talA, "PATCH", "/api/v1/agent-definitions/"+idOrg, map[string]any{
|
||||
"status": "archived",
|
||||
})
|
||||
if talOrgPatch.code != http.StatusForbidden {
|
||||
t.Errorf("talent patch org agent: got %d, want 403", talOrgPatch.code)
|
||||
}
|
||||
|
||||
// Employer can patch org definition -> 200
|
||||
empOrgPatch := r.as(r.empA, "PATCH", "/api/v1/agent-definitions/"+idOrg, map[string]any{
|
||||
"status": "archived",
|
||||
})
|
||||
if empOrgPatch.code != http.StatusOK {
|
||||
t.Errorf("employer patch org agent: got %d, want 200", empOrgPatch.code)
|
||||
}
|
||||
|
||||
// Visibility is immutable after creation -> 422
|
||||
visPatch := r.as(r.admin, "PATCH", "/api/v1/agent-definitions/"+idOrg, map[string]any{
|
||||
"visibility": "personal",
|
||||
})
|
||||
if visPatch.code != http.StatusUnprocessableEntity {
|
||||
t.Errorf("visibility mutation: got %d, want 422 (%v)", visPatch.code, visPatch.body)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6. Delete Tests ──────────────────────────────────────────────────────── */
|
||||
|
||||
func TestDefinitionsDelete(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// Create personal agent for talA
|
||||
createA := r.as(r.talA, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: del-agent-a
|
||||
name: Del Agent A
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "personal",
|
||||
})
|
||||
idA := createA.record(t)["id"].(string)
|
||||
|
||||
// Create org agent
|
||||
createOrg := r.as(r.admin, "POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": `---
|
||||
id: del-agent-org
|
||||
name: Del Agent Org
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`,
|
||||
"visibility": "organization",
|
||||
})
|
||||
idOrg := createOrg.record(t)["id"].(string)
|
||||
|
||||
// 46. Talent cannot delete org definition -> 403
|
||||
talDelOrg := r.as(r.talA, "DELETE", "/api/v1/agent-definitions/"+idOrg, nil)
|
||||
if talDelOrg.code != http.StatusForbidden {
|
||||
t.Errorf("talent delete org agent: got %d, want 403", talDelOrg.code)
|
||||
}
|
||||
|
||||
// 44, 48, 49. Owner can delete personal agent -> 200, returns { "data": { "id": ... } }, subsequent GET -> 404
|
||||
delA := r.as(r.talA, "DELETE", "/api/v1/agent-definitions/"+idA, nil)
|
||||
if delA.code != http.StatusOK {
|
||||
t.Fatalf("delete personal agent: got %d", delA.code)
|
||||
}
|
||||
delRec := delA.record(t)
|
||||
if delRec["id"] != idA {
|
||||
t.Errorf("delete response id = %v, want %s", delRec["id"], idA)
|
||||
}
|
||||
getAAfter := r.as(r.talA, "GET", "/api/v1/agent-definitions/"+idA, nil)
|
||||
if getAAfter.code != http.StatusNotFound {
|
||||
t.Errorf("subsequent GET deleted agent: got %d, want 404", getAAfter.code)
|
||||
}
|
||||
|
||||
// 45. Operator (employer) can delete org definition -> 200
|
||||
delOrg := r.as(r.empA, "DELETE", "/api/v1/agent-definitions/"+idOrg, nil)
|
||||
if delOrg.code != http.StatusOK {
|
||||
t.Fatalf("employer delete org agent: got %d", delOrg.code)
|
||||
}
|
||||
getOrgAfter := r.as(r.admin, "GET", "/api/v1/agent-definitions/"+idOrg, nil)
|
||||
if getOrgAfter.code != http.StatusNotFound {
|
||||
t.Errorf("subsequent GET deleted org agent: got %d, want 404", getOrgAfter.code)
|
||||
}
|
||||
|
||||
// Idempotent delete on non-existent UUID -> 200
|
||||
missingUUID := "00000000-0000-0000-0000-000000000000"
|
||||
delMissing := r.as(r.talA, "DELETE", "/api/v1/agent-definitions/"+missingUUID, nil)
|
||||
if delMissing.code != http.StatusOK {
|
||||
t.Errorf("idempotent delete: got %d, want 200", delMissing.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 7. Security and SQL Injection ────────────────────────────────────────── */
|
||||
|
||||
func TestSecurityAndSQLInjection(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// SQL injection in filter
|
||||
sqliList := r.as(r.talA, "GET", "/api/v1/agent-definitions?definition_id=x'%20OR%20'1'='1", nil)
|
||||
if sqliList.code != http.StatusOK {
|
||||
t.Errorf("sqli filter request failed: %d", sqliList.code)
|
||||
}
|
||||
if len(sqliList.records(t)) != 0 {
|
||||
t.Errorf("sqli in definition_id filter leaked records")
|
||||
}
|
||||
|
||||
// SQL injection in sort
|
||||
sqliSort := r.as(r.talA, "GET", "/api/v1/agent-definitions?sort=name%20DESC%3BDROP%20TABLE%20users%3B", nil)
|
||||
if sqliSort.code != http.StatusBadRequest {
|
||||
t.Errorf("sqli in sort should be rejected as invalid query: got %d (%v)", sqliSort.code, sqliSort.body)
|
||||
}
|
||||
|
||||
// Unauthenticated requests -> 401
|
||||
unauthList := r.doAnon("GET", "/api/v1/agent-definitions", nil)
|
||||
if unauthList.code != http.StatusUnauthorized {
|
||||
t.Errorf("unauth list: got %d, want 401", unauthList.code)
|
||||
}
|
||||
unauthCreate := r.doAnon("POST", "/api/v1/agent-definitions", map[string]any{
|
||||
"markdown": validAgentMD,
|
||||
})
|
||||
if unauthCreate.code != http.StatusUnauthorized {
|
||||
t.Errorf("unauth create: got %d, want 401", unauthCreate.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8. Full CRUD & Projection Consistency Flow ───────────────────────────── */
|
||||
|
||||
func TestFullCRUDFlowAndProjections(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// 58. Create
|
||||
createRes := r.as(r.empA, "POST", "/api/v1/skill-definitions", map[string]any{
|
||||
"markdown": validSkillMD,
|
||||
"visibility": "organization",
|
||||
})
|
||||
if createRes.code != http.StatusCreated {
|
||||
t.Fatalf("create skill failed: %d (%v)", createRes.code, createRes.body)
|
||||
}
|
||||
id := createRes.record(t)["id"].(string)
|
||||
|
||||
// 59. List includes created record
|
||||
listRes := r.as(r.empA, "GET", "/api/v1/skill-definitions?definition_id=test-skill", nil)
|
||||
if listRes.code != http.StatusOK || len(listRes.records(t)) == 0 {
|
||||
t.Fatalf("list skill failed: %d (%v)", listRes.code, listRes.body)
|
||||
}
|
||||
|
||||
// 60. Get created record and verify projections
|
||||
getRes := r.as(r.talA, "GET", "/api/v1/skill-definitions/"+id, nil)
|
||||
if getRes.code != http.StatusOK {
|
||||
t.Fatalf("get skill failed: %d", getRes.code)
|
||||
}
|
||||
rec := getRes.record(t)
|
||||
if rec["definition_id"] != "test-skill" || rec["name"] != "Test Skill" || rec["status"] != "active" {
|
||||
t.Errorf("projection mismatch on get: %v", rec)
|
||||
}
|
||||
if rec["markdown"] != validSkillMD {
|
||||
t.Errorf("markdown not verbatim on get")
|
||||
}
|
||||
|
||||
// 61. Patch
|
||||
newSkillMD := `---
|
||||
id: test-skill
|
||||
name: Updated Skill Name
|
||||
status: inactive
|
||||
pages:
|
||||
- candidates
|
||||
- profile
|
||||
---
|
||||
# Updated Skill Body
|
||||
`
|
||||
patchRes := r.as(r.empA, "PATCH", "/api/v1/skill-definitions/"+id, map[string]any{
|
||||
"markdown": newSkillMD,
|
||||
})
|
||||
if patchRes.code != http.StatusOK {
|
||||
t.Fatalf("patch skill failed: %d (%v)", patchRes.code, patchRes.body)
|
||||
}
|
||||
patchedRec := patchRes.record(t)
|
||||
if patchedRec["name"] != "Updated Skill Name" || patchedRec["status"] != "inactive" {
|
||||
t.Errorf("projection not updated on patch: %v", patchedRec)
|
||||
}
|
||||
if patchedRec["markdown"] != newSkillMD {
|
||||
t.Errorf("markdown not verbatim on patch")
|
||||
}
|
||||
|
||||
// 62. Delete
|
||||
delRes := r.as(r.empA, "DELETE", "/api/v1/skill-definitions/"+id, nil)
|
||||
if delRes.code != http.StatusOK {
|
||||
t.Fatalf("delete skill failed: %d", delRes.code)
|
||||
}
|
||||
getAfterDel := r.as(r.empA, "GET", "/api/v1/skill-definitions/"+id, nil)
|
||||
if getAfterDel.code != http.StatusNotFound {
|
||||
t.Errorf("get after delete: got %d, want 404", getAfterDel.code)
|
||||
}
|
||||
}
|
||||
326
go-api/internal/httpserver/me.go
Normal file
326
go-api/internal/httpserver/me.go
Normal file
@@ -0,0 +1,326 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
// The current-user endpoints. See api-contract.md §9.
|
||||
//
|
||||
// "Current" now means the user behind the session cookie, resolved by the
|
||||
// authentication middleware and read from the request context. Before Phase 3C
|
||||
// it meant the organization's oldest user, found with ORDER BY created_date
|
||||
// LIMIT 1 — a placeholder that was correct only because there was exactly one.
|
||||
|
||||
// preferenceColumns maps the frontend's camelCase preference keys onto their
|
||||
// columns. Anything not listed here lives in user_preferences.extra — which is
|
||||
// where customSkills and customAgents, every account-authored definition,
|
||||
// currently are.
|
||||
var preferenceColumns = map[string]string{
|
||||
"owliverDefault": "owliver_default",
|
||||
"compactDensity": "compact_density",
|
||||
"emailDigest": "email_digest",
|
||||
}
|
||||
|
||||
// userColumns is the projection for a user record.
|
||||
const userColumns = `id::text AS id, legacy_id, full_name, email::text AS email,
|
||||
role, account_type, status,
|
||||
to_char(created_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS created_date,
|
||||
to_char(updated_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS updated_date`
|
||||
|
||||
// updatableUserFields are the only columns PATCH /me may write, and this map is
|
||||
// the entire authority on that: a column absent from it cannot be reached by
|
||||
// this endpoint at all, whatever the request body says.
|
||||
//
|
||||
// full_name the Profile page's display-name field.
|
||||
// account_type which product surface the person is looking at — Employer or
|
||||
//
|
||||
// Talent. Layout.jsx writes it when the viewer switches.
|
||||
// Explicitly NOT an authorization field (Phase 3B decision 1); it
|
||||
// is a display attribute, and users.role is what authorizes.
|
||||
//
|
||||
// Everything else is server-owned. `role` used to be in this map, which meant
|
||||
// any signed-in user could promote themselves to admin with a one-line PATCH
|
||||
// the moment sessions existed. It was harmless while there was no
|
||||
// authentication and a live privilege-escalation path the instant there was.
|
||||
// See serverOwnedUserFields.
|
||||
var updatableUserFields = map[string]string{
|
||||
"full_name": "text",
|
||||
"account_type": "text",
|
||||
}
|
||||
|
||||
// serverOwnedUserFields are the fields a user must never write about
|
||||
// themselves, listed by name so an attempt can be recognised and logged rather
|
||||
// than silently dropped in with every other unknown key.
|
||||
//
|
||||
// They are ignored, not rejected: a PATCH body is "whichever fields the caller
|
||||
// sent", the endpoint has always ignored what it does not own, and the response
|
||||
// returns the user as they actually are — so a caller who asks for a role
|
||||
// change gets a 200 whose body shows the role unchanged. What is new is that
|
||||
// the attempt is now visible in the log, because a client asking to change its
|
||||
// own role is worth knowing about even when the answer is no.
|
||||
var serverOwnedUserFields = map[string]bool{
|
||||
"id": true,
|
||||
"org_id": true,
|
||||
"role": true,
|
||||
"password_hash": true,
|
||||
"status": true,
|
||||
"email": true,
|
||||
"legacy_id": true,
|
||||
"last_login_at": true,
|
||||
"created_date": true,
|
||||
"updated_date": true,
|
||||
}
|
||||
|
||||
func (s *Server) routeMe(mux *http.ServeMux) int {
|
||||
mux.HandleFunc("GET /api/v1/me", s.handleMeGet)
|
||||
mux.HandleFunc("PATCH /api/v1/me", s.handleMePatch)
|
||||
mux.HandleFunc("GET /api/v1/me/preferences", s.handlePreferencesGet)
|
||||
mux.HandleFunc("PATCH /api/v1/me/preferences", s.handlePreferencesPatch)
|
||||
return 4
|
||||
}
|
||||
|
||||
// userRecord reads one user by id, with preferences embedded.
|
||||
//
|
||||
// Embedded rather than a sibling resource because krowHooks.js:42 reads
|
||||
// `user?.preferences` straight off the object returned by auth.me().
|
||||
//
|
||||
// The id is always one this server resolved from a session — never a value off
|
||||
// the request. There is deliberately no variant of this function that takes an
|
||||
// identifier from a caller.
|
||||
func (s *Server) userRecord(ctx context.Context, q repo.Querier, userID string) (domain.Record, error) {
|
||||
rows, err := q.Query(ctx,
|
||||
`SELECT `+userColumns+` FROM users WHERE id = $1::uuid`,
|
||||
userID)
|
||||
if err != nil {
|
||||
return nil, domain.Internal(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
fields := rows.FieldDescriptions()
|
||||
if !rows.Next() {
|
||||
// Unreachable in practice: the middleware just read this row to build
|
||||
// the identity. Reported rather than papered over.
|
||||
return nil, domain.NotFound("User", "current")
|
||||
}
|
||||
vals, err := rows.Values()
|
||||
if err != nil {
|
||||
return nil, domain.Internal(err)
|
||||
}
|
||||
user := make(domain.Record, len(fields)+1)
|
||||
for i, f := range fields {
|
||||
user[string(f.Name)] = vals[i]
|
||||
}
|
||||
rows.Close()
|
||||
|
||||
prefs, err := s.preferences(ctx, q, user["id"].(string))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user["preferences"] = prefs
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// preferences reads the three columns plus the extra blob, flattened into the
|
||||
// single camelCase object the frontend expects.
|
||||
func (s *Server) preferences(ctx context.Context, q repo.Querier, userID string) (map[string]any, error) {
|
||||
var owliver, compact, digest bool
|
||||
var extra []byte
|
||||
err := q.QueryRow(ctx,
|
||||
`SELECT owliver_default, compact_density, email_digest, extra
|
||||
FROM user_preferences WHERE user_id = $1::uuid`, userID).
|
||||
Scan(&owliver, &compact, &digest, &extra)
|
||||
|
||||
out := map[string]any{}
|
||||
if err != nil {
|
||||
// No row yet is a legitimate state: the defaults below are the schema's.
|
||||
return map[string]any{"owliverDefault": true, "compactDensity": false, "emailDigest": true}, nil
|
||||
}
|
||||
if len(extra) > 0 {
|
||||
_ = json.Unmarshal(extra, &out)
|
||||
}
|
||||
out["owliverDefault"] = owliver
|
||||
out["compactDensity"] = compact
|
||||
out["emailDigest"] = digest
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *Server) handleMeGet(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
user, err := s.userRecord(r.Context(), s.db.Pool, id.UserID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, user)
|
||||
}
|
||||
|
||||
func (s *Server) handleMePatch(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
// The row to update is the session's user. Note what is NOT consulted: the
|
||||
// body may contain an "id", and it is ignored — writing to whichever user
|
||||
// the caller names is the whole of the vulnerability this avoids.
|
||||
sets, args := []string{}, []any{id.UserID}
|
||||
details := map[string]string{}
|
||||
for k, v := range body {
|
||||
if k == "preferences" {
|
||||
details[k] = "update preferences through /api/v1/me/preferences"
|
||||
continue
|
||||
}
|
||||
pgType, ok := updatableUserFields[k]
|
||||
if !ok {
|
||||
if serverOwnedUserFields[k] {
|
||||
s.log.Warn("ignored an attempt to write a server-owned user field",
|
||||
"field", k, "user_id", id.UserID, "session_id", id.SessionID)
|
||||
}
|
||||
continue // server-owned, or simply not a column: ignored either way
|
||||
}
|
||||
str, isStr := v.(string)
|
||||
if !isStr {
|
||||
details[k] = "expected a string"
|
||||
continue
|
||||
}
|
||||
args = append(args, str)
|
||||
sets = append(sets, k+" = $"+strconv.Itoa(len(args))+"::"+pgType)
|
||||
}
|
||||
if len(details) > 0 {
|
||||
writeError(w, s.log, domain.Validation("User payload is not valid", details))
|
||||
return
|
||||
}
|
||||
|
||||
if len(sets) > 0 {
|
||||
// Every column name in `sets` came from updatableUserFields, which is a
|
||||
// literal map in this file. No identifier here is caller-supplied; the
|
||||
// values are all bind parameters.
|
||||
q := "UPDATE users SET " + strings.Join(sets, ", ") + ", updated_date = now() WHERE id = $1::uuid"
|
||||
if _, err := s.db.Pool.Exec(r.Context(), q, args...); err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
user, err := s.userRecord(r.Context(), s.db.Pool, id.UserID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeRecord(w, http.StatusOK, user)
|
||||
}
|
||||
|
||||
func (s *Server) handlePreferencesGet(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
prefs, err := s.preferences(r.Context(), s.db.Pool, id.UserID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, envelope{Data: prefs})
|
||||
}
|
||||
|
||||
// handlePreferencesPatch shallow-merges the supplied keys and returns the whole
|
||||
// merged object, matching auth.updatePreferences().
|
||||
func (s *Server) handlePreferencesPatch(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := authctx.MustFrom(r.Context())
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
body, err := decodeBody(r)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
userID := id.UserID
|
||||
current, err := s.preferences(r.Context(), s.db.Pool, userID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
|
||||
merged := current
|
||||
if merged == nil {
|
||||
merged = map[string]any{}
|
||||
}
|
||||
details := map[string]string{}
|
||||
for k, v := range body {
|
||||
if _, isColumn := preferenceColumns[k]; isColumn {
|
||||
if _, ok := v.(bool); !ok {
|
||||
details[k] = "expected a boolean"
|
||||
continue
|
||||
}
|
||||
}
|
||||
merged[k] = v
|
||||
}
|
||||
if len(details) > 0 {
|
||||
writeError(w, s.log, domain.Validation("preferences payload is not valid", details))
|
||||
return
|
||||
}
|
||||
|
||||
// Split the merged object back into its columns and the extra blob.
|
||||
extra := map[string]any{}
|
||||
for k, v := range merged {
|
||||
if _, isColumn := preferenceColumns[k]; !isColumn {
|
||||
extra[k] = v
|
||||
}
|
||||
}
|
||||
extraJSON, err := json.Marshal(extra)
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Validation("preferences are not encodable as JSON", nil))
|
||||
return
|
||||
}
|
||||
|
||||
_, err = s.db.Pool.Exec(r.Context(),
|
||||
`INSERT INTO user_preferences (user_id, owliver_default, compact_density, email_digest, extra, updated_date)
|
||||
VALUES ($1::uuid, $2::boolean, $3::boolean, $4::boolean, $5::jsonb, now())
|
||||
ON CONFLICT (user_id) DO UPDATE SET
|
||||
owliver_default = EXCLUDED.owliver_default,
|
||||
compact_density = EXCLUDED.compact_density,
|
||||
email_digest = EXCLUDED.email_digest,
|
||||
extra = EXCLUDED.extra,
|
||||
updated_date = now()`,
|
||||
userID, truthy(merged["owliverDefault"], true), truthy(merged["compactDensity"], false),
|
||||
truthy(merged["emailDigest"], true), extraJSON)
|
||||
if err != nil {
|
||||
writeError(w, s.log, domain.Internal(err))
|
||||
return
|
||||
}
|
||||
|
||||
prefs, err := s.preferences(r.Context(), s.db.Pool, userID)
|
||||
if err != nil {
|
||||
writeError(w, s.log, err)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, envelope{Data: prefs})
|
||||
}
|
||||
|
||||
func truthy(v any, fallback bool) bool {
|
||||
if b, ok := v.(bool); ok {
|
||||
return b
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
154
go-api/internal/httpserver/ratelimit.go
Normal file
154
go-api/internal/httpserver/ratelimit.go
Normal file
@@ -0,0 +1,154 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Login rate limiting.
|
||||
//
|
||||
// WHAT THIS IS: a fixed-window counter of *failed* sign-in attempts, held in
|
||||
// this process's memory, keyed by client address and by email address. It turns
|
||||
// online password guessing from "as fast as the server can hash" into a handful
|
||||
// of tries per window, which is the whole job.
|
||||
//
|
||||
// WHAT THIS IS NOT, and the limitation to carry into production:
|
||||
//
|
||||
// - It is per-process. Two API instances behind a load balancer each allow
|
||||
// the full budget, so the effective limit is the limit times the instance
|
||||
// count, and a restart clears every counter. A deployment with more than
|
||||
// one instance needs shared state — Redis, or the database — and this
|
||||
// package is the seam where that goes: attemptLimiter is an implementation
|
||||
// detail behind Allow/Fail/Reset.
|
||||
// - It trusts net/http's RemoteAddr for the client address. Behind a reverse
|
||||
// proxy every request appears to come from the proxy, so the per-address
|
||||
// budget becomes global. Reading X-Forwarded-For instead would be worse,
|
||||
// not better, until there is a trusted-proxy list to validate it against —
|
||||
// a client can send that header itself and mint a fresh budget per request.
|
||||
// Deploying behind a proxy means adding that list first.
|
||||
// - It is memory-bounded by pruning, not by a hard cap, so a flood from many
|
||||
// distinct addresses grows the map until the next prune.
|
||||
//
|
||||
// Only failures are counted. A correct password resets the email's counter, so
|
||||
// a person who mistypes twice and then succeeds is not left carrying a penalty.
|
||||
|
||||
const (
|
||||
// loginAttemptLimit is per email address per window. Five is comfortably
|
||||
// above human error and far below useful for guessing.
|
||||
loginAttemptLimit = 5
|
||||
// loginAddressLimit is per client address per window. Higher than the
|
||||
// per-email limit because one address legitimately covers a whole office
|
||||
// behind NAT, where several people may each fumble a password.
|
||||
loginAddressLimit = 20
|
||||
// loginAttemptWindow is how long a counter lives.
|
||||
loginAttemptWindow = 15 * time.Minute
|
||||
)
|
||||
|
||||
// attemptLimiter counts failures per key within a fixed window.
|
||||
type attemptLimiter struct {
|
||||
mu sync.Mutex
|
||||
limit int
|
||||
window time.Duration
|
||||
now func() time.Time
|
||||
buckets map[string]*attemptBucket
|
||||
}
|
||||
|
||||
type attemptBucket struct {
|
||||
count int
|
||||
resetAt time.Time
|
||||
}
|
||||
|
||||
func newAttemptLimiter(limit int, window time.Duration, now func() time.Time) *attemptLimiter {
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return &attemptLimiter{
|
||||
limit: limit, window: window, now: now,
|
||||
buckets: make(map[string]*attemptBucket),
|
||||
}
|
||||
}
|
||||
|
||||
// Allow reports whether another attempt may be made, and if not, how long the
|
||||
// caller should wait. It records nothing: only Fail does.
|
||||
//
|
||||
// Checking and recording are separate so a *successful* login never consumes
|
||||
// budget — the check happens before the password is verified, and the recording
|
||||
// only if it turns out to be wrong.
|
||||
func (l *attemptLimiter) Allow(key string) (bool, time.Duration) {
|
||||
if key == "" {
|
||||
return true, 0
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
b, ok := l.buckets[key]
|
||||
now := l.now()
|
||||
if !ok || !now.Before(b.resetAt) {
|
||||
return true, 0
|
||||
}
|
||||
if b.count < l.limit {
|
||||
return true, 0
|
||||
}
|
||||
return false, b.resetAt.Sub(now)
|
||||
}
|
||||
|
||||
// Fail records one failed attempt.
|
||||
func (l *attemptLimiter) Fail(key string) {
|
||||
if key == "" {
|
||||
return
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
now := l.now()
|
||||
l.pruneLocked(now)
|
||||
|
||||
b, ok := l.buckets[key]
|
||||
if !ok || !now.Before(b.resetAt) {
|
||||
l.buckets[key] = &attemptBucket{count: 1, resetAt: now.Add(l.window)}
|
||||
return
|
||||
}
|
||||
b.count++
|
||||
}
|
||||
|
||||
// Reset clears a key's counter. Called on a successful sign-in.
|
||||
func (l *attemptLimiter) Reset(key string) {
|
||||
if key == "" {
|
||||
return
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
delete(l.buckets, key)
|
||||
}
|
||||
|
||||
// pruneMinimum is the size below which pruning is not worth the walk.
|
||||
const pruneMinimum = 1024
|
||||
|
||||
// pruneLocked drops expired buckets once the map is large enough to be worth
|
||||
// walking. Called from Fail, which is the only path that grows the map.
|
||||
func (l *attemptLimiter) pruneLocked(now time.Time) {
|
||||
if len(l.buckets) < pruneMinimum {
|
||||
return
|
||||
}
|
||||
for key, b := range l.buckets {
|
||||
if !now.Before(b.resetAt) {
|
||||
delete(l.buckets, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// clientAddr is the key for per-address limiting.
|
||||
//
|
||||
// The port is stripped: a browser uses a new source port for every connection,
|
||||
// so keying on host:port would give each attempt its own budget and limit
|
||||
// nothing at all.
|
||||
func clientAddr(r *http.Request) string {
|
||||
host, _, err := net.SplitHostPort(strings.TrimSpace(r.RemoteAddr))
|
||||
if err != nil {
|
||||
return strings.TrimSpace(r.RemoteAddr)
|
||||
}
|
||||
return host
|
||||
}
|
||||
730
go-api/internal/httpserver/rbac_test.go
Normal file
730
go-api/internal/httpserver/rbac_test.go
Normal file
@@ -0,0 +1,730 @@
|
||||
package httpserver_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/httpserver"
|
||||
)
|
||||
|
||||
// Phase 3D authorization tests.
|
||||
//
|
||||
// Two questions are under test and they are deliberately kept apart, because
|
||||
// conflating them is how authorization bugs hide:
|
||||
//
|
||||
// MAY THIS ROLE CALL THIS ENDPOINT AT ALL? → checked in the handler, 403.
|
||||
// WHICH ROWS DOES THIS CALLER SEE? → a SQL predicate, so a row that
|
||||
// is not theirs is absent, 404.
|
||||
//
|
||||
// The row question is tested against the database rather than against a mock,
|
||||
// because the answer lives in a WHERE clause. A test that stubbed the
|
||||
// repository would prove the policy table is well-formed and nothing about
|
||||
// whether talent B can read talent A's application.
|
||||
|
||||
/* ── Fixture ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// rbac is one organization holding one of each role, a second employer and a
|
||||
// second talent to test isolation between peers, and a user in another
|
||||
// organization entirely.
|
||||
type rbac struct {
|
||||
*api
|
||||
admin, empA, empB, talA, talB actor
|
||||
|
||||
otherOrgID string
|
||||
outsider actor // admin in another organization
|
||||
|
||||
activePosting string
|
||||
draftPosting string
|
||||
}
|
||||
|
||||
func newRBAC(t *testing.T) *rbac {
|
||||
t.Helper()
|
||||
a := newAPI(t) // signs in as the seeded user, whose role is admin
|
||||
ctx := context.Background()
|
||||
r := &rbac{api: a}
|
||||
|
||||
r.admin = actor{name: "admin", id: a.userID, email: a.email, role: "admin", cookie: a.cookie}
|
||||
r.empA = signInAs(t, a.handler, a.h.Pool, a.orgID, "employerA", "employer-a@example.test", "employer")
|
||||
r.empB = signInAs(t, a.handler, a.h.Pool, a.orgID, "employerB", "employer-b@example.test", "employer")
|
||||
r.talA = signInAs(t, a.handler, a.h.Pool, a.orgID, "talentA", "talent-a@example.test", "talent")
|
||||
r.talB = signInAs(t, a.handler, a.h.Pool, a.orgID, "talentB", "talent-b@example.test", "talent")
|
||||
|
||||
if err := a.h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Other Tenant','other-tenant') RETURNING id::text`).
|
||||
Scan(&r.otherOrgID); err != nil {
|
||||
t.Fatalf("create the second organization: %v", err)
|
||||
}
|
||||
// An ADMIN in the other organization: cross-organization isolation must
|
||||
// hold on its own, without a role restriction doing the work for it.
|
||||
r.outsider = signInAs(t, a.handler, a.h.Pool, r.otherOrgID, "outsider", "outsider@example.test", "admin")
|
||||
|
||||
// One active posting and one draft, for the talent visibility rule.
|
||||
r.activePosting = createPosting(t, r, "Open Role", "active")
|
||||
r.draftPosting = createPosting(t, r, "Unannounced Role", "draft")
|
||||
return r
|
||||
}
|
||||
|
||||
func createPosting(t *testing.T, r *rbac, title, status string) string {
|
||||
t.Helper()
|
||||
got := r.as(r.admin, "POST", "/api/v1/job-postings", map[string]any{
|
||||
"title": title, "status": status,
|
||||
})
|
||||
if got.code != http.StatusCreated {
|
||||
t.Fatalf("create %s posting: %d (%v)", status, got.code, got.body)
|
||||
}
|
||||
return got.body["data"].(map[string]any)["id"].(string)
|
||||
}
|
||||
|
||||
func (r *rbac) ids(t *testing.T, act actor, path string) map[string]bool {
|
||||
t.Helper()
|
||||
got := r.as(act, "GET", path, nil)
|
||||
if got.code != http.StatusOK {
|
||||
t.Fatalf("%s GET %s = %d (%v)", act.name, path, got.code, got.body)
|
||||
}
|
||||
out := map[string]bool{}
|
||||
for _, rec := range got.records(t) {
|
||||
if id, ok := rec["id"].(string); ok {
|
||||
out[id] = true
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
/* ── 1. The role matrix ─────────────────────────────────────────────────── */
|
||||
|
||||
// Every endpoint against every role. The assertion is only about the role gate:
|
||||
// 403 means refused, anything else means the gate let the request through to be
|
||||
// judged on its merits. A 422 from a deliberately thin body still proves the
|
||||
// caller was allowed in, which is what this test is about.
|
||||
func TestRoleMatrix(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
type call struct {
|
||||
method, path string
|
||||
body any
|
||||
}
|
||||
// forbidden lists the roles that must be refused. Every other role must get
|
||||
// past the gate.
|
||||
cases := []struct {
|
||||
call
|
||||
forbidden []string
|
||||
}{
|
||||
{call{"GET", "/api/v1/job-postings", nil}, nil},
|
||||
{call{"GET", "/api/v1/job-postings/" + r.activePosting, nil}, nil},
|
||||
{call{"POST", "/api/v1/job-postings", map[string]any{"title": "X"}}, []string{"talent"}},
|
||||
{call{"PATCH", "/api/v1/job-postings/" + r.activePosting, map[string]any{"location": "Here"}}, []string{"talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/job-applications", nil}, nil},
|
||||
{call{"POST", "/api/v1/job-applications", map[string]any{
|
||||
"job_posting_id": r.activePosting, "applicant_name": "A", "email": "someone@example.test"}}, nil},
|
||||
{call{"PATCH", "/api/v1/job-applications/" + zeroUUID, map[string]any{"phone": "1"}}, []string{"talent"}},
|
||||
{call{"DELETE", "/api/v1/job-applications/" + zeroUUID, nil}, []string{"talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/ai-interviews", nil}, nil},
|
||||
{call{"POST", "/api/v1/ai-interviews", map[string]any{
|
||||
"application_id": zeroUUID, "job_posting_id": r.activePosting}}, nil},
|
||||
|
||||
{call{"GET", "/api/v1/staff", nil}, []string{"talent"}},
|
||||
{call{"POST", "/api/v1/staff", map[string]any{
|
||||
"name": "N", "email": "s@example.test", "hire_date": "2026-01-01"}}, []string{"talent"}},
|
||||
{call{"PATCH", "/api/v1/staff/" + zeroUUID, map[string]any{"phone": "1"}}, []string{"talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/worker-profiles", nil}, nil},
|
||||
{call{"POST", "/api/v1/worker-profiles", map[string]any{
|
||||
"full_name": "W", "email": "w@example.test"}}, nil},
|
||||
{call{"PATCH", "/api/v1/worker-profiles/" + zeroUUID, map[string]any{"phone": "1"}}, nil},
|
||||
|
||||
{call{"GET", "/api/v1/assignments", nil}, nil},
|
||||
{call{"POST", "/api/v1/assignments", map[string]any{
|
||||
"job_posting_id": r.activePosting, "worker_email": "w@example.test",
|
||||
"starts_at": "2026-01-01T00:00:00.000Z"}}, []string{"talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/shift-records", nil}, nil},
|
||||
|
||||
{call{"GET", "/api/v1/courses", nil}, nil},
|
||||
{call{"POST", "/api/v1/courses", map[string]any{"title": "C"}}, []string{"employer", "talent"}},
|
||||
{call{"PATCH", "/api/v1/courses/" + zeroUUID, map[string]any{"title": "C2"}}, []string{"employer", "talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/learning-paths", nil}, nil},
|
||||
|
||||
{call{"GET", "/api/v1/role-categories", nil}, nil},
|
||||
{call{"POST", "/api/v1/role-categories", map[string]any{"name": "RC"}}, []string{"talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/certifications", nil}, nil},
|
||||
{call{"POST", "/api/v1/certifications", map[string]any{"name": "Cert"}}, []string{"talent"}},
|
||||
{call{"DELETE", "/api/v1/certifications/" + zeroUUID, nil}, []string{"employer", "talent"}},
|
||||
|
||||
{call{"GET", "/api/v1/user-activity", nil}, nil},
|
||||
{call{"POST", "/api/v1/user-activity", map[string]any{"event_type": "test"}}, nil},
|
||||
|
||||
{call{"GET", "/api/v1/evidence", nil}, nil},
|
||||
{call{"POST", "/api/v1/evidence", map[string]any{"type": "photo_identify", "worker_email": "w@example.test"}}, nil},
|
||||
{call{"PATCH", "/api/v1/evidence/" + zeroUUID, map[string]any{"notes": "n"}}, []string{"talent"}},
|
||||
|
||||
// /me is every authenticated role's own business.
|
||||
{call{"GET", "/api/v1/me", nil}, nil},
|
||||
{call{"PATCH", "/api/v1/me", map[string]any{"full_name": "Renamed"}}, nil},
|
||||
{call{"GET", "/api/v1/me/preferences", nil}, nil},
|
||||
{call{"PATCH", "/api/v1/me/preferences", map[string]any{"emailDigest": true}}, nil},
|
||||
}
|
||||
|
||||
actors := map[string]actor{"admin": r.admin, "employer": r.empA, "talent": r.talA}
|
||||
|
||||
for _, tc := range cases {
|
||||
for role, act := range actors {
|
||||
name := fmt.Sprintf("%s %s as %s", tc.method, tc.path, role)
|
||||
t.Run(name, func(t *testing.T) {
|
||||
got := r.as(act, tc.method, tc.path, tc.body)
|
||||
denied := listsRole(tc.forbidden, role)
|
||||
|
||||
if denied {
|
||||
if got.code != http.StatusForbidden {
|
||||
t.Errorf("= %d (%s), want 403 forbidden", got.code, got.codeOrEmpty())
|
||||
}
|
||||
return
|
||||
}
|
||||
if got.code == http.StatusForbidden {
|
||||
t.Errorf("= 403, but %s should be allowed through the role gate", role)
|
||||
}
|
||||
if got.code == http.StatusUnauthorized {
|
||||
t.Errorf("= 401 — the session was rejected, which is not what this tests")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const zeroUUID = "00000000-0000-0000-0000-000000000000"
|
||||
|
||||
func listsRole(set []string, v string) bool {
|
||||
for _, s := range set {
|
||||
if s == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
/* ── 2. Ownership isolation between two talent users ────────────────────── */
|
||||
|
||||
// Talent A's records are invisible to talent B across every owned resource,
|
||||
// and visible to the organization's operators.
|
||||
func TestTalentSeesOnlyTheirOwnRecords(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
ctx := context.Background()
|
||||
|
||||
own := map[string]string{} // resource path → the id talent A owns
|
||||
|
||||
// Created through the API by talent A, so the ownership column is whatever
|
||||
// the server derived — not what the test asked for.
|
||||
own["worker-profiles"] = mustCreate(t, r, r.talA, "/api/v1/worker-profiles",
|
||||
map[string]any{"full_name": "Talent A", "email": r.talA.email})
|
||||
own["job-applications"] = mustCreate(t, r, r.talA, "/api/v1/job-applications",
|
||||
map[string]any{"job_posting_id": r.activePosting, "applicant_name": "Talent A"})
|
||||
own["evidence"] = mustCreate(t, r, r.talA, "/api/v1/evidence",
|
||||
map[string]any{"type": "photo_identify"})
|
||||
own["user-activity"] = mustCreate(t, r, r.talA, "/api/v1/user-activity",
|
||||
map[string]any{"event_type": "viewed_something"})
|
||||
own["ai-interviews"] = mustCreate(t, r, r.talA, "/api/v1/ai-interviews",
|
||||
map[string]any{"application_id": own["job-applications"], "job_posting_id": r.activePosting})
|
||||
|
||||
// Assignments are created by operators; shift records only by the seeder.
|
||||
own["assignments"] = mustCreate(t, r, r.admin, "/api/v1/assignments", map[string]any{
|
||||
"job_posting_id": r.activePosting, "worker_email": r.talA.email,
|
||||
"starts_at": "2026-01-01T00:00:00.000Z"})
|
||||
var shiftID string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO shift_records
|
||||
(org_id, worker_email, shift_date, scheduled_start, scheduled_end, scheduled_hours, created_date)
|
||||
VALUES ($1::uuid, $2::citext, '2026-01-02',
|
||||
'2026-01-02T09:00:00Z', '2026-01-02T17:00:00Z', 8, now())
|
||||
RETURNING id::text`, r.orgID, r.talA.email).Scan(&shiftID); err != nil {
|
||||
t.Fatalf("insert a shift record: %v", err)
|
||||
}
|
||||
own["shift-records"] = shiftID
|
||||
|
||||
// Talent B also has records of their own, so "B sees nothing" cannot pass
|
||||
// by the endpoint simply being broken.
|
||||
mustCreate(t, r, r.talB, "/api/v1/worker-profiles",
|
||||
map[string]any{"full_name": "Talent B", "email": r.talB.email})
|
||||
mustCreate(t, r, r.talB, "/api/v1/user-activity", map[string]any{"event_type": "b_event"})
|
||||
|
||||
for path, id := range own {
|
||||
t.Run(path, func(t *testing.T) {
|
||||
if !r.ids(t, r.talA, "/api/v1/"+path+"?limit=500")[id] {
|
||||
t.Errorf("talent A cannot see their own %s record", path)
|
||||
}
|
||||
if r.ids(t, r.talB, "/api/v1/"+path+"?limit=500")[id] {
|
||||
t.Errorf("talent B can see talent A's %s record", path)
|
||||
}
|
||||
if !r.ids(t, r.admin, "/api/v1/"+path+"?limit=500")[id] {
|
||||
t.Errorf("the organization's admin cannot see the %s record", path)
|
||||
}
|
||||
if !r.ids(t, r.empA, "/api/v1/"+path+"?limit=500")[id] {
|
||||
t.Errorf("the organization's employer cannot see the %s record", path)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The count must respect ownership too. A total computed over the whole
|
||||
// organization would leak how many records exist even with the rows hidden.
|
||||
t.Run("meta total respects ownership", func(t *testing.T) {
|
||||
got := r.as(r.talB, "GET", "/api/v1/worker-profiles?limit=500", nil)
|
||||
meta := got.meta(t)
|
||||
if n, _ := meta["total"].(float64); n != 1 {
|
||||
t.Errorf("talent B's worker-profiles total = %v, want 1 (their own)", meta["total"])
|
||||
}
|
||||
})
|
||||
|
||||
// Talent A cannot reach talent B's profile by PATCHing its id either: the
|
||||
// ownership predicate is in the UPDATE's WHERE clause, so the row is not
|
||||
// found rather than refused.
|
||||
t.Run("PATCH another talent's profile is 404", func(t *testing.T) {
|
||||
var bProfile string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT id::text FROM worker_profiles WHERE user_id = $1::uuid`, r.talB.id).Scan(&bProfile); err != nil {
|
||||
t.Fatalf("find talent B's profile: %v", err)
|
||||
}
|
||||
got := r.as(r.talA, "PATCH", "/api/v1/worker-profiles/"+bProfile, map[string]any{"phone": "hijacked"})
|
||||
if got.code != http.StatusNotFound {
|
||||
t.Errorf("= %d, want 404 (absent, not forbidden — existence must not leak)", got.code)
|
||||
}
|
||||
var phone string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT phone FROM worker_profiles WHERE id = $1::uuid`, bProfile).Scan(&phone); err != nil {
|
||||
t.Fatalf("re-read talent B's profile: %v", err)
|
||||
}
|
||||
if phone == "hijacked" {
|
||||
t.Fatal("talent A modified talent B's worker profile")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func mustCreate(t *testing.T, r *rbac, act actor, path string, body map[string]any) string {
|
||||
t.Helper()
|
||||
got := r.as(act, "POST", path, body)
|
||||
if got.code != http.StatusCreated {
|
||||
t.Fatalf("%s POST %s = %d (%v)", act.name, path, got.code, got.body)
|
||||
}
|
||||
return got.body["data"].(map[string]any)["id"].(string)
|
||||
}
|
||||
|
||||
/* ── 3. Mass assignment ─────────────────────────────────────────────────── */
|
||||
|
||||
// Identity a caller supplies is ignored; identity the server derives wins.
|
||||
//
|
||||
// This is the test that makes the ownership predicates above mean anything. If
|
||||
// a talent user could name someone else in the ownership column, every "own
|
||||
// records only" rule would be bypassable by the same request it constrains.
|
||||
func TestServerOwnedIdentityCannotBeSupplied(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("worker_profiles.user_id", func(t *testing.T) {
|
||||
id := mustCreate(t, r, r.talA, "/api/v1/worker-profiles", map[string]any{
|
||||
"full_name": "Claimed", "email": r.talA.email,
|
||||
"user_id": r.talB.id, // naming somebody else
|
||||
})
|
||||
var owner string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT COALESCE(user_id::text,'') FROM worker_profiles WHERE id = $1::uuid`, id).Scan(&owner); err != nil {
|
||||
t.Fatalf("read the profile: %v", err)
|
||||
}
|
||||
if owner != r.talA.id {
|
||||
t.Errorf("user_id = %q, want the creating talent %q", owner, r.talA.id)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("worker_profiles.user_id is NOT the admin when an operator creates one", func(t *testing.T) {
|
||||
// The subject of an operator-created profile is a candidate, not the
|
||||
// operator. Deriving it unconditionally would file every candidate's
|
||||
// record under whoever typed it in.
|
||||
id := mustCreate(t, r, r.admin, "/api/v1/worker-profiles", map[string]any{
|
||||
"full_name": "Candidate", "email": "candidate@example.test",
|
||||
})
|
||||
var owner string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT COALESCE(user_id::text,'') FROM worker_profiles WHERE id = $1::uuid`, id).Scan(&owner); err != nil {
|
||||
t.Fatalf("read the profile: %v", err)
|
||||
}
|
||||
if owner != "" {
|
||||
t.Errorf("user_id = %q, want empty — an operator-created profile has no claimant yet", owner)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("job_applications.email", func(t *testing.T) {
|
||||
id := mustCreate(t, r, r.talA, "/api/v1/job-applications", map[string]any{
|
||||
"job_posting_id": r.activePosting, "applicant_name": "A",
|
||||
"email": r.talB.email, // applying as somebody else
|
||||
})
|
||||
var email string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT email::text FROM job_applications WHERE id = $1::uuid`, id).Scan(&email); err != nil {
|
||||
t.Fatalf("read the application: %v", err)
|
||||
}
|
||||
if email != r.talA.email {
|
||||
t.Errorf("email = %q, want the applying talent %q", email, r.talA.email)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("evidence.worker_email", func(t *testing.T) {
|
||||
id := mustCreate(t, r, r.talA, "/api/v1/evidence", map[string]any{
|
||||
"type": "photo_identify", "worker_email": r.talB.email,
|
||||
})
|
||||
var email string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT worker_email::text FROM evidence WHERE id = $1::uuid`, id).Scan(&email); err != nil {
|
||||
t.Fatalf("read the evidence: %v", err)
|
||||
}
|
||||
if email != r.talA.email {
|
||||
t.Errorf("worker_email = %q, want %q", email, r.talA.email)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("user_activity identity is entirely server-derived", func(t *testing.T) {
|
||||
id := mustCreate(t, r, r.talA, "/api/v1/user-activity", map[string]any{
|
||||
"event_type": "forged",
|
||||
"user_id": r.admin.id,
|
||||
"user_email": r.admin.email,
|
||||
"user_name": "The Administrator",
|
||||
"account_type": "admin",
|
||||
})
|
||||
var uid, email, name, acct string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT COALESCE(user_id::text,''), user_email::text, user_name, account_type
|
||||
FROM user_activity WHERE id::text = $1`, id).Scan(&uid, &email, &name, &acct); err != nil {
|
||||
t.Fatalf("read the activity row: %v", err)
|
||||
}
|
||||
if uid != r.talA.id || email != r.talA.email {
|
||||
t.Errorf("activity attributed to %s/%s, want talent A %s/%s", uid, email, r.talA.id, r.talA.email)
|
||||
}
|
||||
if name == "The Administrator" || acct == "admin" {
|
||||
t.Errorf("client-supplied user_name/account_type were stored: %q / %q", name, acct)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("job_postings.created_by", func(t *testing.T) {
|
||||
got := r.as(r.empA, "POST", "/api/v1/job-postings", map[string]any{
|
||||
"title": "Attributed", "created_by": r.admin.id,
|
||||
})
|
||||
if got.code != http.StatusCreated {
|
||||
t.Fatalf("create = %d (%v)", got.code, got.body)
|
||||
}
|
||||
id := got.body["data"].(map[string]any)["id"].(string)
|
||||
var by string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT COALESCE(created_by::text,'') FROM job_postings WHERE id = $1::uuid`, id).Scan(&by); err != nil {
|
||||
t.Fatalf("read the posting: %v", err)
|
||||
}
|
||||
if by != r.empA.id {
|
||||
t.Errorf("created_by = %q, want the actual creator %q", by, r.empA.id)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("org_id and role still cannot be supplied", func(t *testing.T) {
|
||||
id := mustCreate(t, r, r.empA, "/api/v1/job-postings", map[string]any{
|
||||
"title": "Tenancy", "org_id": r.otherOrgID,
|
||||
})
|
||||
var org string
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT org_id::text FROM job_postings WHERE id = $1::uuid`, id).Scan(&org); err != nil {
|
||||
t.Fatalf("read the posting: %v", err)
|
||||
}
|
||||
if org != r.orgID {
|
||||
t.Errorf("org_id = %q, want the session's organization %q", org, r.orgID)
|
||||
}
|
||||
|
||||
// And a talent cannot promote themselves through /me.
|
||||
if got := r.as(r.talA, "PATCH", "/api/v1/me", map[string]any{"role": "admin"}); got.code != http.StatusOK {
|
||||
t.Fatalf("PATCH /me = %d", got.code)
|
||||
}
|
||||
var role string
|
||||
if err := r.h.Pool.QueryRow(ctx, `SELECT role FROM users WHERE id = $1::uuid`, r.talA.id).Scan(&role); err != nil {
|
||||
t.Fatalf("read the user: %v", err)
|
||||
}
|
||||
if role != "talent" {
|
||||
t.Fatalf("role = %q — a talent user promoted themselves", role)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// A talent user cannot attach an interview to somebody else's application.
|
||||
// Ownership here is by reference, so it is checked against the application.
|
||||
func TestTalentCannotInterviewForAnotherApplication(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
othersApplication := mustCreate(t, r, r.talB, "/api/v1/job-applications",
|
||||
map[string]any{"job_posting_id": r.activePosting, "applicant_name": "Talent B"})
|
||||
|
||||
got := r.as(r.talA, "POST", "/api/v1/ai-interviews", map[string]any{
|
||||
"application_id": othersApplication, "job_posting_id": r.activePosting,
|
||||
})
|
||||
if got.code != http.StatusNotFound {
|
||||
t.Errorf("= %d (%s), want 404 — the same answer an application that does not exist gives",
|
||||
got.code, got.codeOrEmpty())
|
||||
}
|
||||
|
||||
// Their own application is accepted, so the guard is not simply refusing
|
||||
// everything.
|
||||
mine := mustCreate(t, r, r.talA, "/api/v1/job-applications",
|
||||
map[string]any{"job_posting_id": r.activePosting, "applicant_name": "Talent A"})
|
||||
if ok := r.as(r.talA, "POST", "/api/v1/ai-interviews", map[string]any{
|
||||
"application_id": mine, "job_posting_id": r.activePosting,
|
||||
}); ok.code != http.StatusCreated {
|
||||
t.Errorf("interviewing for their own application = %d (%v)", ok.code, ok.body)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4. Talent posting visibility ───────────────────────────────────────── */
|
||||
|
||||
func TestTalentSeesOnlyActivePostings(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
talent := r.ids(t, r.talA, "/api/v1/job-postings?limit=200")
|
||||
if !talent[r.activePosting] {
|
||||
t.Error("talent cannot see an active posting")
|
||||
}
|
||||
if talent[r.draftPosting] {
|
||||
t.Error("talent can see a draft posting")
|
||||
}
|
||||
|
||||
for _, act := range []actor{r.admin, r.empA} {
|
||||
seen := r.ids(t, act, "/api/v1/job-postings?limit=200")
|
||||
if !seen[r.draftPosting] {
|
||||
t.Errorf("%s cannot see the organization's draft posting", act.name)
|
||||
}
|
||||
}
|
||||
|
||||
// By id, too — and as a 404, so the draft's existence is not disclosed.
|
||||
if got := r.as(r.talA, "GET", "/api/v1/job-postings/"+r.draftPosting, nil); got.code != http.StatusNotFound {
|
||||
t.Errorf("talent GET of a draft posting = %d, want 404", got.code)
|
||||
}
|
||||
if got := r.as(r.talA, "GET", "/api/v1/job-postings/"+r.activePosting, nil); got.code != http.StatusOK {
|
||||
t.Errorf("talent GET of an active posting = %d, want 200", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 5. Cross-organization isolation ────────────────────────────────────── */
|
||||
|
||||
// The outsider is an ADMIN in another organization, so nothing here is being
|
||||
// done by a role restriction.
|
||||
func TestCrossOrganizationIsolation(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
ctx := context.Background()
|
||||
|
||||
appID := mustCreate(t, r, r.admin, "/api/v1/job-applications", map[string]any{
|
||||
"job_posting_id": r.activePosting, "applicant_name": "Insider", "email": "insider@example.test"})
|
||||
|
||||
t.Run("cannot read", func(t *testing.T) {
|
||||
if r.ids(t, r.outsider, "/api/v1/job-postings?limit=200")[r.activePosting] {
|
||||
t.Error("an outsider can list another organization's posting")
|
||||
}
|
||||
if got := r.as(r.outsider, "GET", "/api/v1/job-postings/"+r.activePosting, nil); got.code != http.StatusNotFound {
|
||||
t.Errorf("GET by id = %d, want 404", got.code)
|
||||
}
|
||||
if n := len(r.ids(t, r.outsider, "/api/v1/job-applications?limit=200")); n != 0 {
|
||||
t.Errorf("an outsider sees %d applications from another organization", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cannot update", func(t *testing.T) {
|
||||
got := r.as(r.outsider, "PATCH", "/api/v1/job-postings/"+r.activePosting,
|
||||
map[string]any{"title": "Hijacked"})
|
||||
if got.code != http.StatusNotFound {
|
||||
t.Errorf("= %d, want 404", got.code)
|
||||
}
|
||||
var title string
|
||||
if err := r.h.Pool.QueryRow(ctx, `SELECT title FROM job_postings WHERE id = $1::uuid`,
|
||||
r.activePosting).Scan(&title); err != nil {
|
||||
t.Fatalf("re-read: %v", err)
|
||||
}
|
||||
if title == "Hijacked" {
|
||||
t.Fatal("an outsider modified another organization's posting")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cannot delete", func(t *testing.T) {
|
||||
// DELETE reports success whether or not a row matched — a deliberate
|
||||
// contract choice (§12.7) that reveals nothing. What matters is that
|
||||
// the row survives.
|
||||
r.as(r.outsider, "DELETE", "/api/v1/job-applications/"+appID, nil)
|
||||
var alive int
|
||||
if err := r.h.Pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM job_applications WHERE id = $1::uuid`, appID).Scan(&alive); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if alive != 1 {
|
||||
t.Fatal("an outsider deleted another organization's application")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── 6. 403 versus 404 ──────────────────────────────────────────────────── */
|
||||
|
||||
// The discipline: a refused ROLE is 403; a row outside the caller's visibility
|
||||
// is 404, whether it is another tenant's or another person's.
|
||||
func TestForbiddenVersusNotFound(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
t.Run("role refused is 403", func(t *testing.T) {
|
||||
got := r.as(r.talA, "GET", "/api/v1/staff", nil)
|
||||
if got.code != http.StatusForbidden || got.codeOrEmpty() != "forbidden" {
|
||||
t.Errorf("= %d (%s), want 403 forbidden", got.code, got.codeOrEmpty())
|
||||
}
|
||||
// And the message must not name the roles that would have worked.
|
||||
body, _ := got.body["error"].(map[string]any)
|
||||
msg, _ := body["message"].(string)
|
||||
for _, leak := range []string{"admin", "employer", "talent", "role"} {
|
||||
if containsFold(msg, leak) {
|
||||
t.Errorf("the 403 message names %q: %q", leak, msg)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("another tenant's row is 404", func(t *testing.T) {
|
||||
if got := r.as(r.outsider, "GET", "/api/v1/job-postings/"+r.activePosting, nil); got.code != http.StatusNotFound {
|
||||
t.Errorf("= %d, want 404", got.code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("another person's row is 404", func(t *testing.T) {
|
||||
bProfile := mustCreate(t, r, r.talB, "/api/v1/worker-profiles",
|
||||
map[string]any{"full_name": "B", "email": r.talB.email})
|
||||
if got := r.as(r.talA, "PATCH", "/api/v1/worker-profiles/"+bProfile,
|
||||
map[string]any{"phone": "x"}); got.code != http.StatusNotFound {
|
||||
t.Errorf("= %d, want 404", got.code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unauthenticated is still 401", func(t *testing.T) {
|
||||
if got := r.doAnon("GET", "/api/v1/staff", nil); got.code != http.StatusUnauthorized {
|
||||
t.Errorf("= %d, want 401", got.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func containsFold(haystack, needle string) bool {
|
||||
h, n := []rune(haystack), []rune(needle)
|
||||
lower := func(r rune) rune {
|
||||
if r >= 'A' && r <= 'Z' {
|
||||
return r + 32
|
||||
}
|
||||
return r
|
||||
}
|
||||
for i := 0; i+len(n) <= len(h); i++ {
|
||||
ok := true
|
||||
for j := range n {
|
||||
if lower(h[i+j]) != lower(n[j]) {
|
||||
ok = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
/* ── 7. Admin regression ────────────────────────────────────────────────── */
|
||||
|
||||
// Everything the admin console does today must still work. The endpoints below
|
||||
// are the ones the frontend actually calls, taken from the Phase 3D audit's
|
||||
// call-site inventory.
|
||||
func TestAdminRegression(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
for _, path := range []string{
|
||||
"job-postings", "job-applications", "ai-interviews", "staff", "worker-profiles",
|
||||
"courses", "learning-paths", "certifications", "role-categories",
|
||||
"user-activity", "evidence", "assignments", "shift-records",
|
||||
} {
|
||||
if got := r.as(r.admin, "GET", "/api/v1/"+path+"?limit=5", nil); got.code != http.StatusOK {
|
||||
t.Errorf("admin GET /api/v1/%s = %d (%v)", path, got.code, got.body)
|
||||
}
|
||||
}
|
||||
|
||||
// The seeded dataset is still fully visible to an admin: ownership scoping
|
||||
// must not have narrowed the operator view.
|
||||
if n := len(r.ids(t, r.admin, "/api/v1/job-postings?limit=200")); n < 8 {
|
||||
t.Errorf("admin sees %d job postings, want at least the 8 seeded", n)
|
||||
}
|
||||
|
||||
// A representative write of each shape.
|
||||
posting := mustCreate(t, r, r.admin, "/api/v1/job-postings", map[string]any{"title": "Admin Wrote This"})
|
||||
if got := r.as(r.admin, "PATCH", "/api/v1/job-postings/"+posting,
|
||||
map[string]any{"location": "Somewhere"}); got.code != http.StatusOK {
|
||||
t.Errorf("admin PATCH = %d (%v)", got.code, got.body)
|
||||
}
|
||||
app := mustCreate(t, r, r.admin, "/api/v1/job-applications", map[string]any{
|
||||
"job_posting_id": posting, "applicant_name": "C", "email": "c@example.test"})
|
||||
if got := r.as(r.admin, "DELETE", "/api/v1/job-applications/"+app, nil); got.code != http.StatusOK {
|
||||
t.Errorf("admin DELETE = %d", got.code)
|
||||
}
|
||||
if got := r.as(r.admin, "GET", "/api/v1/me", nil); got.code != http.StatusOK {
|
||||
t.Errorf("admin GET /me = %d", got.code)
|
||||
}
|
||||
if got := r.doAnon("GET", "/health", nil); got.code != http.StatusOK {
|
||||
t.Errorf("GET /health = %d, want 200 and still public", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8. Employer boundaries ─────────────────────────────────────────────── */
|
||||
|
||||
func TestEmployerBoundaries(t *testing.T) {
|
||||
r := newRBAC(t)
|
||||
|
||||
// Employer runs the organization's hiring: the operator surface works.
|
||||
for _, path := range []string{"job-postings", "job-applications", "staff", "worker-profiles", "user-activity"} {
|
||||
if got := r.as(r.empA, "GET", "/api/v1/"+path+"?limit=5", nil); got.code != http.StatusOK {
|
||||
t.Errorf("employer GET /api/v1/%s = %d", path, got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// Admin-only operations are refused. Course authoring is admin's because a
|
||||
// NULL-org course is the shared platform library and reaches every tenant.
|
||||
for _, tc := range []struct{ method, path string }{
|
||||
{"POST", "/api/v1/courses"},
|
||||
{"PATCH", "/api/v1/courses/" + zeroUUID},
|
||||
{"DELETE", "/api/v1/certifications/" + zeroUUID},
|
||||
} {
|
||||
got := r.as(r.empA, tc.method, tc.path, map[string]any{"title": "X"})
|
||||
if got.code != http.StatusForbidden {
|
||||
t.Errorf("employer %s %s = %d, want 403", tc.method, tc.path, got.code)
|
||||
}
|
||||
}
|
||||
|
||||
// Two employers in one organization see the same rows: the ownership
|
||||
// predicate must not have leaked onto the operator roles.
|
||||
posting := mustCreate(t, r, r.empA, "/api/v1/job-postings", map[string]any{"title": "By A"})
|
||||
if !r.ids(t, r.empB, "/api/v1/job-postings?limit=200")[posting] {
|
||||
t.Error("employer B cannot see employer A's posting — operators share the organization")
|
||||
}
|
||||
if got := r.as(r.empB, "PATCH", "/api/v1/job-postings/"+posting,
|
||||
map[string]any{"location": "Edited by B"}); got.code != http.StatusOK {
|
||||
t.Errorf("employer B editing employer A's posting = %d, want 200", got.code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 9. Session expiry still governs everything ─────────────────────────── */
|
||||
|
||||
// Authorization does not replace authentication: an expired session is refused
|
||||
// before any role is consulted.
|
||||
func TestExpiredSessionIsRefusedBeforeRoleCheck(t *testing.T) {
|
||||
now := time.Date(2026, 8, 22, 9, 0, 0, 0, time.UTC)
|
||||
a := newAPI(t,
|
||||
httpserver.WithClock(func() time.Time { return now }),
|
||||
httpserver.WithSessionPolicy(shortSessions))
|
||||
|
||||
if got := a.do("GET", "/api/v1/job-postings", nil); got.code != http.StatusOK {
|
||||
t.Fatalf("while live = %d", got.code)
|
||||
}
|
||||
now = now.Add(shortSessions.IdleLifetime + time.Minute)
|
||||
got := a.do("GET", "/api/v1/job-postings", nil)
|
||||
if got.code != http.StatusUnauthorized {
|
||||
t.Errorf("= %d (%s), want 401 — not 403", got.code, got.codeOrEmpty())
|
||||
}
|
||||
}
|
||||
110
go-api/internal/httpserver/response.go
Normal file
110
go-api/internal/httpserver/response.go
Normal file
@@ -0,0 +1,110 @@
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
)
|
||||
|
||||
// envelope is the success shape from api-contract.md §4.
|
||||
type envelope struct {
|
||||
Data any `json:"data"`
|
||||
Meta *meta `json:"meta,omitempty"`
|
||||
}
|
||||
|
||||
// meta accompanies a collection. `truncated` exists so the silent-truncation
|
||||
// problem in §12.2 is fixable without another contract change.
|
||||
type meta struct {
|
||||
Total int `json:"total"`
|
||||
Limit int `json:"limit"`
|
||||
Offset int `json:"offset"`
|
||||
Returned int `json:"returned"`
|
||||
Truncated bool `json:"truncated"`
|
||||
}
|
||||
|
||||
// errorEnvelope is the failure shape from api-contract.md §5.
|
||||
type errorEnvelope struct {
|
||||
Error errorBody `json:"error"`
|
||||
}
|
||||
|
||||
type errorBody struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Details map[string]string `json:"details"`
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, code int, body any) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(code)
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
_ = enc.Encode(body)
|
||||
}
|
||||
|
||||
func writeRecord(w http.ResponseWriter, code int, rec domain.Record) {
|
||||
writeJSON(w, code, envelope{Data: rec})
|
||||
}
|
||||
|
||||
func writePage(w http.ResponseWriter, page *domain.Page) {
|
||||
writeJSON(w, http.StatusOK, envelope{
|
||||
Data: page.Records,
|
||||
Meta: &meta{
|
||||
Total: page.Total,
|
||||
Limit: page.Limit,
|
||||
Offset: page.Offset,
|
||||
Returned: len(page.Records),
|
||||
Truncated: page.Total > page.Offset+len(page.Records),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// statusFor maps the contract's error codes onto HTTP status codes.
|
||||
func statusFor(code string) int {
|
||||
switch code {
|
||||
case "unauthorized":
|
||||
return http.StatusUnauthorized
|
||||
case "forbidden":
|
||||
return http.StatusForbidden
|
||||
case "rate_limited":
|
||||
return http.StatusTooManyRequests
|
||||
case "not_found":
|
||||
return http.StatusNotFound
|
||||
case "validation_failed":
|
||||
return http.StatusUnprocessableEntity
|
||||
case "invalid_query":
|
||||
return http.StatusBadRequest
|
||||
case "conflict":
|
||||
return http.StatusConflict
|
||||
default:
|
||||
return http.StatusInternalServerError
|
||||
}
|
||||
}
|
||||
|
||||
// writeError renders any error as the documented envelope.
|
||||
//
|
||||
// An unrecognised error is deliberately flattened to a generic message: the
|
||||
// detail goes to the log, not to the client.
|
||||
func writeError(w http.ResponseWriter, log *slog.Logger, err error) {
|
||||
var de *domain.Error
|
||||
if !errors.As(err, &de) {
|
||||
log.Error("unhandled error", "error", err)
|
||||
writeJSON(w, http.StatusInternalServerError, errorEnvelope{Error: errorBody{
|
||||
Code: "internal", Message: "internal error", Details: map[string]string{},
|
||||
}})
|
||||
return
|
||||
}
|
||||
if de.Code == "internal" {
|
||||
log.Error("internal error", "error", de.Unwrap())
|
||||
}
|
||||
details := de.Details
|
||||
if details == nil {
|
||||
details = map[string]string{}
|
||||
}
|
||||
writeJSON(w, statusFor(de.Code), errorEnvelope{Error: errorBody{
|
||||
Code: de.Code, Message: de.Message, Details: details,
|
||||
}})
|
||||
}
|
||||
383
go-api/internal/httpserver/server.go
Normal file
383
go-api/internal/httpserver/server.go
Normal file
@@ -0,0 +1,383 @@
|
||||
// Package httpserver holds the HTTP surface.
|
||||
//
|
||||
// It serves /health, the sign-in endpoints, the entity endpoints described in
|
||||
// docs/api-contract.md, and the current-user endpoints.
|
||||
//
|
||||
// Phase 3C replaced the development identity with real authentication. Every
|
||||
// request outside the small public allowlist in auth.go must carry a session
|
||||
// cookie; the middleware resolves it to a user row and puts that user, and
|
||||
// their organization, on the request context. Nothing downstream changed —
|
||||
// every service and repository already took the organization as a parameter,
|
||||
// which is what devOrgMiddleware existed to make true.
|
||||
//
|
||||
// Authorization is NOT here. A signed-in user reaches every endpoint they could
|
||||
// reach before; deciding which roles may do what is Phase 3D.
|
||||
package httpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/config"
|
||||
"github.com/krow/krow-backend/go-api/internal/db"
|
||||
"github.com/krow/krow-backend/go-api/internal/service"
|
||||
)
|
||||
|
||||
// Server binds the router, the pool, authentication and the lifecycle together.
|
||||
type Server struct {
|
||||
cfg *config.Config
|
||||
db *db.DB
|
||||
api *service.Registry
|
||||
definitions *service.DefinitionsService
|
||||
log *slog.Logger
|
||||
http *http.Server
|
||||
started time.Time
|
||||
endpoints int
|
||||
|
||||
// The authentication surface. sessions owns the lifecycle, users is the
|
||||
// read side of the users table, credentials verifies a password against it,
|
||||
// and the two limiters bound how often that may be attempted.
|
||||
//
|
||||
// Two limiters, not one, because the budgets are different sizes on
|
||||
// purpose: an email is one account and gets a tight budget, while an
|
||||
// address may be a whole office behind NAT and gets a loose one. Sharing a
|
||||
// limiter would force the office to live within one person's budget.
|
||||
sessions *auth.Manager
|
||||
users auth.UserStore
|
||||
credentials *auth.Credentials
|
||||
loginByEmail *attemptLimiter
|
||||
loginByAddr *attemptLimiter
|
||||
|
||||
// now is injectable so tests can drive expiry without sleeping.
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
// Option adjusts the server before it is wired. Production passes none.
|
||||
type Option func(*serverOptions)
|
||||
|
||||
type serverOptions struct {
|
||||
policy auth.Policy
|
||||
now func() time.Time
|
||||
perEmail int
|
||||
perAddress int
|
||||
loginWindow time.Duration
|
||||
}
|
||||
|
||||
// WithSessionPolicy overrides the session lifetimes. For tests that need to
|
||||
// reach an expiry without waiting twelve hours for it.
|
||||
func WithSessionPolicy(p auth.Policy) Option {
|
||||
return func(o *serverOptions) { o.policy = p }
|
||||
}
|
||||
|
||||
// WithClock replaces the clock used for session expiry and last_login_at.
|
||||
func WithClock(now func() time.Time) Option {
|
||||
return func(o *serverOptions) {
|
||||
if now != nil {
|
||||
o.now = now
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithLoginRateLimit overrides the failed-attempt budgets and their window.
|
||||
//
|
||||
// perEmail bounds attempts against one account; perAddress bounds attempts from
|
||||
// one client address across all accounts. Both are consulted on every attempt.
|
||||
func WithLoginRateLimit(perEmail, perAddress int, window time.Duration) Option {
|
||||
return func(o *serverOptions) {
|
||||
o.perEmail, o.perAddress, o.loginWindow = perEmail, perAddress, window
|
||||
}
|
||||
}
|
||||
|
||||
// New wires the routes and returns a server that has not yet been started.
|
||||
//
|
||||
// Authentication is built here rather than passed in, so there is exactly one
|
||||
// construction of the session manager and no way to start a server with the
|
||||
// middleware wired to a different store than the login handler.
|
||||
func New(cfg *config.Config, database *db.DB, log *slog.Logger, opts ...Option) (*Server, error) {
|
||||
o := serverOptions{
|
||||
policy: auth.DefaultPolicy,
|
||||
now: time.Now,
|
||||
perEmail: loginAttemptLimit,
|
||||
perAddress: loginAddressLimit,
|
||||
loginWindow: loginAttemptWindow,
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(&o)
|
||||
}
|
||||
|
||||
sessions, err := auth.NewManager(auth.NewPGStore(database.Pool), o.policy)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build session manager: %w", err)
|
||||
}
|
||||
sessions.WithClock(o.now)
|
||||
|
||||
users := auth.NewPGUserStore(database.Pool)
|
||||
s := &Server{
|
||||
cfg: cfg, db: database, log: log,
|
||||
api: service.NewRegistry(database.Pool),
|
||||
definitions: service.NewDefinitions(database.Pool),
|
||||
started: o.now(),
|
||||
sessions: sessions,
|
||||
users: users,
|
||||
credentials: auth.NewCredentials(users),
|
||||
loginByEmail: newAttemptLimiter(o.perEmail, o.loginWindow, o.now),
|
||||
loginByAddr: newAttemptLimiter(o.perAddress, o.loginWindow, o.now),
|
||||
now: o.now,
|
||||
}
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("GET /health", s.handleHealth)
|
||||
s.endpoints = s.routeAuth(mux) + s.routeResources(mux) + s.routeMe(mux) + s.routeDefinitions(mux)
|
||||
|
||||
handler := jsonErrors(mux)
|
||||
// Authentication sits where devOrgMiddleware used to, so every route below
|
||||
// it — including the mux's own 404 — is behind the allowlist.
|
||||
handler = s.authenticate(handler)
|
||||
handler = recoverer(log)(handler)
|
||||
// CORS sits outside the recoverer so a preflight is answered without
|
||||
// touching the router, and inside the logger so refused origins are still
|
||||
// visible in the log. With no allowlist configured it is not installed at
|
||||
// all, which is the same-origin default.
|
||||
if len(cfg.HTTP.CORSOrigins) > 0 {
|
||||
handler = cors(cfg.HTTP.CORSOrigins)(handler)
|
||||
}
|
||||
handler = requestLogger(log)(handler)
|
||||
|
||||
s.http = &http.Server{
|
||||
Addr: net.JoinHostPort(cfg.HTTP.Host, strconv.Itoa(cfg.HTTP.Port)),
|
||||
Handler: handler,
|
||||
ReadTimeout: cfg.HTTP.ReadTimeout,
|
||||
WriteTimeout: cfg.HTTP.WriteTimeout,
|
||||
IdleTimeout: cfg.HTTP.IdleTimeout,
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Sessions exposes the session manager, so the process can sweep expired rows
|
||||
// and tests can drive the clock.
|
||||
func (s *Server) Sessions() *auth.Manager { return s.sessions }
|
||||
|
||||
// Handler exposes the routed handler so tests can drive it without a listener.
|
||||
func (s *Server) Handler() http.Handler { return s.http.Handler }
|
||||
|
||||
// Endpoints is how many routes were registered.
|
||||
func (s *Server) Endpoints() int { return s.endpoints }
|
||||
|
||||
// Addr is the address the server listens on.
|
||||
func (s *Server) Addr() string { return s.http.Addr }
|
||||
|
||||
// Start blocks until the server stops accepting connections.
|
||||
func (s *Server) Start() error {
|
||||
err := s.http.ListenAndServe()
|
||||
if errors.Is(err, http.ErrServerClosed) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Shutdown drains in-flight requests, then gives up after the configured grace.
|
||||
func (s *Server) Shutdown(ctx context.Context) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, s.cfg.HTTP.ShutdownTimeout)
|
||||
defer cancel()
|
||||
return s.http.Shutdown(ctx)
|
||||
}
|
||||
|
||||
// healthResponse is the entire public /health body: one field, deliberately.
|
||||
//
|
||||
// /health is unauthenticated and reachable by anyone who can reach the port,
|
||||
// so it is treated as a public document rather than as an operator's console.
|
||||
// Everything an unauthenticated caller legitimately needs is the answer to
|
||||
// "should traffic be sent here", and that fits in a status string plus the
|
||||
// HTTP status code.
|
||||
//
|
||||
// What used to be here and is now deliberately absent: the PostgreSQL version,
|
||||
// the database name, the schema name, the applied migration version, the table
|
||||
// count, the connection error text, the deployment environment and the process
|
||||
// uptime. Individually each is small; together they are a free reconnaissance
|
||||
// report — the server version to look up known CVEs against, the migration
|
||||
// version to date the deployment, the table count and error text to infer
|
||||
// shape and topology. None of it is diagnostic to anyone who could not already
|
||||
// read it from the database directly.
|
||||
//
|
||||
// The check itself is unchanged. db.Check still runs on every request and
|
||||
// still decides the answer; its full detail now goes to the server log, where
|
||||
// the operator is, instead of into the response, where the internet is. See
|
||||
// logHealth.
|
||||
type healthResponse struct {
|
||||
Status string `json:"status"`
|
||||
}
|
||||
|
||||
// handleHealth reports whether this instance should be sent traffic.
|
||||
//
|
||||
// 200 "ok" serving normally
|
||||
// 200 "degraded" the process is healthy, the schema is not: unmigrated,
|
||||
// or a migration left the version dirty. Still 200,
|
||||
// because the fault is the database's and taking the
|
||||
// instance out of rotation would not fix it.
|
||||
// 503 "unavailable" the database is unreachable, so a load balancer can act
|
||||
// on the status code alone without parsing the body.
|
||||
//
|
||||
// The three status words are a coarse operational signal, not infrastructure
|
||||
// detail: they say what a caller should do, and nothing about what is running.
|
||||
func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) {
|
||||
ctx, cancel := context.WithTimeout(r.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
health := s.db.Check(ctx)
|
||||
|
||||
status, code := "ok", http.StatusOK
|
||||
switch {
|
||||
case !health.Reachable:
|
||||
status, code = "unavailable", http.StatusServiceUnavailable
|
||||
case health.MigrationDirty, !health.SchemaPresent:
|
||||
status = "degraded"
|
||||
}
|
||||
|
||||
s.logHealth(status, health)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(code)
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
_ = enc.Encode(healthResponse{Status: status})
|
||||
}
|
||||
|
||||
// logHealth writes the detail the response body used to carry.
|
||||
//
|
||||
// This is the "internally" half of the change: nothing was deleted from
|
||||
// db.Check, and nothing it learns is thrown away — the audience moved from the
|
||||
// response to the log, which is already authenticated by virtue of being on
|
||||
// the host.
|
||||
//
|
||||
// A load balancer polls this endpoint every few seconds, so a healthy check
|
||||
// logs at debug and a bad one at warn. Anything other than "ok" is worth
|
||||
// seeing without turning debug on.
|
||||
func (s *Server) logHealth(status string, h db.Health) {
|
||||
attrs := []any{
|
||||
"status", status,
|
||||
"env", s.cfg.AppEnv,
|
||||
"uptime_seconds", int64(time.Since(s.started).Seconds()),
|
||||
"reachable", h.Reachable,
|
||||
"schema", h.Schema,
|
||||
"schema_present", h.SchemaPresent,
|
||||
"table_count", h.TableCount,
|
||||
"migration_dirty", h.MigrationDirty,
|
||||
"latency_ms", h.LatencyMS,
|
||||
}
|
||||
if h.Database != "" {
|
||||
attrs = append(attrs, "database", h.Database, "postgres_version", h.Version)
|
||||
}
|
||||
if h.AppliedMigration != nil {
|
||||
attrs = append(attrs, "applied_migration", *h.AppliedMigration)
|
||||
}
|
||||
if h.Error != "" {
|
||||
attrs = append(attrs, "error", h.Error)
|
||||
}
|
||||
|
||||
if status == "ok" {
|
||||
s.log.Debug("health", attrs...)
|
||||
return
|
||||
}
|
||||
s.log.Warn("health", attrs...)
|
||||
}
|
||||
|
||||
type statusRecorder struct {
|
||||
http.ResponseWriter
|
||||
code int
|
||||
}
|
||||
|
||||
func (r *statusRecorder) WriteHeader(code int) {
|
||||
r.code = code
|
||||
r.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
func requestLogger(log *slog.Logger) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
started := time.Now()
|
||||
rec := &statusRecorder{ResponseWriter: w, code: http.StatusOK}
|
||||
next.ServeHTTP(rec, r)
|
||||
log.Info("request",
|
||||
"method", r.Method, "path", r.URL.Path,
|
||||
"status", rec.code, "duration_ms", time.Since(started).Milliseconds())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// jsonErrors converts net/http's own plain-text 404 and 405 replies into the
|
||||
// documented error envelope.
|
||||
//
|
||||
// ServeMux writes those itself, before any handler of ours runs, so a client
|
||||
// that hit a wrong path or method would otherwise get "404 page not found" in
|
||||
// text/plain while every other response is JSON. Only the mux's own replies are
|
||||
// rewritten: anything that set a content type has already answered properly.
|
||||
func jsonErrors(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
iw := &interceptor{ResponseWriter: w}
|
||||
next.ServeHTTP(iw, r)
|
||||
if !iw.rewritten || iw.wrote {
|
||||
return
|
||||
}
|
||||
errCode, message := "not_found", "resource not found"
|
||||
if iw.code == http.StatusMethodNotAllowed {
|
||||
errCode = "method_not_allowed"
|
||||
message = r.Method + " is not supported for this resource"
|
||||
}
|
||||
writeJSON(w, iw.code, errorEnvelope{Error: errorBody{
|
||||
Code: errCode, Message: message, Details: map[string]string{},
|
||||
}})
|
||||
})
|
||||
}
|
||||
|
||||
// interceptor defers the mux's plain-text 404/405 body so it can be replaced.
|
||||
type interceptor struct {
|
||||
http.ResponseWriter
|
||||
code int
|
||||
rewritten bool // this is a mux-generated 404/405 we intend to replace
|
||||
wrote bool // a body already went to the client
|
||||
}
|
||||
|
||||
func (i *interceptor) WriteHeader(code int) {
|
||||
i.code = code
|
||||
if code == http.StatusNotFound || code == http.StatusMethodNotAllowed {
|
||||
if i.Header().Get("Content-Type") != "application/json; charset=utf-8" {
|
||||
i.rewritten = true
|
||||
return // hold the header back; jsonErrors writes its own
|
||||
}
|
||||
}
|
||||
i.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
func (i *interceptor) Write(b []byte) (int, error) {
|
||||
if i.rewritten {
|
||||
return len(b), nil // swallow the mux's plain-text body
|
||||
}
|
||||
i.wrote = true
|
||||
return i.ResponseWriter.Write(b)
|
||||
}
|
||||
|
||||
// recoverer turns a panic into a logged 500 rather than a dropped connection.
|
||||
func recoverer(log *slog.Logger) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() {
|
||||
if v := recover(); v != nil {
|
||||
log.Error("panic", "value", v, "path", r.URL.Path)
|
||||
writeJSON(w, http.StatusInternalServerError, errorEnvelope{Error: errorBody{
|
||||
Code: "internal", Message: "internal error", Details: map[string]string{},
|
||||
}})
|
||||
}
|
||||
}()
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
52
go-api/internal/orgctx/orgctx.go
Normal file
52
go-api/internal/orgctx/orgctx.go
Normal file
@@ -0,0 +1,52 @@
|
||||
// Package orgctx carries the organization a request operates on.
|
||||
//
|
||||
// Phase 2C has no authentication, so there is nothing to derive a tenant from.
|
||||
// Rather than defaulting org_id deep inside the SQL — where it would have to be
|
||||
// unpicked from fourteen repositories once auth arrives — the value is put on
|
||||
// the request context by one middleware and threaded explicitly through the
|
||||
// service and repository boundaries.
|
||||
//
|
||||
// When authentication lands, DevMiddleware is replaced by one that reads the
|
||||
// organization off the authenticated session. Nothing below this package
|
||||
// changes: every caller already takes an org id as a parameter.
|
||||
package orgctx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
type key struct{}
|
||||
|
||||
// DevOrgSlug identifies the single organization every Phase 2C request runs as.
|
||||
// The seeder creates it; nothing else does.
|
||||
//
|
||||
// THIS IS NOT AUTHENTICATION. It is a fixed development identity with no
|
||||
// credential, no session and no verification behind it.
|
||||
const DevOrgSlug = "krow-dev"
|
||||
|
||||
// DevOrgName is that organization's display name.
|
||||
const DevOrgName = "Krow Development"
|
||||
|
||||
// ErrNoOrg means the context reached a scoped operation without an organization,
|
||||
// which is a programming error rather than a client one.
|
||||
var ErrNoOrg = errors.New("no organization in context")
|
||||
|
||||
// With returns a context carrying an organization id.
|
||||
func With(ctx context.Context, orgID string) context.Context {
|
||||
return context.WithValue(ctx, key{}, orgID)
|
||||
}
|
||||
|
||||
// From reads the organization id, reporting whether one was present.
|
||||
func From(ctx context.Context) (string, bool) {
|
||||
v, ok := ctx.Value(key{}).(string)
|
||||
return v, ok && v != ""
|
||||
}
|
||||
|
||||
// MustFrom reads the organization id or returns ErrNoOrg.
|
||||
func MustFrom(ctx context.Context) (string, error) {
|
||||
if v, ok := From(ctx); ok {
|
||||
return v, nil
|
||||
}
|
||||
return "", ErrNoOrg
|
||||
}
|
||||
553
go-api/internal/repo/definitions.go
Normal file
553
go-api/internal/repo/definitions.go
Normal file
@@ -0,0 +1,553 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
)
|
||||
|
||||
// DefinitionListParams holds the query parameters for listing authored definitions.
|
||||
type DefinitionListParams struct {
|
||||
Visibility string
|
||||
Status string
|
||||
DefinitionID string
|
||||
Sort string
|
||||
Desc bool
|
||||
Limit int
|
||||
Offset int
|
||||
}
|
||||
|
||||
// AgentInsertInput holds the fields to insert into agent_definitions.
|
||||
type AgentInsertInput struct {
|
||||
DefinitionID string
|
||||
OrgID string
|
||||
Visibility string
|
||||
OwnerUserID *string
|
||||
CreatedBy *string
|
||||
Markdown string
|
||||
Status string
|
||||
Version int
|
||||
Name string
|
||||
Description string
|
||||
Pages []string
|
||||
}
|
||||
|
||||
// AgentUpdateInput holds the fields to update in agent_definitions.
|
||||
type AgentUpdateInput struct {
|
||||
DefinitionID *string
|
||||
Name *string
|
||||
Description *string
|
||||
Status *string
|
||||
Version *int
|
||||
Pages []string
|
||||
Markdown *string
|
||||
Visibility *string
|
||||
OwnerUserID *string
|
||||
}
|
||||
|
||||
// SkillInsertInput holds the fields to insert into skill_definitions.
|
||||
type SkillInsertInput struct {
|
||||
DefinitionID string
|
||||
OrgID string
|
||||
Visibility string
|
||||
OwnerUserID *string
|
||||
CreatedBy *string
|
||||
Markdown string
|
||||
Status string
|
||||
Name string
|
||||
Description string
|
||||
Pages []string
|
||||
}
|
||||
|
||||
// SkillUpdateInput holds the fields to update in skill_definitions.
|
||||
type SkillUpdateInput struct {
|
||||
DefinitionID *string
|
||||
Name *string
|
||||
Description *string
|
||||
Status *string
|
||||
Pages []string
|
||||
Markdown *string
|
||||
Visibility *string
|
||||
OwnerUserID *string
|
||||
}
|
||||
|
||||
const agentDefinitionColumns = `id::text AS id,
|
||||
definition_id,
|
||||
org_id::text AS org_id,
|
||||
visibility,
|
||||
owner_user_id::text AS owner_user_id,
|
||||
created_by::text AS created_by,
|
||||
markdown,
|
||||
status,
|
||||
version,
|
||||
name,
|
||||
description,
|
||||
pages,
|
||||
to_char(created_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS created_date,
|
||||
to_char(updated_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS updated_date`
|
||||
|
||||
const skillDefinitionColumns = `id::text AS id,
|
||||
definition_id,
|
||||
org_id::text AS org_id,
|
||||
visibility,
|
||||
owner_user_id::text AS owner_user_id,
|
||||
created_by::text AS created_by,
|
||||
markdown,
|
||||
status,
|
||||
name,
|
||||
description,
|
||||
pages,
|
||||
to_char(created_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS created_date,
|
||||
to_char(updated_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"') AS updated_date`
|
||||
|
||||
// DefinitionsRepo handles persistence for agent_definitions and skill_definitions.
|
||||
type DefinitionsRepo struct {
|
||||
db Querier
|
||||
}
|
||||
|
||||
// NewDefinitionsRepo builds a repository over a pool or transaction.
|
||||
func NewDefinitionsRepo(db Querier) *DefinitionsRepo {
|
||||
return &DefinitionsRepo{db: db}
|
||||
}
|
||||
|
||||
/* ── Agents ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// ListAgents lists agent definitions matching the criteria with tenant isolation.
|
||||
func (r *DefinitionsRepo) ListAgents(ctx context.Context, ident authctx.Identity, p DefinitionListParams) ([]domain.Record, int, error) {
|
||||
b := &builder{}
|
||||
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", b.add(ident.OrgID)))
|
||||
|
||||
switch p.Visibility {
|
||||
case "personal":
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'personal' AND owner_user_id = %s::uuid)", b.add(ident.UserID)))
|
||||
case "organization":
|
||||
b.where = append(b.where, "visibility = 'organization'")
|
||||
default:
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = %s::uuid))", b.add(ident.UserID)))
|
||||
}
|
||||
|
||||
if p.Status != "" {
|
||||
b.where = append(b.where, fmt.Sprintf("status = %s::text", b.add(p.Status)))
|
||||
}
|
||||
if p.DefinitionID != "" {
|
||||
b.where = append(b.where, fmt.Sprintf("definition_id = %s::text", b.add(p.DefinitionID)))
|
||||
}
|
||||
|
||||
countSQL := "SELECT count(*)::int FROM agent_definitions" + b.clause()
|
||||
var total int
|
||||
if err := r.db.QueryRow(ctx, countSQL, b.args...).Scan(&total); err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
|
||||
sortCol := "created_date"
|
||||
switch p.Sort {
|
||||
case "created_date", "updated_date", "name", "definition_id", "status", "version":
|
||||
sortCol = p.Sort
|
||||
}
|
||||
dir := "ASC"
|
||||
if p.Desc {
|
||||
dir = "DESC"
|
||||
}
|
||||
orderClause := fmt.Sprintf(" ORDER BY %s %s, id %s", sortCol, dir, dir)
|
||||
|
||||
limit := p.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
offset := p.Offset
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
|
||||
limitClause := fmt.Sprintf(" LIMIT %s OFFSET %s", b.add(limit), b.add(offset))
|
||||
querySQL := "SELECT " + agentDefinitionColumns + " FROM agent_definitions" + b.clause() + orderClause + limitClause
|
||||
|
||||
rows, err := r.db.Query(ctx, querySQL, b.args...)
|
||||
if err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
return records, total, nil
|
||||
}
|
||||
|
||||
// GetAgent retrieves one agent definition by id within the caller's tenant and ownership scope.
|
||||
func (r *DefinitionsRepo) GetAgent(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`SELECT %s FROM agent_definitions
|
||||
WHERE id = $1::uuid
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))
|
||||
LIMIT 1`,
|
||||
agentDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, id, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// GetAgentByDefinitionID retrieves an agent definition by definition_id within the caller's scope,
|
||||
// giving personal definitions precedence over organization definitions when shadowed.
|
||||
func (r *DefinitionsRepo) GetAgentByDefinitionID(ctx context.Context, ident authctx.Identity, defID string) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`SELECT %s FROM agent_definitions
|
||||
WHERE definition_id = $1::text
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))
|
||||
ORDER BY CASE WHEN visibility = 'personal' THEN 1 ELSE 2 END
|
||||
LIMIT 1`,
|
||||
agentDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, defID, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// InsertAgent writes an agent definition and returns the stored row.
|
||||
func (r *DefinitionsRepo) InsertAgent(ctx context.Context, ident authctx.Identity, in AgentInsertInput) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`INSERT INTO agent_definitions (
|
||||
definition_id, org_id, visibility, owner_user_id, created_by,
|
||||
markdown, status, version, name, description, pages
|
||||
) VALUES (
|
||||
$1::text, $2::uuid, $3::text, $4::uuid, $5::uuid,
|
||||
$6::text, $7::text, $8::integer, $9::text, $10::text, $11::text[]
|
||||
) RETURNING %s`, agentDefinitionColumns)
|
||||
|
||||
pages := in.Pages
|
||||
if pages == nil {
|
||||
pages = []string{}
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(ctx, q,
|
||||
in.DefinitionID, in.OrgID, in.Visibility, in.OwnerUserID, in.CreatedBy,
|
||||
in.Markdown, in.Status, in.Version, in.Name, in.Description, pages)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, fmt.Errorf("insert agent_definitions returned no row")
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// UpdateAgent updates fields on an agent definition and returns the full updated record.
|
||||
func (r *DefinitionsRepo) UpdateAgent(ctx context.Context, ident authctx.Identity, id string, in AgentUpdateInput) (domain.Record, error) {
|
||||
b := &builder{}
|
||||
sets := []string{"updated_date = now()"}
|
||||
|
||||
if in.DefinitionID != nil {
|
||||
sets = append(sets, fmt.Sprintf("definition_id = %s::text", b.add(*in.DefinitionID)))
|
||||
}
|
||||
if in.Name != nil {
|
||||
sets = append(sets, fmt.Sprintf("name = %s::text", b.add(*in.Name)))
|
||||
}
|
||||
if in.Description != nil {
|
||||
sets = append(sets, fmt.Sprintf("description = %s::text", b.add(*in.Description)))
|
||||
}
|
||||
if in.Status != nil {
|
||||
sets = append(sets, fmt.Sprintf("status = %s::text", b.add(*in.Status)))
|
||||
}
|
||||
if in.Version != nil {
|
||||
sets = append(sets, fmt.Sprintf("version = %s::integer", b.add(*in.Version)))
|
||||
}
|
||||
if in.Pages != nil {
|
||||
sets = append(sets, fmt.Sprintf("pages = %s::text[]", b.add(in.Pages)))
|
||||
}
|
||||
if in.Markdown != nil {
|
||||
sets = append(sets, fmt.Sprintf("markdown = %s::text", b.add(*in.Markdown)))
|
||||
}
|
||||
if in.Visibility != nil {
|
||||
sets = append(sets, fmt.Sprintf("visibility = %s::text", b.add(*in.Visibility)))
|
||||
sets = append(sets, fmt.Sprintf("owner_user_id = %s::uuid", b.add(in.OwnerUserID)))
|
||||
}
|
||||
|
||||
b.where = append(b.where, fmt.Sprintf("id = %s::uuid", b.add(id)))
|
||||
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", b.add(ident.OrgID)))
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = %s::uuid))", b.add(ident.UserID)))
|
||||
|
||||
q := fmt.Sprintf("UPDATE agent_definitions SET %s%s RETURNING %s",
|
||||
strings.Join(sets, ", "), b.clause(), agentDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// DeleteAgent removes an agent definition matching id and access scope.
|
||||
func (r *DefinitionsRepo) DeleteAgent(ctx context.Context, ident authctx.Identity, id string) (int64, error) {
|
||||
q := `DELETE FROM agent_definitions
|
||||
WHERE id = $1::uuid
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))`
|
||||
|
||||
tag, err := r.db.Exec(ctx, q, id, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return 0, translate(err)
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
|
||||
/* ── Skills ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// ListSkills lists skill definitions matching the criteria with tenant isolation.
|
||||
func (r *DefinitionsRepo) ListSkills(ctx context.Context, ident authctx.Identity, p DefinitionListParams) ([]domain.Record, int, error) {
|
||||
b := &builder{}
|
||||
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", b.add(ident.OrgID)))
|
||||
|
||||
switch p.Visibility {
|
||||
case "personal":
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'personal' AND owner_user_id = %s::uuid)", b.add(ident.UserID)))
|
||||
case "organization":
|
||||
b.where = append(b.where, "visibility = 'organization'")
|
||||
default:
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = %s::uuid))", b.add(ident.UserID)))
|
||||
}
|
||||
|
||||
if p.Status != "" {
|
||||
b.where = append(b.where, fmt.Sprintf("status = %s::text", b.add(p.Status)))
|
||||
}
|
||||
if p.DefinitionID != "" {
|
||||
b.where = append(b.where, fmt.Sprintf("definition_id = %s::text", b.add(p.DefinitionID)))
|
||||
}
|
||||
|
||||
countSQL := "SELECT count(*)::int FROM skill_definitions" + b.clause()
|
||||
var total int
|
||||
if err := r.db.QueryRow(ctx, countSQL, b.args...).Scan(&total); err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
|
||||
sortCol := "created_date"
|
||||
switch p.Sort {
|
||||
case "created_date", "updated_date", "name", "definition_id", "status":
|
||||
sortCol = p.Sort
|
||||
}
|
||||
dir := "ASC"
|
||||
if p.Desc {
|
||||
dir = "DESC"
|
||||
}
|
||||
orderClause := fmt.Sprintf(" ORDER BY %s %s, id %s", sortCol, dir, dir)
|
||||
|
||||
limit := p.Limit
|
||||
if limit <= 0 {
|
||||
limit = 100
|
||||
}
|
||||
offset := p.Offset
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
|
||||
limitClause := fmt.Sprintf(" LIMIT %s OFFSET %s", b.add(limit), b.add(offset))
|
||||
querySQL := "SELECT " + skillDefinitionColumns + " FROM skill_definitions" + b.clause() + orderClause + limitClause
|
||||
|
||||
rows, err := r.db.Query(ctx, querySQL, b.args...)
|
||||
if err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, 0, translate(err)
|
||||
}
|
||||
return records, total, nil
|
||||
}
|
||||
|
||||
// GetSkill retrieves one skill definition by id within the caller's tenant and ownership scope.
|
||||
func (r *DefinitionsRepo) GetSkill(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`SELECT %s FROM skill_definitions
|
||||
WHERE id = $1::uuid
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))
|
||||
LIMIT 1`,
|
||||
skillDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, id, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// GetSkillByDefinitionID retrieves a skill definition by definition_id within the caller's scope,
|
||||
// giving personal definitions precedence over organization definitions when shadowed.
|
||||
func (r *DefinitionsRepo) GetSkillByDefinitionID(ctx context.Context, ident authctx.Identity, defID string) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`SELECT %s FROM skill_definitions
|
||||
WHERE definition_id = $1::text
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))
|
||||
ORDER BY CASE WHEN visibility = 'personal' THEN 1 ELSE 2 END
|
||||
LIMIT 1`,
|
||||
skillDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, defID, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// InsertSkill writes a skill definition and returns the stored row.
|
||||
func (r *DefinitionsRepo) InsertSkill(ctx context.Context, ident authctx.Identity, in SkillInsertInput) (domain.Record, error) {
|
||||
q := fmt.Sprintf(
|
||||
`INSERT INTO skill_definitions (
|
||||
definition_id, org_id, visibility, owner_user_id, created_by,
|
||||
markdown, status, name, description, pages
|
||||
) VALUES (
|
||||
$1::text, $2::uuid, $3::text, $4::uuid, $5::uuid,
|
||||
$6::text, $7::text, $8::text, $9::text, $10::text[]
|
||||
) RETURNING %s`, skillDefinitionColumns)
|
||||
|
||||
pages := in.Pages
|
||||
if pages == nil {
|
||||
pages = []string{}
|
||||
}
|
||||
|
||||
rows, err := r.db.Query(ctx, q,
|
||||
in.DefinitionID, in.OrgID, in.Visibility, in.OwnerUserID, in.CreatedBy,
|
||||
in.Markdown, in.Status, in.Name, in.Description, pages)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, fmt.Errorf("insert skill_definitions returned no row")
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// UpdateSkill updates fields on a skill definition and returns the full updated record.
|
||||
func (r *DefinitionsRepo) UpdateSkill(ctx context.Context, ident authctx.Identity, id string, in SkillUpdateInput) (domain.Record, error) {
|
||||
b := &builder{}
|
||||
sets := []string{"updated_date = now()"}
|
||||
|
||||
if in.DefinitionID != nil {
|
||||
sets = append(sets, fmt.Sprintf("definition_id = %s::text", b.add(*in.DefinitionID)))
|
||||
}
|
||||
if in.Name != nil {
|
||||
sets = append(sets, fmt.Sprintf("name = %s::text", b.add(*in.Name)))
|
||||
}
|
||||
if in.Description != nil {
|
||||
sets = append(sets, fmt.Sprintf("description = %s::text", b.add(*in.Description)))
|
||||
}
|
||||
if in.Status != nil {
|
||||
sets = append(sets, fmt.Sprintf("status = %s::text", b.add(*in.Status)))
|
||||
}
|
||||
if in.Pages != nil {
|
||||
sets = append(sets, fmt.Sprintf("pages = %s::text[]", b.add(in.Pages)))
|
||||
}
|
||||
if in.Markdown != nil {
|
||||
sets = append(sets, fmt.Sprintf("markdown = %s::text", b.add(*in.Markdown)))
|
||||
}
|
||||
if in.Visibility != nil {
|
||||
sets = append(sets, fmt.Sprintf("visibility = %s::text", b.add(*in.Visibility)))
|
||||
sets = append(sets, fmt.Sprintf("owner_user_id = %s::uuid", b.add(in.OwnerUserID)))
|
||||
}
|
||||
|
||||
b.where = append(b.where, fmt.Sprintf("id = %s::uuid", b.add(id)))
|
||||
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", b.add(ident.OrgID)))
|
||||
b.where = append(b.where, fmt.Sprintf("(visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = %s::uuid))", b.add(ident.UserID)))
|
||||
|
||||
q := fmt.Sprintf("UPDATE skill_definitions SET %s%s RETURNING %s",
|
||||
strings.Join(sets, ", "), b.clause(), skillDefinitionColumns)
|
||||
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// DeleteSkill removes a skill definition matching id and access scope.
|
||||
func (r *DefinitionsRepo) DeleteSkill(ctx context.Context, ident authctx.Identity, id string) (int64, error) {
|
||||
q := `DELETE FROM skill_definitions
|
||||
WHERE id = $1::uuid
|
||||
AND org_id = $2::uuid
|
||||
AND (visibility = 'organization' OR (visibility = 'personal' AND owner_user_id = $3::uuid))`
|
||||
|
||||
tag, err := r.db.Exec(ctx, q, id, ident.OrgID, ident.UserID)
|
||||
if err != nil {
|
||||
return 0, translate(err)
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
668
go-api/internal/repo/repo.go
Normal file
668
go-api/internal/repo/repo.go
Normal file
@@ -0,0 +1,668 @@
|
||||
// Package repo is the PostgreSQL access layer.
|
||||
//
|
||||
// Every statement is built from a *domain.Resource: column lists are explicit
|
||||
// and come from the generated descriptors, and every value reaches the database
|
||||
// as a bind parameter cast to its declared type. No identifier is ever taken
|
||||
// from user input — a filter or sort name is resolved to a *domain.Column
|
||||
// first, and an unresolved name is rejected before any SQL is assembled.
|
||||
package repo
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
)
|
||||
|
||||
// Querier is satisfied by both *pgxpool.Pool and pgx.Tx, so every method here
|
||||
// works inside or outside a transaction.
|
||||
type Querier interface {
|
||||
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
|
||||
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
||||
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
|
||||
}
|
||||
|
||||
// Repo reads and writes one resource.
|
||||
type Repo struct {
|
||||
res *domain.Resource
|
||||
db Querier
|
||||
}
|
||||
|
||||
// New builds a repository for a resource over a pool or transaction.
|
||||
func New(res *domain.Resource, db Querier) *Repo { return &Repo{res: res, db: db} }
|
||||
|
||||
// Resource is the descriptor this repository serves.
|
||||
func (r *Repo) Resource() *domain.Resource { return r.res }
|
||||
|
||||
/* ── Projection ─────────────────────────────────────────────────────────── */
|
||||
|
||||
func (r *Repo) selectList() string {
|
||||
parts := make([]string, 0, len(r.res.Columns))
|
||||
for _, c := range r.res.Columns {
|
||||
parts = append(parts, c.SelectExpr())
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
}
|
||||
|
||||
/* ── Predicates ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type builder struct {
|
||||
args []any
|
||||
where []string
|
||||
}
|
||||
|
||||
func (b *builder) add(v any) string {
|
||||
b.args = append(b.args, v)
|
||||
return "$" + strconv.Itoa(len(b.args))
|
||||
}
|
||||
|
||||
// scope applies organization scoping. A NULL org_id on an OrgNullable table
|
||||
// means the shared platform library, which every organization can read.
|
||||
func (b *builder) scope(res *domain.Resource, orgID string) {
|
||||
p := b.add(orgID)
|
||||
if res.OrgNullable {
|
||||
b.where = append(b.where, fmt.Sprintf("(org_id = %s::uuid OR org_id IS NULL)", p))
|
||||
return
|
||||
}
|
||||
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", p))
|
||||
}
|
||||
|
||||
// ownership narrows the rows a caller may touch, beyond their organization.
|
||||
//
|
||||
// This is the second half of authorization and it lives HERE, in the WHERE
|
||||
// clause, rather than in a loop over fetched rows. The difference is not
|
||||
// stylistic: List runs `count(*)` over the same predicate, so a filtered-in-Go
|
||||
// approach would return the right page and the wrong total, and would fetch
|
||||
// rows the caller may not see in order to discard them. A row outside the
|
||||
// predicate is never read.
|
||||
//
|
||||
// Only talent is scoped — Policy.ScopeFor decides that, not this function.
|
||||
// Admin and employer see the whole organization, which is what an operator
|
||||
// console is for.
|
||||
//
|
||||
// A role the API does not recognise matches nothing. That should be
|
||||
// unreachable — the handler answers 403 before any query runs — but a
|
||||
// scope-narrowing function whose failure mode is "see everything" is the wrong
|
||||
// shape to leave lying around.
|
||||
func (b *builder) ownership(res *domain.Resource, ident authctx.Identity) {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
|
||||
scope := res.Policy.ScopeFor(role)
|
||||
switch scope.Kind {
|
||||
case domain.ScopeNone:
|
||||
return
|
||||
|
||||
case domain.ScopeUserID:
|
||||
col, ok := res.Column(scope.Column)
|
||||
if !ok {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
b.where = append(b.where,
|
||||
fmt.Sprintf("%s = %s::%s", col.Name, b.add(ident.UserID), col.PGType))
|
||||
|
||||
case domain.ScopeEmail:
|
||||
col, ok := res.Column(scope.Column)
|
||||
if !ok {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
// citext, so the comparison is case-insensitive — the same equality the
|
||||
// rest of the domain uses for email.
|
||||
b.where = append(b.where,
|
||||
fmt.Sprintf("%s = %s::%s", col.Name, b.add(ident.Email), col.PGType))
|
||||
|
||||
case domain.ScopeActivePostings:
|
||||
col, ok := res.Column(scope.Column)
|
||||
if !ok {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
b.where = append(b.where,
|
||||
fmt.Sprintf("%s = %s::%s", col.Name, b.add("active"), col.PGType))
|
||||
|
||||
case domain.ScopeOwnApplications:
|
||||
// Ownership by reference: the row names an application, the application
|
||||
// names a person. The subquery is scoped to the organization as well as
|
||||
// the email, so it cannot reach across tenants even if an id from
|
||||
// another one were supplied.
|
||||
apps, ok := domain.ResourceByPath["job-applications"]
|
||||
if !ok {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
col, colOK := res.Column(scope.Column)
|
||||
if !colOK {
|
||||
b.where = append(b.where, "false")
|
||||
return
|
||||
}
|
||||
b.where = append(b.where, fmt.Sprintf(
|
||||
"%s IN (SELECT id FROM %s WHERE org_id = %s::uuid AND email = %s::citext)",
|
||||
col.Name, apps.Table, b.add(ident.OrgID), b.add(ident.Email)))
|
||||
|
||||
default:
|
||||
b.where = append(b.where, "false")
|
||||
}
|
||||
}
|
||||
|
||||
// filters renders the contract's two operators and nothing else: equality for a
|
||||
// single value, membership for several. See api-contract.md §6.
|
||||
func (b *builder) filters(fs []domain.Filter) error {
|
||||
for _, f := range fs {
|
||||
if len(f.Values) == 1 {
|
||||
v, err := bindValue(*f.Column, f.Values[0])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
b.where = append(b.where,
|
||||
fmt.Sprintf("%s = %s::%s", f.Column.Name, b.add(v), f.Column.PGType))
|
||||
continue
|
||||
}
|
||||
vals := make([]string, 0, len(f.Values))
|
||||
vals = append(vals, f.Values...)
|
||||
b.where = append(b.where,
|
||||
fmt.Sprintf("%s = ANY(%s::%s[])", f.Column.Name, b.add(vals), arrayElem(f.Column)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func arrayElem(c *domain.Column) string {
|
||||
if c.Kind == domain.KindEnum {
|
||||
return c.PGType
|
||||
}
|
||||
switch c.PGType {
|
||||
case "citext":
|
||||
return "citext"
|
||||
case "uuid":
|
||||
return "uuid"
|
||||
case "int", "bigint":
|
||||
return c.PGType
|
||||
default:
|
||||
return "text"
|
||||
}
|
||||
}
|
||||
|
||||
func (b *builder) clause() string {
|
||||
if len(b.where) == 0 {
|
||||
return ""
|
||||
}
|
||||
return " WHERE " + strings.Join(b.where, " AND ")
|
||||
}
|
||||
|
||||
/* ── Ordering ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// orderBy renders the two ordering rules the contract calls out.
|
||||
//
|
||||
// NULLS LAST in *both* directions, because store.js's comparator returns before
|
||||
// the descending negation is applied — PostgreSQL's default would put nulls
|
||||
// first on DESC. And `, id` as a final tiebreaker, because JavaScript's sort is
|
||||
// stable and PostgreSQL's is not. See api-contract.md §7.1 and §7.3.
|
||||
//
|
||||
// Names are qualified with the table so they bind to the real column rather
|
||||
// than to a same-named output alias from the projection.
|
||||
func (r *Repo) orderBy(p domain.ListParams) string {
|
||||
if p.Sort == "" {
|
||||
return fmt.Sprintf(" ORDER BY %s.id", r.res.Table)
|
||||
}
|
||||
dir := "ASC"
|
||||
if p.Desc {
|
||||
dir = "DESC"
|
||||
}
|
||||
return fmt.Sprintf(" ORDER BY %s.%s %s NULLS LAST, %s.id",
|
||||
r.res.Table, p.Sort, dir, r.res.Table)
|
||||
}
|
||||
|
||||
/* ── Reads ──────────────────────────────────────────────────────────────── */
|
||||
|
||||
// List returns one page plus the total matching count.
|
||||
func (r *Repo) List(ctx context.Context, ident authctx.Identity, p domain.ListParams) (*domain.Page, error) {
|
||||
b := &builder{}
|
||||
b.scope(r.res, ident.OrgID)
|
||||
b.ownership(r.res, ident)
|
||||
if err := b.filters(p.Filters); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
where := b.clause()
|
||||
|
||||
var total int
|
||||
countSQL := "SELECT count(*) FROM " + r.res.Table + where
|
||||
if err := r.db.QueryRow(ctx, countSQL, b.args...).Scan(&total); err != nil {
|
||||
return nil, fmt.Errorf("count %s: %w", r.res.Table, err)
|
||||
}
|
||||
|
||||
limitP := b.add(p.Limit)
|
||||
offsetP := b.add(p.Offset)
|
||||
q := "SELECT " + r.selectList() + " FROM " + r.res.Table + where +
|
||||
r.orderBy(p) + " LIMIT " + limitP + " OFFSET " + offsetP
|
||||
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list %s: %w", r.res.Table, err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &domain.Page{Records: records, Total: total, Limit: p.Limit, Offset: p.Offset}, nil
|
||||
}
|
||||
|
||||
// Get returns one record, or a nil record when nothing matches.
|
||||
func (r *Repo) Get(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
b := &builder{}
|
||||
b.scope(r.res, ident.OrgID)
|
||||
b.ownership(r.res, ident)
|
||||
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
|
||||
|
||||
q := "SELECT " + r.selectList() + " FROM " + r.res.Table + b.clause() + " LIMIT 1"
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get %s: %w", r.res.Table, err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
records, err := collect(rows)
|
||||
if err != nil || len(records) == 0 {
|
||||
return nil, err
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
/* ── Writes ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// derivedValues are the columns this caller does not get to choose.
|
||||
//
|
||||
// Every column here is also ReadOnly in the descriptors, so a value in the
|
||||
// request body has already been dropped by service.validate. This supplies what
|
||||
// goes in instead: a fact about the authenticated session.
|
||||
//
|
||||
// Two shapes share this mechanism and it is worth keeping them apart. `created_by`
|
||||
// and the user_activity identity columns record WHO ACTED — always the session
|
||||
// user, whatever their role. The TalentOnly ones record WHO THE ROW IS ABOUT,
|
||||
// and only a talent caller is necessarily writing about themselves; when an
|
||||
// admin creates a candidate's worker profile, the subject is the candidate, so
|
||||
// nothing is derived and the column is left to the database's default (NULL).
|
||||
func (r *Repo) derivedValues(ident authctx.Identity) map[string]any {
|
||||
if r.res.Policy == nil || len(r.res.Policy.Derived) == 0 {
|
||||
return nil
|
||||
}
|
||||
isTalent := ident.Role == string(domain.RoleTalent)
|
||||
|
||||
out := make(map[string]any, len(r.res.Policy.Derived))
|
||||
for _, d := range r.res.Policy.Derived {
|
||||
if d.TalentOnly && !isTalent {
|
||||
continue
|
||||
}
|
||||
switch d.Source {
|
||||
case domain.DeriveUserID:
|
||||
out[d.Column] = ident.UserID
|
||||
case domain.DeriveEmail:
|
||||
out[d.Column] = ident.Email
|
||||
case domain.DeriveFullName:
|
||||
out[d.Column] = ident.FullName
|
||||
case domain.DeriveAccountType:
|
||||
out[d.Column] = ident.AccountType
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// guardInsert refuses a create whose ownership is expressed by reference rather
|
||||
// than by a column of its own.
|
||||
//
|
||||
// Only ai_interviews needs this: the row names an application, and a talent
|
||||
// caller may only interview for an application of theirs. There is no column on
|
||||
// the interview to derive, so the reference itself has to be checked.
|
||||
//
|
||||
// The refusal is the same 404 an application that does not exist would produce.
|
||||
// Answering "that application is not yours" would confirm it exists.
|
||||
func (r *Repo) guardInsert(ctx context.Context, ident authctx.Identity, in domain.Record) error {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known {
|
||||
return domain.Forbidden()
|
||||
}
|
||||
scope := r.res.Policy.ScopeFor(role)
|
||||
if scope.Kind != domain.ScopeOwnApplications {
|
||||
return nil
|
||||
}
|
||||
|
||||
apps, ok := domain.ResourceByPath["job-applications"]
|
||||
if !ok {
|
||||
return domain.Internal(errors.New("repo: job-applications resource is missing"))
|
||||
}
|
||||
ref, _ := in[scope.Column].(string)
|
||||
if ref == "" {
|
||||
// A required reference that is absent is a validation problem, and the
|
||||
// service's Required check has already reported it.
|
||||
return nil
|
||||
}
|
||||
|
||||
var owns bool
|
||||
q := fmt.Sprintf(
|
||||
`SELECT EXISTS (SELECT 1 FROM %s WHERE org_id = $1::uuid AND id = $2::uuid AND email = $3::citext)`,
|
||||
apps.Table)
|
||||
if err := r.db.QueryRow(ctx, q, ident.OrgID, ref, ident.Email).Scan(&owns); err != nil {
|
||||
return fmt.Errorf("check application ownership: %w", err)
|
||||
}
|
||||
if !owns {
|
||||
return domain.NotFound(apps.Name, ref)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Insert writes a record and returns it as the API represents it.
|
||||
func (r *Repo) Insert(ctx context.Context, ident authctx.Identity, in domain.Record) (domain.Record, error) {
|
||||
if err := r.guardInsert(ctx, ident, in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cols := []string{"org_id"}
|
||||
b := &builder{}
|
||||
// The organization is the session's, never the body's. org_id is ReadOnly
|
||||
// on every resource, so this is the only way a value reaches the column.
|
||||
vals := []string{b.add(ident.OrgID) + "::uuid"}
|
||||
|
||||
derived := r.derivedValues(ident)
|
||||
for _, c := range r.res.Columns {
|
||||
// A derived column is written whether or not it is ReadOnly, and it
|
||||
// overrides anything the caller sent — which is what closes the
|
||||
// attribution holes: a talent caller cannot file an application, a
|
||||
// profile or an audit entry under somebody else's name.
|
||||
if v, ok := derived[c.Name]; ok {
|
||||
cols = append(cols, c.Name)
|
||||
vals = append(vals, b.add(v)+"::"+c.PGType)
|
||||
continue
|
||||
}
|
||||
if c.ReadOnly {
|
||||
continue
|
||||
}
|
||||
v, present := in[c.Name]
|
||||
if !present {
|
||||
continue // let the column default apply
|
||||
}
|
||||
bound, err := bindValue(c, v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cols = append(cols, c.Name)
|
||||
vals = append(vals, b.add(bound)+"::"+c.PGType)
|
||||
}
|
||||
|
||||
q := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s) RETURNING %s",
|
||||
r.res.Table, strings.Join(cols, ", "), strings.Join(vals, ", "), r.selectList())
|
||||
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
// pgx does not execute until the rows are read, so a constraint violation
|
||||
// arrives here rather than from Query above. Both paths must translate.
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, fmt.Errorf("insert %s returned no row", r.res.Table)
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// Update applies a shallow merge: only the supplied columns are written.
|
||||
//
|
||||
// A key that is absent is left alone; a key present with null sets NULL. There
|
||||
// is no deep merge — store.js spreads one level, and useSubmitChallenge relies
|
||||
// on whole arrays being replaced rather than appended to. See api-contract.md §3.2.
|
||||
func (r *Repo) Update(ctx context.Context, ident authctx.Identity, id string, patch domain.Record) (domain.Record, error) {
|
||||
b := &builder{}
|
||||
sets := []string{}
|
||||
|
||||
for _, c := range r.res.Columns {
|
||||
if c.ReadOnly {
|
||||
continue
|
||||
}
|
||||
v, present := patch[c.Name]
|
||||
if !present {
|
||||
continue
|
||||
}
|
||||
bound, err := bindValue(c, v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sets = append(sets, fmt.Sprintf("%s = %s::%s", c.Name, b.add(bound), c.PGType))
|
||||
}
|
||||
|
||||
// updated_date is always server-owned. user_activity is append-only and has
|
||||
// no such column, so this is conditional rather than assumed.
|
||||
if _, ok := r.res.Column("updated_date"); ok {
|
||||
sets = append(sets, "updated_date = now()")
|
||||
}
|
||||
if len(sets) == 0 {
|
||||
return r.Get(ctx, ident, id)
|
||||
}
|
||||
|
||||
b.scope(r.res, ident.OrgID)
|
||||
b.ownership(r.res, ident)
|
||||
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
|
||||
|
||||
q := fmt.Sprintf("UPDATE %s SET %s%s RETURNING %s",
|
||||
r.res.Table, strings.Join(sets, ", "), b.clause(), r.selectList())
|
||||
|
||||
rows, err := r.db.Query(ctx, q, b.args...)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
records, err := collect(rows)
|
||||
if err != nil {
|
||||
return nil, translate(err)
|
||||
}
|
||||
if len(records) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return records[0], nil
|
||||
}
|
||||
|
||||
// Delete removes a record and reports how many rows went.
|
||||
//
|
||||
// A miss is not an error: store.js filters its array and returns { id }
|
||||
// whether or not anything matched. See api-contract.md §12.7.
|
||||
func (r *Repo) Delete(ctx context.Context, ident authctx.Identity, id string) (int64, error) {
|
||||
b := &builder{}
|
||||
b.scope(r.res, ident.OrgID)
|
||||
b.ownership(r.res, ident)
|
||||
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
|
||||
|
||||
tag, err := r.db.Exec(ctx, "DELETE FROM "+r.res.Table+b.clause(), b.args...)
|
||||
if err != nil {
|
||||
return 0, translate(err)
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
|
||||
/* ── Scanning ───────────────────────────────────────────────────────────── */
|
||||
|
||||
func collect(rows pgx.Rows) ([]domain.Record, error) {
|
||||
fields := rows.FieldDescriptions()
|
||||
out := make([]domain.Record, 0, 16)
|
||||
for rows.Next() {
|
||||
vals, err := rows.Values()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rec := make(domain.Record, len(fields))
|
||||
for i, f := range fields {
|
||||
rec[string(f.Name)] = normalise(vals[i])
|
||||
}
|
||||
out = append(out, rec)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// normalise flattens the few pgx representations that would not JSON-encode the
|
||||
// way the frontend expects. Most values arrive ready to use because the
|
||||
// projection casts them (see Column.SelectExpr).
|
||||
func normalise(v any) any {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
return nil
|
||||
case [16]byte: // a uuid that slipped through without a ::text cast
|
||||
return fmt.Sprintf("%x-%x-%x-%x-%x", t[0:4], t[4:6], t[6:8], t[8:10], t[10:16])
|
||||
case []any:
|
||||
out := make([]any, len(t))
|
||||
for i, e := range t {
|
||||
out[i] = normalise(e)
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Binding ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// bindValue converts a decoded-JSON value into something pgx can send for a
|
||||
// column of this type. Everything is explicitly cast in the SQL, so the job
|
||||
// here is only to pick a Go representation PostgreSQL will accept.
|
||||
func bindValue(c domain.Column, v any) (any, error) {
|
||||
if v == nil {
|
||||
return nil, nil
|
||||
}
|
||||
switch c.Kind {
|
||||
case domain.KindTextArray:
|
||||
switch t := v.(type) {
|
||||
case []any:
|
||||
out := make([]string, 0, len(t))
|
||||
for _, e := range t {
|
||||
s, ok := e.(string)
|
||||
if !ok {
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s must be an array of strings", c.Name),
|
||||
map[string]string{c.Name: "expected string elements"})
|
||||
}
|
||||
out = append(out, s)
|
||||
}
|
||||
return out, nil
|
||||
case []string:
|
||||
return t, nil
|
||||
default:
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s must be an array", c.Name),
|
||||
map[string]string{c.Name: "expected an array"})
|
||||
}
|
||||
|
||||
case domain.KindJSON:
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, domain.Validation(fmt.Sprintf("%s is not encodable as JSON", c.Name), nil)
|
||||
}
|
||||
return raw, nil
|
||||
|
||||
case domain.KindInt:
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
if t != float64(int64(t)) {
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s must be a whole number", c.Name),
|
||||
map[string]string{c.Name: "expected an integer"})
|
||||
}
|
||||
return int64(t), nil
|
||||
case string:
|
||||
n, err := strconv.ParseInt(t, 10, 64)
|
||||
if err != nil {
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s must be a whole number", c.Name),
|
||||
map[string]string{c.Name: "expected an integer"})
|
||||
}
|
||||
return n, nil
|
||||
case int64:
|
||||
return t, nil
|
||||
case int:
|
||||
return int64(t), nil
|
||||
}
|
||||
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
|
||||
|
||||
case domain.KindFloat:
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
return t, nil
|
||||
case string:
|
||||
f, err := strconv.ParseFloat(t, 64)
|
||||
if err != nil {
|
||||
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
|
||||
|
||||
case domain.KindBool:
|
||||
switch t := v.(type) {
|
||||
case bool:
|
||||
return t, nil
|
||||
case string:
|
||||
b, err := strconv.ParseBool(t)
|
||||
if err != nil {
|
||||
return nil, domain.Validation(fmt.Sprintf("%s must be a boolean", c.Name), nil)
|
||||
}
|
||||
return b, nil
|
||||
}
|
||||
return nil, domain.Validation(fmt.Sprintf("%s must be a boolean", c.Name), nil)
|
||||
|
||||
default: // strings, enums, uuids, dates, timestamps
|
||||
s, ok := v.(string)
|
||||
if !ok {
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s must be a string", c.Name),
|
||||
map[string]string{c.Name: "expected a string"})
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Error translation ──────────────────────────────────────────────────── */
|
||||
|
||||
// translate maps PostgreSQL's SQLSTATE codes onto the contract's error codes,
|
||||
// so a constraint the database enforces surfaces as the documented API error
|
||||
// rather than as a 500.
|
||||
func translate(err error) error {
|
||||
var pg *pgconn.PgError
|
||||
if !errors.As(err, &pg) {
|
||||
return err
|
||||
}
|
||||
switch pg.Code {
|
||||
case "23505": // unique_violation
|
||||
return domain.Conflict(pg.Detail)
|
||||
case "23503": // foreign_key_violation
|
||||
return domain.Validation("referenced record does not exist",
|
||||
map[string]string{constraintField(pg): "no such record"})
|
||||
case "23514": // check_violation
|
||||
return domain.Validation("value violates constraint "+pg.ConstraintName,
|
||||
map[string]string{constraintField(pg): "constraint " + pg.ConstraintName})
|
||||
case "23502": // not_null_violation
|
||||
return domain.Validation(pg.ColumnName+" must not be null",
|
||||
map[string]string{pg.ColumnName: "required"})
|
||||
case "22P02", "22007", "22008": // invalid text representation / datetime
|
||||
return domain.Validation("value is not valid for its column type", nil)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func constraintField(pg *pgconn.PgError) string {
|
||||
if pg.ColumnName != "" {
|
||||
return pg.ColumnName
|
||||
}
|
||||
return pg.ConstraintName
|
||||
}
|
||||
129
go-api/internal/runtime/executor.go
Normal file
129
go-api/internal/runtime/executor.go
Normal file
@@ -0,0 +1,129 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
// AgentExecutor is the boundary interface for executing an authored agent.
|
||||
type AgentExecutor interface {
|
||||
ExecuteAgent(ctx context.Context, agent *Agent, input ExecutionInput) (*ExecutionResult, error)
|
||||
}
|
||||
|
||||
// SkillExecutor is the boundary interface for executing an authored skill.
|
||||
type SkillExecutor interface {
|
||||
ExecuteSkill(ctx context.Context, skill *Skill, input ExecutionInput) (*ExecutionResult, error)
|
||||
}
|
||||
|
||||
// UnavailableExecutor is the Phase 4F default stub executor that explicitly refuses execution
|
||||
// until real AI / Owliver / LangGraph executors are installed in Phase 5.
|
||||
type UnavailableExecutor struct{}
|
||||
|
||||
// ExecuteAgent implements AgentExecutor by returning ErrExecutorUnavailable.
|
||||
func (u *UnavailableExecutor) ExecuteAgent(_ context.Context, agent *Agent, _ ExecutionInput) (*ExecutionResult, error) {
|
||||
skillIDs := make([]string, len(agent.ResolvedSkills))
|
||||
for i, s := range agent.ResolvedSkills {
|
||||
skillIDs[i] = s.ID
|
||||
}
|
||||
|
||||
return &ExecutionResult{
|
||||
Success: false,
|
||||
AgentID: agent.ID,
|
||||
AgentVersion: agent.Version,
|
||||
ResolvedSkills: skillIDs,
|
||||
Error: ErrExecutorUnavailable,
|
||||
}, ErrExecutorUnavailable
|
||||
}
|
||||
|
||||
// ExecuteSkill implements SkillExecutor by returning ErrExecutorUnavailable.
|
||||
func (u *UnavailableExecutor) ExecuteSkill(_ context.Context, skill *Skill, _ ExecutionInput) (*ExecutionResult, error) {
|
||||
return &ExecutionResult{
|
||||
Success: false,
|
||||
Error: ErrExecutorUnavailable,
|
||||
}, ErrExecutorUnavailable
|
||||
}
|
||||
|
||||
// Engine coordinates runtime loading, eligibility checks, dependency resolution and execution.
|
||||
type Engine struct {
|
||||
Loader *Loader
|
||||
AgentExec AgentExecutor
|
||||
SkillExec SkillExecutor
|
||||
}
|
||||
|
||||
// Option configures the runtime engine.
|
||||
type Option func(*Engine)
|
||||
|
||||
// WithAgentExecutor overrides the agent executor implementation.
|
||||
func WithAgentExecutor(exec AgentExecutor) Option {
|
||||
return func(e *Engine) {
|
||||
if exec != nil {
|
||||
e.AgentExec = exec
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithSkillExecutor overrides the skill executor implementation.
|
||||
func WithSkillExecutor(exec SkillExecutor) Option {
|
||||
return func(e *Engine) {
|
||||
if exec != nil {
|
||||
e.SkillExec = exec
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NewEngine builds a runtime engine over a database querier.
|
||||
func NewEngine(db repo.Querier, opts ...Option) *Engine {
|
||||
e := &Engine{
|
||||
Loader: NewLoader(db),
|
||||
AgentExec: &UnavailableExecutor{},
|
||||
SkillExec: &UnavailableExecutor{},
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(e)
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
// RunAgent loads an executable agent with dependencies and dispatches to the executor boundary.
|
||||
func (e *Engine) RunAgent(ctx context.Context, ident authctx.Identity, idOrDefID string, input ExecutionInput) (*ExecutionResult, error) {
|
||||
agent, err := e.Loader.LoadExecutableAgent(ctx, ident, idOrDefID)
|
||||
if err != nil {
|
||||
return &ExecutionResult{
|
||||
Success: false,
|
||||
Error: err,
|
||||
}, err
|
||||
}
|
||||
|
||||
res, err := e.AgentExec.ExecuteAgent(ctx, agent, input)
|
||||
if res == nil {
|
||||
res = &ExecutionResult{
|
||||
Success: err == nil,
|
||||
AgentID: agent.ID,
|
||||
AgentVersion: agent.Version,
|
||||
Error: err,
|
||||
}
|
||||
}
|
||||
return res, err
|
||||
}
|
||||
|
||||
// RunSkill loads an executable skill and dispatches to the executor boundary.
|
||||
func (e *Engine) RunSkill(ctx context.Context, ident authctx.Identity, idOrDefID string, input ExecutionInput) (*ExecutionResult, error) {
|
||||
skill, err := e.Loader.LoadExecutableSkill(ctx, ident, idOrDefID)
|
||||
if err != nil {
|
||||
return &ExecutionResult{
|
||||
Success: false,
|
||||
Error: err,
|
||||
}, err
|
||||
}
|
||||
|
||||
res, err := e.SkillExec.ExecuteSkill(ctx, skill, input)
|
||||
if res == nil {
|
||||
res = &ExecutionResult{
|
||||
Success: err == nil,
|
||||
Error: err,
|
||||
}
|
||||
}
|
||||
return res, err
|
||||
}
|
||||
237
go-api/internal/runtime/loader.go
Normal file
237
go-api/internal/runtime/loader.go
Normal file
@@ -0,0 +1,237 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/definition"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
var uuidPattern = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`)
|
||||
|
||||
func isUUID(s string) bool {
|
||||
return uuidPattern.MatchString(s)
|
||||
}
|
||||
|
||||
// Loader loads and validates authored definitions into runtime representations with tenant isolation.
|
||||
type Loader struct {
|
||||
repo *repo.DefinitionsRepo
|
||||
}
|
||||
|
||||
// NewLoader builds a runtime definition loader over a storage repository.
|
||||
func NewLoader(db repo.Querier) *Loader {
|
||||
return &Loader{repo: repo.NewDefinitionsRepo(db)}
|
||||
}
|
||||
|
||||
// LoadAgent loads an agent definition by id or definition_id, parsing it into a runtime representation.
|
||||
func (l *Loader) LoadAgent(ctx context.Context, ident authctx.Identity, idOrDefID string) (*Agent, error) {
|
||||
var (
|
||||
rec domain.Record
|
||||
err error
|
||||
)
|
||||
|
||||
if isUUID(idOrDefID) {
|
||||
rec, err = l.repo.GetAgent(ctx, ident, idOrDefID)
|
||||
} else {
|
||||
rec, err = l.repo.GetAgentByDefinitionID(ctx, ident, idOrDefID)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, fmt.Errorf("%w: agent %q", ErrNotFound, idOrDefID)
|
||||
}
|
||||
|
||||
rawMD, ok := rec["markdown"].(string)
|
||||
if !ok || rawMD == "" {
|
||||
return nil, fmt.Errorf("%w: missing markdown payload for agent %q", ErrInvalidDefinition, idOrDefID)
|
||||
}
|
||||
|
||||
if err := definition.ValidateAgent(rawMD); err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidDefinition, err)
|
||||
}
|
||||
|
||||
parsed, err := definition.ParseAgent(rawMD, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidDefinition, err)
|
||||
}
|
||||
|
||||
agent := &Agent{
|
||||
ID: parsed.ID,
|
||||
DatabaseID: rec["id"].(string),
|
||||
Name: parsed.Name,
|
||||
Description: parsed.Description,
|
||||
Status: parsed.Status,
|
||||
Version: parsed.Version,
|
||||
Visibility: rec["visibility"].(string),
|
||||
Pages: parsed.Pages,
|
||||
Icon: parsed.Icon,
|
||||
Reasoning: parsed.Reasoning,
|
||||
Trigger: parsed.Trigger,
|
||||
WebSearch: parsed.WebSearch,
|
||||
Instructions: parsed.Instructions,
|
||||
Skills: parsed.Skills,
|
||||
Subagents: parsed.Subagents,
|
||||
RawMarkdown: rawMD,
|
||||
}
|
||||
|
||||
if rec["owner_user_id"] != nil {
|
||||
if uid, ok := rec["owner_user_id"].(string); ok && uid != "" {
|
||||
agent.OwnerUserID = &uid
|
||||
}
|
||||
}
|
||||
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// LoadSkill loads a skill definition by id or definition_id, parsing it into a runtime representation.
|
||||
func (l *Loader) LoadSkill(ctx context.Context, ident authctx.Identity, idOrDefID string) (*Skill, error) {
|
||||
var (
|
||||
rec domain.Record
|
||||
err error
|
||||
)
|
||||
|
||||
if isUUID(idOrDefID) {
|
||||
rec, err = l.repo.GetSkill(ctx, ident, idOrDefID)
|
||||
} else {
|
||||
rec, err = l.repo.GetSkillByDefinitionID(ctx, ident, idOrDefID)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, fmt.Errorf("%w: skill %q", ErrNotFound, idOrDefID)
|
||||
}
|
||||
|
||||
rawMD, ok := rec["markdown"].(string)
|
||||
if !ok || rawMD == "" {
|
||||
return nil, fmt.Errorf("%w: missing markdown payload for skill %q", ErrInvalidDefinition, idOrDefID)
|
||||
}
|
||||
|
||||
if err := definition.ValidateSkill(rawMD); err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidDefinition, err)
|
||||
}
|
||||
|
||||
parsed, err := definition.ParseSkill(rawMD, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrInvalidDefinition, err)
|
||||
}
|
||||
|
||||
skill := &Skill{
|
||||
ID: parsed.ID,
|
||||
DatabaseID: rec["id"].(string),
|
||||
Name: parsed.Name,
|
||||
Description: parsed.Description,
|
||||
Status: parsed.Status,
|
||||
Visibility: rec["visibility"].(string),
|
||||
Pages: parsed.Pages,
|
||||
Kind: parsed.Kind,
|
||||
Category: parsed.Category,
|
||||
Actions: parsed.Actions,
|
||||
Triggers: parsed.Triggers,
|
||||
Prompt: parsed.Prompt,
|
||||
SkillID: parsed.SkillID,
|
||||
Body: parsed.Body,
|
||||
RawMarkdown: rawMD,
|
||||
}
|
||||
|
||||
if rec["owner_user_id"] != nil {
|
||||
if uid, ok := rec["owner_user_id"].(string); ok && uid != "" {
|
||||
skill.OwnerUserID = &uid
|
||||
}
|
||||
}
|
||||
|
||||
return skill, nil
|
||||
}
|
||||
|
||||
// ResolveAgentDependencies resolves all skill dependencies referenced by the agent within caller scope.
|
||||
func (l *Loader) ResolveAgentDependencies(ctx context.Context, ident authctx.Identity, agent *Agent) error {
|
||||
if len(agent.Skills) == 0 {
|
||||
agent.ResolvedSkills = []*Skill{}
|
||||
return nil
|
||||
}
|
||||
|
||||
visited := make(map[string]*Skill)
|
||||
inProgress := make(map[string]bool)
|
||||
resolved := make([]*Skill, 0, len(agent.Skills))
|
||||
|
||||
for _, skillID := range agent.Skills {
|
||||
if _, ok := visited[skillID]; ok {
|
||||
// Deterministic deduplication
|
||||
continue
|
||||
}
|
||||
if inProgress[skillID] {
|
||||
return fmt.Errorf("%w: skill %q", ErrCircularDependency, skillID)
|
||||
}
|
||||
inProgress[skillID] = true
|
||||
|
||||
skill, err := l.LoadSkill(ctx, ident, skillID)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
return fmt.Errorf("%w: skill %q", ErrDependencyMissing, skillID)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if skill.Status != "active" {
|
||||
return fmt.Errorf("%w: skill %q has status %q", ErrDependencyInactive, skillID, skill.Status)
|
||||
}
|
||||
|
||||
inProgress[skillID] = false
|
||||
visited[skillID] = skill
|
||||
resolved = append(resolved, skill)
|
||||
}
|
||||
|
||||
agent.ResolvedSkills = resolved
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadExecutableAgent loads an agent, verifies its published status, and resolves all active dependencies.
|
||||
func (l *Loader) LoadExecutableAgent(ctx context.Context, ident authctx.Identity, idOrDefID string) (*Agent, error) {
|
||||
agent, err := l.LoadAgent(ctx, ident, idOrDefID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch agent.Status {
|
||||
case "published":
|
||||
// Eligible
|
||||
case "draft":
|
||||
return nil, fmt.Errorf("%w: agent %q is in draft status", ErrDraftAgent, agent.ID)
|
||||
case "archived":
|
||||
return nil, fmt.Errorf("%w: agent %q is archived", ErrArchivedAgent, agent.ID)
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: agent %q has unsupported status %q", ErrNotExecutable, agent.ID, agent.Status)
|
||||
}
|
||||
|
||||
if err := l.ResolveAgentDependencies(ctx, ident, agent); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
// LoadExecutableSkill loads a skill and verifies its active status.
|
||||
func (l *Loader) LoadExecutableSkill(ctx context.Context, ident authctx.Identity, idOrDefID string) (*Skill, error) {
|
||||
skill, err := l.LoadSkill(ctx, ident, idOrDefID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch skill.Status {
|
||||
case "active":
|
||||
// Eligible
|
||||
case "inactive":
|
||||
return nil, fmt.Errorf("%w: skill %q is inactive", ErrInactiveSkill, skill.ID)
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: skill %q has unsupported status %q", ErrNotExecutable, skill.ID, skill.Status)
|
||||
}
|
||||
|
||||
return skill, nil
|
||||
}
|
||||
890
go-api/internal/runtime/runtime_test.go
Normal file
890
go-api/internal/runtime/runtime_test.go
Normal file
@@ -0,0 +1,890 @@
|
||||
package runtime_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
type fixture struct {
|
||||
h *testutil.Harness
|
||||
loader *runtime.Loader
|
||||
engine *runtime.Engine
|
||||
defRepo *repo.DefinitionsRepo
|
||||
org1 string
|
||||
org2 string
|
||||
userA authctx.Identity
|
||||
userB authctx.Identity
|
||||
userOther authctx.Identity
|
||||
}
|
||||
|
||||
func newFixture(t *testing.T) *fixture {
|
||||
t.Helper()
|
||||
h := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var org2 string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Second Org', 'second-org') RETURNING id::text`).
|
||||
Scan(&org2); err != nil {
|
||||
t.Fatalf("create second org: %v", err)
|
||||
}
|
||||
|
||||
var userAID, userBID, userOtherID string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, full_name, email, role, status) VALUES ($1::uuid, 'User A', 'user-a@example.test', 'admin', 'active') RETURNING id::text`,
|
||||
h.OrgID).Scan(&userAID); err != nil {
|
||||
t.Fatalf("create userA: %v", err)
|
||||
}
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, full_name, email, role, status) VALUES ($1::uuid, 'User B', 'user-b@example.test', 'talent', 'active') RETURNING id::text`,
|
||||
h.OrgID).Scan(&userBID); err != nil {
|
||||
t.Fatalf("create userB: %v", err)
|
||||
}
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, full_name, email, role, status) VALUES ($1::uuid, 'User Other', 'user-other@example.test', 'admin', 'active') RETURNING id::text`,
|
||||
org2).Scan(&userOtherID); err != nil {
|
||||
t.Fatalf("create userOther: %v", err)
|
||||
}
|
||||
|
||||
f := &fixture{
|
||||
h: h,
|
||||
loader: runtime.NewLoader(h.Pool),
|
||||
engine: runtime.NewEngine(h.Pool),
|
||||
defRepo: repo.NewDefinitionsRepo(h.Pool),
|
||||
org1: h.OrgID,
|
||||
org2: org2,
|
||||
userA: authctx.Identity{
|
||||
UserID: userAID,
|
||||
OrgID: h.OrgID,
|
||||
Role: "admin",
|
||||
Email: "user-a@example.test",
|
||||
},
|
||||
userB: authctx.Identity{
|
||||
UserID: userBID,
|
||||
OrgID: h.OrgID,
|
||||
Role: "talent",
|
||||
Email: "user-b@example.test",
|
||||
},
|
||||
userOther: authctx.Identity{
|
||||
UserID: userOtherID,
|
||||
OrgID: org2,
|
||||
Role: "admin",
|
||||
Email: "user-other@example.test",
|
||||
},
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
/* ── 1. Loader Tests ──────────────────────────────────────────────────────── */
|
||||
|
||||
func TestRuntimeLoader_Agent(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 1. Published org agent
|
||||
agentMD := `---
|
||||
id: talent-scout
|
||||
name: Talent Scout
|
||||
description: Discovers matching candidates
|
||||
status: published
|
||||
version: 2
|
||||
pages:
|
||||
- candidates
|
||||
icon: sparkles
|
||||
reasoning: deep
|
||||
trigger: manual
|
||||
webSearch: true
|
||||
skills:
|
||||
- resume-evaluator
|
||||
subagents:
|
||||
- profile-enricher
|
||||
---
|
||||
|
||||
## Instructions
|
||||
Review candidate profiles with diligence.
|
||||
`
|
||||
rec, err := f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "talent-scout",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: agentMD,
|
||||
Status: "published",
|
||||
Version: 2,
|
||||
Name: "Talent Scout",
|
||||
Description: "Discovers matching candidates",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("insert agent: %v", err)
|
||||
}
|
||||
dbUUID := rec["id"].(string)
|
||||
|
||||
// Load by definition_id
|
||||
agentByDefID, err := f.loader.LoadAgent(ctx, f.userB, "talent-scout")
|
||||
if err != nil {
|
||||
t.Fatalf("load agent by definition_id: %v", err)
|
||||
}
|
||||
if agentByDefID.ID != "talent-scout" || agentByDefID.DatabaseID != dbUUID {
|
||||
t.Errorf("agent ID mismatch: %+v", agentByDefID)
|
||||
}
|
||||
if agentByDefID.Name != "Talent Scout" || agentByDefID.Version != 2 || agentByDefID.Status != "published" {
|
||||
t.Errorf("agent fields mismatch: %+v", agentByDefID)
|
||||
}
|
||||
if agentByDefID.Icon != "sparkles" || agentByDefID.Reasoning != "deep" || !agentByDefID.WebSearch {
|
||||
t.Errorf("agent config mismatch: %+v", agentByDefID)
|
||||
}
|
||||
if len(agentByDefID.Skills) != 1 || agentByDefID.Skills[0] != "resume-evaluator" {
|
||||
t.Errorf("agent skills mismatch: %v", agentByDefID.Skills)
|
||||
}
|
||||
if agentByDefID.Instructions == "" || agentByDefID.RawMarkdown != agentMD {
|
||||
t.Errorf("agent markdown/instructions mismatch")
|
||||
}
|
||||
|
||||
// Load by UUID
|
||||
agentByUUID, err := f.loader.LoadAgent(ctx, f.userA, dbUUID)
|
||||
if err != nil {
|
||||
t.Fatalf("load agent by UUID: %v", err)
|
||||
}
|
||||
if agentByUUID.ID != "talent-scout" {
|
||||
t.Errorf("agent by UUID ID mismatch: %+v", agentByUUID)
|
||||
}
|
||||
|
||||
// Personal agent for User B
|
||||
persMD := `---
|
||||
id: personal-agent
|
||||
name: Personal Agent
|
||||
status: draft
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
## Instructions
|
||||
Personal instructions.
|
||||
`
|
||||
persRec, err := f.defRepo.InsertAgent(ctx, f.userB, repo.AgentInsertInput{
|
||||
DefinitionID: "personal-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "personal",
|
||||
OwnerUserID: &f.userB.UserID,
|
||||
CreatedBy: &f.userB.UserID,
|
||||
Markdown: persMD,
|
||||
Status: "draft",
|
||||
Version: 1,
|
||||
Name: "Personal Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("insert personal agent: %v", err)
|
||||
}
|
||||
persUUID := persRec["id"].(string)
|
||||
|
||||
// Owner (User B) can load personal agent
|
||||
persAgent, err := f.loader.LoadAgent(ctx, f.userB, "personal-agent")
|
||||
if err != nil {
|
||||
t.Fatalf("load own personal agent: %v", err)
|
||||
}
|
||||
if persAgent.DatabaseID != persUUID {
|
||||
t.Errorf("personal agent uuid mismatch: %s != %s", persAgent.DatabaseID, persUUID)
|
||||
}
|
||||
|
||||
// Other user in same org (User A) CANNOT load User B's personal agent -> ErrNotFound
|
||||
_, err = f.loader.LoadAgent(ctx, f.userA, "personal-agent")
|
||||
if !errors.Is(err, runtime.ErrNotFound) {
|
||||
t.Errorf("userA loading userB personal agent: got error %v, want ErrNotFound", err)
|
||||
}
|
||||
|
||||
// Outsider CANNOT load org agent from org 1 -> ErrNotFound
|
||||
_, err = f.loader.LoadAgent(ctx, f.userOther, "talent-scout")
|
||||
if !errors.Is(err, runtime.ErrNotFound) {
|
||||
t.Errorf("outsider loading org1 agent: got error %v, want ErrNotFound", err)
|
||||
}
|
||||
|
||||
// Missing agent -> ErrNotFound
|
||||
_, err = f.loader.LoadAgent(ctx, f.userA, "nonexistent-agent")
|
||||
if !errors.Is(err, runtime.ErrNotFound) {
|
||||
t.Errorf("load nonexistent agent: got %v, want ErrNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeLoader_Skill(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
skillMD := `---
|
||||
id: resume-evaluator
|
||||
name: Resume Evaluator
|
||||
description: Evaluates candidate resume text
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
category: screening
|
||||
actions:
|
||||
- score
|
||||
- summarize
|
||||
triggers:
|
||||
- resume
|
||||
- cv
|
||||
prompt: Evaluate the candidate resume thoroughly.
|
||||
---
|
||||
|
||||
# Resume Evaluator Body
|
||||
Detailed skill instructions.
|
||||
`
|
||||
rec, err := f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "resume-evaluator",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: skillMD,
|
||||
Status: "active",
|
||||
Name: "Resume Evaluator",
|
||||
Description: "Evaluates candidate resume text",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("insert skill: %v", err)
|
||||
}
|
||||
skillUUID := rec["id"].(string)
|
||||
|
||||
// Load by definition_id
|
||||
skillByDefID, err := f.loader.LoadSkill(ctx, f.userB, "resume-evaluator")
|
||||
if err != nil {
|
||||
t.Fatalf("load skill by definition_id: %v", err)
|
||||
}
|
||||
if skillByDefID.ID != "resume-evaluator" || skillByDefID.DatabaseID != skillUUID {
|
||||
t.Errorf("skill ID mismatch: %+v", skillByDefID)
|
||||
}
|
||||
if skillByDefID.Name != "Resume Evaluator" || skillByDefID.Status != "active" {
|
||||
t.Errorf("skill fields mismatch: %+v", skillByDefID)
|
||||
}
|
||||
if skillByDefID.Category != "screening" || len(skillByDefID.Actions) != 2 || len(skillByDefID.Triggers) != 2 {
|
||||
t.Errorf("skill actions/triggers mismatch: %+v", skillByDefID)
|
||||
}
|
||||
if skillByDefID.Prompt == nil || *skillByDefID.Prompt != "Evaluate the candidate resume thoroughly." {
|
||||
t.Errorf("skill prompt mismatch: %v", skillByDefID.Prompt)
|
||||
}
|
||||
if skillByDefID.RawMarkdown != skillMD {
|
||||
t.Errorf("skill raw markdown not verbatim")
|
||||
}
|
||||
|
||||
// Load by UUID
|
||||
skillByUUID, err := f.loader.LoadSkill(ctx, f.userA, skillUUID)
|
||||
if err != nil {
|
||||
t.Fatalf("load skill by UUID: %v", err)
|
||||
}
|
||||
if skillByUUID.ID != "resume-evaluator" {
|
||||
t.Errorf("skill by UUID mismatch: %+v", skillByUUID)
|
||||
}
|
||||
|
||||
// Missing skill -> ErrNotFound
|
||||
_, err = f.loader.LoadSkill(ctx, f.userA, "nonexistent-skill")
|
||||
if !errors.Is(err, runtime.ErrNotFound) {
|
||||
t.Errorf("missing skill: got %v, want ErrNotFound", err)
|
||||
}
|
||||
|
||||
// Cross-org skill -> ErrNotFound
|
||||
_, err = f.loader.LoadSkill(ctx, f.userOther, "resume-evaluator")
|
||||
if !errors.Is(err, runtime.ErrNotFound) {
|
||||
t.Errorf("cross-org skill: got %v, want ErrNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 2. Status Eligibility Tests ─────────────────────────────────────────── */
|
||||
|
||||
func TestRuntime_StatusEligibility(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 1. Draft Agent
|
||||
draftMD := `---
|
||||
id: draft-agent
|
||||
name: Draft Agent
|
||||
status: draft
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "draft-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: draftMD,
|
||||
Status: "draft",
|
||||
Version: 1,
|
||||
Name: "Draft Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err := f.loader.LoadExecutableAgent(ctx, f.userA, "draft-agent")
|
||||
if !errors.Is(err, runtime.ErrDraftAgent) {
|
||||
t.Errorf("draft agent: got %v, want ErrDraftAgent", err)
|
||||
}
|
||||
|
||||
// 2. Archived Agent
|
||||
archivedMD := `---
|
||||
id: archived-agent
|
||||
name: Archived Agent
|
||||
status: archived
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "archived-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: archivedMD,
|
||||
Status: "archived",
|
||||
Version: 1,
|
||||
Name: "Archived Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err = f.loader.LoadExecutableAgent(ctx, f.userA, "archived-agent")
|
||||
if !errors.Is(err, runtime.ErrArchivedAgent) {
|
||||
t.Errorf("archived agent: got %v, want ErrArchivedAgent", err)
|
||||
}
|
||||
|
||||
// 3. Published Agent without dependencies -> Success
|
||||
pubMD := `---
|
||||
id: published-agent
|
||||
name: Published Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "published-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: pubMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Published Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
pubAgent, err := f.loader.LoadExecutableAgent(ctx, f.userA, "published-agent")
|
||||
if err != nil {
|
||||
t.Fatalf("load published agent: %v", err)
|
||||
}
|
||||
if pubAgent.Status != "published" {
|
||||
t.Errorf("published agent status: %s", pubAgent.Status)
|
||||
}
|
||||
|
||||
// 4. Inactive Skill
|
||||
inactiveSkillMD := `---
|
||||
id: inactive-skill
|
||||
name: Inactive Skill
|
||||
status: inactive
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "inactive-skill",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: inactiveSkillMD,
|
||||
Status: "inactive",
|
||||
Name: "Inactive Skill",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err = f.loader.LoadExecutableSkill(ctx, f.userA, "inactive-skill")
|
||||
if !errors.Is(err, runtime.ErrInactiveSkill) {
|
||||
t.Errorf("inactive skill: got %v, want ErrInactiveSkill", err)
|
||||
}
|
||||
|
||||
// 5. Active Skill -> Success
|
||||
activeSkillMD := `---
|
||||
id: active-skill
|
||||
name: Active Skill
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "active-skill",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: activeSkillMD,
|
||||
Status: "active",
|
||||
Name: "Active Skill",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
activeSkill, err := f.loader.LoadExecutableSkill(ctx, f.userA, "active-skill")
|
||||
if err != nil {
|
||||
t.Fatalf("load active skill: %v", err)
|
||||
}
|
||||
if activeSkill.Status != "active" {
|
||||
t.Errorf("active skill status: %s", activeSkill.Status)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 3. Version Semantics Tests ──────────────────────────────────────────── */
|
||||
|
||||
func TestRuntime_VersionSemantics(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
v5MD := `---
|
||||
id: v5-agent
|
||||
name: Version 5 Agent
|
||||
version: 5
|
||||
status: published
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, err := f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "v5-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: v5MD,
|
||||
Status: "published",
|
||||
Version: 5,
|
||||
Name: "Version 5 Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("insert v5 agent: %v", err)
|
||||
}
|
||||
|
||||
agent, err := f.loader.LoadAgent(ctx, f.userA, "v5-agent")
|
||||
if err != nil {
|
||||
t.Fatalf("load v5 agent: %v", err)
|
||||
}
|
||||
if agent.Version != 5 {
|
||||
t.Errorf("agent version = %d, want 5", agent.Version)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4. Dependency Resolution Tests ──────────────────────────────────────── */
|
||||
|
||||
func TestRuntime_DependencyResolution(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 1. Create active skill 1
|
||||
skill1MD := `---
|
||||
id: skill-one
|
||||
name: Skill One
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "skill-one",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: skill1MD,
|
||||
Status: "active",
|
||||
Name: "Skill One",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// 2. Create active skill 2 (personal for userB)
|
||||
skill2MD := `---
|
||||
id: skill-two
|
||||
name: Skill Two
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userB, repo.SkillInsertInput{
|
||||
DefinitionID: "skill-two",
|
||||
OrgID: f.org1,
|
||||
Visibility: "personal",
|
||||
OwnerUserID: &f.userB.UserID,
|
||||
CreatedBy: &f.userB.UserID,
|
||||
Markdown: skill2MD,
|
||||
Status: "active",
|
||||
Name: "Skill Two",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// 3. Create inactive skill 3
|
||||
skill3MD := `---
|
||||
id: skill-inactive
|
||||
name: Skill Inactive
|
||||
status: inactive
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "skill-inactive",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: skill3MD,
|
||||
Status: "inactive",
|
||||
Name: "Skill Inactive",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// Agent with valid and duplicate dependencies
|
||||
validDepMD := `---
|
||||
id: dep-agent
|
||||
name: Dependency Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
skills:
|
||||
- skill-one
|
||||
- skill-two
|
||||
- skill-one
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userB, repo.AgentInsertInput{
|
||||
DefinitionID: "dep-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "personal",
|
||||
OwnerUserID: &f.userB.UserID,
|
||||
CreatedBy: &f.userB.UserID,
|
||||
Markdown: validDepMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Dependency Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// User B resolves dep-agent: sees skill-one (org) and skill-two (personal), deduplicates skill-one
|
||||
loadedAgent, err := f.loader.LoadExecutableAgent(ctx, f.userB, "dep-agent")
|
||||
if err != nil {
|
||||
t.Fatalf("load executable agent with deps: %v", err)
|
||||
}
|
||||
if len(loadedAgent.ResolvedSkills) != 2 {
|
||||
t.Fatalf("expected 2 resolved skills (deduplicated), got %d", len(loadedAgent.ResolvedSkills))
|
||||
}
|
||||
if loadedAgent.ResolvedSkills[0].ID != "skill-one" || loadedAgent.ResolvedSkills[1].ID != "skill-two" {
|
||||
t.Errorf("resolved skill IDs mismatch: %v, %v", loadedAgent.ResolvedSkills[0].ID, loadedAgent.ResolvedSkills[1].ID)
|
||||
}
|
||||
|
||||
// Agent with missing dependency
|
||||
missingDepMD := `---
|
||||
id: missing-dep-agent
|
||||
name: Missing Dep Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
skills:
|
||||
- nonexistent-skill
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "missing-dep-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: missingDepMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Missing Dep Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err = f.loader.LoadExecutableAgent(ctx, f.userA, "missing-dep-agent")
|
||||
if !errors.Is(err, runtime.ErrDependencyMissing) {
|
||||
t.Errorf("missing dep agent: got %v, want ErrDependencyMissing", err)
|
||||
}
|
||||
|
||||
// Agent with inactive dependency
|
||||
inactiveDepMD := `---
|
||||
id: inactive-dep-agent
|
||||
name: Inactive Dep Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
skills:
|
||||
- skill-inactive
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "inactive-dep-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: inactiveDepMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Inactive Dep Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err = f.loader.LoadExecutableAgent(ctx, f.userA, "inactive-dep-agent")
|
||||
if !errors.Is(err, runtime.ErrDependencyInactive) {
|
||||
t.Errorf("inactive dep agent: got %v, want ErrDependencyInactive", err)
|
||||
}
|
||||
|
||||
// Agent with cross-org dependency: outsider creates agent referencing org1's skill
|
||||
outsiderAgentMD := `---
|
||||
id: outsider-agent
|
||||
name: Outsider Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
skills:
|
||||
- skill-one
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userOther, repo.AgentInsertInput{
|
||||
DefinitionID: "outsider-agent",
|
||||
OrgID: f.org2,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userOther.UserID,
|
||||
Markdown: outsiderAgentMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Outsider Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
_, err = f.loader.LoadExecutableAgent(ctx, f.userOther, "outsider-agent")
|
||||
if !errors.Is(err, runtime.ErrDependencyMissing) {
|
||||
t.Errorf("outsider cross-org dep: got %v, want ErrDependencyMissing", err)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 5. Executor Boundary Tests ──────────────────────────────────────────── */
|
||||
|
||||
type testEchoExecutor struct {
|
||||
lastAgent *runtime.Agent
|
||||
lastInput runtime.ExecutionInput
|
||||
}
|
||||
|
||||
func (e *testEchoExecutor) ExecuteAgent(_ context.Context, agent *runtime.Agent, input runtime.ExecutionInput) (*runtime.ExecutionResult, error) {
|
||||
e.lastAgent = agent
|
||||
e.lastInput = input
|
||||
return &runtime.ExecutionResult{
|
||||
Success: true,
|
||||
Output: fmt.Sprintf("executed %s with %s", agent.Name, input.Input),
|
||||
AgentID: agent.ID,
|
||||
AgentVersion: agent.Version,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (e *testEchoExecutor) ExecuteSkill(_ context.Context, skill *runtime.Skill, input runtime.ExecutionInput) (*runtime.ExecutionResult, error) {
|
||||
return &runtime.ExecutionResult{
|
||||
Success: true,
|
||||
Output: fmt.Sprintf("executed skill %s", skill.Name),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func TestRuntime_ExecutorBoundary(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
pubMD := `---
|
||||
id: exec-agent
|
||||
name: Exec Agent
|
||||
status: published
|
||||
version: 1
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertAgent(ctx, f.userA, repo.AgentInsertInput{
|
||||
DefinitionID: "exec-agent",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: pubMD,
|
||||
Status: "published",
|
||||
Version: 1,
|
||||
Name: "Exec Agent",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// 1. Default engine without configured executor returns ErrExecutorUnavailable
|
||||
res, err := f.engine.RunAgent(ctx, f.userA, "exec-agent", runtime.ExecutionInput{
|
||||
Identity: f.userA,
|
||||
Input: "Hello agent",
|
||||
})
|
||||
if !errors.Is(err, runtime.ErrExecutorUnavailable) {
|
||||
t.Errorf("default engine run agent: got %v, want ErrExecutorUnavailable", err)
|
||||
}
|
||||
if res.Success {
|
||||
t.Errorf("default engine must not succeed without executor")
|
||||
}
|
||||
|
||||
// 2. Custom executor plugged in
|
||||
echo := &testEchoExecutor{}
|
||||
customEngine := runtime.NewEngine(f.h.Pool, runtime.WithAgentExecutor(echo), runtime.WithSkillExecutor(echo))
|
||||
|
||||
resCustom, err := customEngine.RunAgent(ctx, f.userA, "exec-agent", runtime.ExecutionInput{
|
||||
Identity: f.userA,
|
||||
Input: "Hello agent",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("custom engine run agent failed: %v", err)
|
||||
}
|
||||
if !resCustom.Success || resCustom.Output != "executed Exec Agent with Hello agent" {
|
||||
t.Errorf("custom executor output mismatch: %+v", resCustom)
|
||||
}
|
||||
if echo.lastAgent.ID != "exec-agent" || echo.lastInput.Input != "Hello agent" {
|
||||
t.Errorf("executor received incorrect arguments: agent=%v, input=%v", echo.lastAgent, echo.lastInput)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntime_PersonalSkillShadowing(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Org skill
|
||||
orgSkillMD := `---
|
||||
id: shadow-skill
|
||||
name: Org Shadow Skill
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "shadow-skill",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: orgSkillMD,
|
||||
Status: "active",
|
||||
Name: "Org Shadow Skill",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// User B's personal skill with the SAME definition_id
|
||||
persSkillMD := `---
|
||||
id: shadow-skill
|
||||
name: User B Personal Shadow Skill
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
persRec, _ := f.defRepo.InsertSkill(ctx, f.userB, repo.SkillInsertInput{
|
||||
DefinitionID: "shadow-skill",
|
||||
OrgID: f.org1,
|
||||
Visibility: "personal",
|
||||
OwnerUserID: &f.userB.UserID,
|
||||
CreatedBy: &f.userB.UserID,
|
||||
Markdown: persSkillMD,
|
||||
Status: "active",
|
||||
Name: "User B Personal Shadow Skill",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
persUUID := persRec["id"].(string)
|
||||
|
||||
// When User A loads shadow-skill -> gets Org Shadow Skill
|
||||
skillA, err := f.loader.LoadSkill(ctx, f.userA, "shadow-skill")
|
||||
if err != nil {
|
||||
t.Fatalf("userA load shadow skill: %v", err)
|
||||
}
|
||||
if skillA.Name != "Org Shadow Skill" {
|
||||
t.Errorf("userA got %s, want Org Shadow Skill", skillA.Name)
|
||||
}
|
||||
|
||||
// When User B loads shadow-skill -> gets User B Personal Shadow Skill (shadow precedence)
|
||||
skillB, err := f.loader.LoadSkill(ctx, f.userB, "shadow-skill")
|
||||
if err != nil {
|
||||
t.Fatalf("userB load shadow skill: %v", err)
|
||||
}
|
||||
if skillB.Name != "User B Personal Shadow Skill" || skillB.DatabaseID != persUUID {
|
||||
t.Errorf("userB got %s, want User B Personal Shadow Skill", skillB.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntime_MalformedMarkdown(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Broken YAML
|
||||
brokenMD := `---
|
||||
id: [broken-id
|
||||
name: Invalid
|
||||
---
|
||||
`
|
||||
_, _ = f.h.Pool.Exec(ctx, `INSERT INTO agent_definitions (definition_id, org_id, visibility, created_by, markdown, status, version, name, description, pages)
|
||||
VALUES ('broken-agent', $1::uuid, 'organization', $2::uuid, $3::text, 'published', 1, 'Broken', '', ARRAY['candidates'])`,
|
||||
f.org1, f.userA.UserID, brokenMD)
|
||||
|
||||
_, err := f.loader.LoadAgent(ctx, f.userA, "broken-agent")
|
||||
if !errors.Is(err, runtime.ErrInvalidDefinition) {
|
||||
t.Errorf("load broken agent: got %v, want ErrInvalidDefinition", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntime_SkillExecution(t *testing.T) {
|
||||
f := newFixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
skillMD := `---
|
||||
id: run-skill
|
||||
name: Runnable Skill
|
||||
status: active
|
||||
pages:
|
||||
- candidates
|
||||
---
|
||||
`
|
||||
_, _ = f.defRepo.InsertSkill(ctx, f.userA, repo.SkillInsertInput{
|
||||
DefinitionID: "run-skill",
|
||||
OrgID: f.org1,
|
||||
Visibility: "organization",
|
||||
CreatedBy: &f.userA.UserID,
|
||||
Markdown: skillMD,
|
||||
Status: "active",
|
||||
Name: "Runnable Skill",
|
||||
Pages: []string{"candidates"},
|
||||
})
|
||||
|
||||
// Default engine
|
||||
res, err := f.engine.RunSkill(ctx, f.userA, "run-skill", runtime.ExecutionInput{
|
||||
Identity: f.userA,
|
||||
Input: "skill input",
|
||||
})
|
||||
if !errors.Is(err, runtime.ErrExecutorUnavailable) {
|
||||
t.Errorf("default engine run skill: got %v, want ErrExecutorUnavailable", err)
|
||||
}
|
||||
if res.Success {
|
||||
t.Errorf("default engine run skill must not succeed")
|
||||
}
|
||||
|
||||
// Custom executor
|
||||
echo := &testEchoExecutor{}
|
||||
customEngine := runtime.NewEngine(f.h.Pool, runtime.WithSkillExecutor(echo))
|
||||
resCustom, err := customEngine.RunSkill(ctx, f.userA, "run-skill", runtime.ExecutionInput{
|
||||
Identity: f.userA,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("custom engine run skill failed: %v", err)
|
||||
}
|
||||
if !resCustom.Success || resCustom.Output != "executed skill Runnable Skill" {
|
||||
t.Errorf("custom skill output mismatch: %+v", resCustom)
|
||||
}
|
||||
}
|
||||
103
go-api/internal/runtime/types.go
Normal file
103
go-api/internal/runtime/types.go
Normal file
@@ -0,0 +1,103 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// Standard runtime errors.
|
||||
var (
|
||||
ErrNotFound = errors.New("runtime: definition not found")
|
||||
ErrUnauthorized = errors.New("runtime: unauthorized")
|
||||
ErrInvalidDefinition = errors.New("runtime: invalid definition")
|
||||
ErrDraftAgent = errors.New("runtime: agent is in draft status and cannot be executed")
|
||||
ErrArchivedAgent = errors.New("runtime: agent is archived and cannot be executed")
|
||||
ErrInactiveSkill = errors.New("runtime: skill is inactive and cannot be executed")
|
||||
ErrNotExecutable = errors.New("runtime: definition is not eligible for execution")
|
||||
ErrDependencyMissing = errors.New("runtime: required skill dependency not found")
|
||||
ErrDependencyInactive = errors.New("runtime: required skill dependency is inactive")
|
||||
ErrCircularDependency = errors.New("runtime: circular dependency detected in skills")
|
||||
ErrExecutorUnavailable = errors.New("runtime: AI executor is unavailable (deferred to Phase 5)")
|
||||
)
|
||||
|
||||
// Agent represents an authored agent prepared for runtime execution.
|
||||
type Agent struct {
|
||||
ID string `json:"id"`
|
||||
DatabaseID string `json:"databaseId"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Version int `json:"version"`
|
||||
Visibility string `json:"visibility"`
|
||||
OwnerUserID *string `json:"ownerUserId,omitempty"`
|
||||
Pages []string `json:"pages"`
|
||||
Icon string `json:"icon,omitempty"`
|
||||
Reasoning string `json:"reasoning,omitempty"`
|
||||
Trigger string `json:"trigger,omitempty"`
|
||||
WebSearch bool `json:"webSearch,omitempty"`
|
||||
Instructions string `json:"instructions"`
|
||||
Skills []string `json:"skills"`
|
||||
ResolvedSkills []*Skill `json:"resolvedSkills,omitempty"`
|
||||
Subagents []string `json:"subagents,omitempty"`
|
||||
RawMarkdown string `json:"rawMarkdown"`
|
||||
}
|
||||
|
||||
// Skill represents an authored skill prepared for runtime execution.
|
||||
type Skill struct {
|
||||
ID string `json:"id"`
|
||||
DatabaseID string `json:"databaseId"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Visibility string `json:"visibility"`
|
||||
OwnerUserID *string `json:"ownerUserId,omitempty"`
|
||||
Pages []string `json:"pages"`
|
||||
Kind string `json:"kind,omitempty"`
|
||||
Category string `json:"category,omitempty"`
|
||||
Actions []string `json:"actions,omitempty"`
|
||||
Triggers []string `json:"triggers,omitempty"`
|
||||
Prompt *string `json:"prompt,omitempty"`
|
||||
SkillID *string `json:"skillId,omitempty"`
|
||||
Body string `json:"body"`
|
||||
RawMarkdown string `json:"rawMarkdown"`
|
||||
}
|
||||
|
||||
// ExecutionInput provides caller context and payload to the execution boundary.
|
||||
type ExecutionInput struct {
|
||||
Identity authctx.Identity `json:"identity"`
|
||||
TargetID string `json:"targetId"`
|
||||
Input string `json:"input"`
|
||||
Parameters map[string]any `json:"parameters,omitempty"`
|
||||
Context map[string]any `json:"context,omitempty"`
|
||||
}
|
||||
|
||||
// ExecutionResult captures the outcome of an execution attempt.
|
||||
type ExecutionResult struct {
|
||||
Success bool `json:"success"`
|
||||
Output string `json:"output,omitempty"`
|
||||
AgentID string `json:"agentId,omitempty"`
|
||||
AgentVersion int `json:"agentVersion,omitempty"`
|
||||
ResolvedSkills []string `json:"resolvedSkills,omitempty"`
|
||||
Error error `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// RuntimeError is a structured error containing context for execution failures.
|
||||
type RuntimeError struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Target string `json:"target,omitempty"`
|
||||
Cause error `json:"-"`
|
||||
}
|
||||
|
||||
func (e *RuntimeError) Error() string {
|
||||
if e.Target != "" {
|
||||
return fmt.Sprintf("%s: %s (%s)", e.Code, e.Message, e.Target)
|
||||
}
|
||||
return fmt.Sprintf("%s: %s", e.Code, e.Message)
|
||||
}
|
||||
|
||||
func (e *RuntimeError) Unwrap() error {
|
||||
return e.Cause
|
||||
}
|
||||
525
go-api/internal/seeder/seeder.go
Normal file
525
go-api/internal/seeder/seeder.go
Normal file
@@ -0,0 +1,525 @@
|
||||
// Package seeder loads the frontend's demo dataset into PostgreSQL.
|
||||
//
|
||||
// Source of truth is the frontend, not this package. `seed/fixtures/seed.json`
|
||||
// is produced by executing src/api/seed.js through Vite and serialising what it
|
||||
// exports, so ids, dates, numbers and enum values arrive exactly as the demo
|
||||
// has them — no transcription step, and nothing to drift. Shift records are the
|
||||
// one exception: they are generated (see shifts.go) because their dates are
|
||||
// anchored to now.
|
||||
//
|
||||
// Idempotency strategy: EXPLICIT UPSERT inside a single transaction.
|
||||
//
|
||||
// Every record's primary key is derived deterministically from its source id
|
||||
// (uuid v5 over a fixed namespace), so re-running the seeder targets exactly the
|
||||
// same rows and `ON CONFLICT (id) DO UPDATE` restores each one to its seeded
|
||||
// values. Records created through the API survive a re-seed. A column the
|
||||
// fixture does not carry is left as it is.
|
||||
//
|
||||
// Shift records are the one collection that is also PRUNED — see
|
||||
// pruneShiftRecords. Upsert alone cannot converge a rolling window, and shift
|
||||
// records are the only collection that is a rolling window.
|
||||
package seeder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/orgctx"
|
||||
)
|
||||
|
||||
// namespace is a fixed UUID used to derive record ids from source ids. Changing
|
||||
// it re-keys the entire dataset, so it is a constant, not configuration.
|
||||
var namespace = [16]byte{
|
||||
0x6b, 0x72, 0x6f, 0x77, 0x2d, 0x73, 0x65, 0x65,
|
||||
0x64, 0x2d, 0x76, 0x31, 0x00, 0x00, 0x00, 0x01,
|
||||
}
|
||||
|
||||
// DeterministicUUID derives a stable v5 UUID from a source id.
|
||||
func DeterministicUUID(name string) string {
|
||||
h := sha1.New()
|
||||
h.Write(namespace[:])
|
||||
h.Write([]byte(name))
|
||||
var b [16]byte
|
||||
copy(b[:], h.Sum(nil))
|
||||
b[6] = (b[6] & 0x0f) | 0x50 // version 5
|
||||
b[8] = (b[8] & 0x3f) | 0x80 // RFC 4122 variant
|
||||
return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:16])
|
||||
}
|
||||
|
||||
// Fixture is the serialised frontend dataset.
|
||||
type Fixture struct {
|
||||
DemoUser map[string]any `json:"demoUser"`
|
||||
Entities map[string][]map[string]any `json:"entities"`
|
||||
}
|
||||
|
||||
// Result counts what was written, by entity.
|
||||
type Result struct {
|
||||
OrgID string
|
||||
Counts map[string]int
|
||||
// Pruned is how many stale shift records this run removed. Reported rather
|
||||
// than silent: a delete during a seed should never be something you have to
|
||||
// read the source to discover.
|
||||
Pruned int
|
||||
}
|
||||
|
||||
// entityOrder is insertion order, chosen so every foreign key is satisfied by
|
||||
// the time it is referenced.
|
||||
var entityOrder = []string{
|
||||
"RoleCategory", "Certification", "Badge", "Course", "LearningPath",
|
||||
"JobPosting", "WorkerProfile", "JobApplication", "AIInterview", "Staff",
|
||||
"Assignment", "ShiftRecord", "Evidence", "UserActivity",
|
||||
}
|
||||
|
||||
// entityTable maps a frontend entity name to its table.
|
||||
var entityTable = map[string]string{
|
||||
"RoleCategory": "role_categories", "Certification": "certifications",
|
||||
"Badge": "badges", "Course": "courses", "LearningPath": "learning_paths",
|
||||
"JobPosting": "job_postings", "WorkerProfile": "worker_profiles",
|
||||
"JobApplication": "job_applications", "AIInterview": "ai_interviews",
|
||||
"Staff": "staff", "Assignment": "assignments", "ShiftRecord": "shift_records",
|
||||
"Evidence": "evidence", "UserActivity": "user_activity",
|
||||
}
|
||||
|
||||
// referenceFields are columns holding a source id that must be rewritten to the
|
||||
// derived UUID. Every one of these is a real reference in the frontend data.
|
||||
var referenceFields = map[string]bool{
|
||||
"job_posting_id": true, "application_id": true, "course_id": true,
|
||||
"staff_id": true, "assignment_id": true, "worker_profile_id": true,
|
||||
"interview_id": true, "position_id": true, "candidate_id": true,
|
||||
"user_id": true, "created_by": true,
|
||||
}
|
||||
|
||||
// droppedFields are keys the fixture carries that no column exists for and no
|
||||
// frontend code reads. Dropping them is deliberate and recorded here rather
|
||||
// than being silent.
|
||||
//
|
||||
// _order — a positional index used only while seed.js builds its course list
|
||||
// (src/api/seed.js:1326). Nothing reads it.
|
||||
var droppedFields = map[string]bool{"_order": true}
|
||||
|
||||
// Seeder loads a fixture into a database.
|
||||
type Seeder struct {
|
||||
pool *pgxpool.Pool
|
||||
fixture *Fixture
|
||||
now time.Time
|
||||
}
|
||||
|
||||
// Load reads a fixture from disk.
|
||||
func Load(path string) (*Fixture, error) {
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read fixture %s: %w", path, err)
|
||||
}
|
||||
var f Fixture
|
||||
if err := json.Unmarshal(raw, &f); err != nil {
|
||||
return nil, fmt.Errorf("parse fixture %s: %w", path, err)
|
||||
}
|
||||
return &f, nil
|
||||
}
|
||||
|
||||
// New builds a seeder. `now` anchors the generated shift records.
|
||||
func New(pool *pgxpool.Pool, fixture *Fixture, now time.Time) *Seeder {
|
||||
return &Seeder{pool: pool, fixture: fixture, now: now}
|
||||
}
|
||||
|
||||
// Run seeds everything in one transaction: either the whole dataset lands or
|
||||
// none of it does.
|
||||
func (s *Seeder) Run(ctx context.Context) (*Result, error) {
|
||||
tx, err := s.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback(ctx) }()
|
||||
|
||||
orgID, err := s.upsertOrganization(ctx, tx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := &Result{OrgID: orgID, Counts: map[string]int{}}
|
||||
|
||||
n, err := s.upsertUser(ctx, tx, orgID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.Counts["User"] = n
|
||||
|
||||
for _, entity := range entityOrder {
|
||||
records := s.fixture.Entities[entity]
|
||||
if entity == "ShiftRecord" {
|
||||
records = BuildShifts(s.now)
|
||||
}
|
||||
count, err := s.upsertEntity(ctx, tx, orgID, entity, records)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("seed %s: %w", entity, err)
|
||||
}
|
||||
result.Counts[entity] = count
|
||||
|
||||
if entity == "ShiftRecord" {
|
||||
pruned, err := s.pruneShiftRecords(ctx, tx, orgID, records)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prune ShiftRecord: %w", err)
|
||||
}
|
||||
result.Pruned = pruned
|
||||
}
|
||||
}
|
||||
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// upsertOrganization creates the development organization the whole dataset
|
||||
// belongs to. See internal/orgctx — this is not a tenant, it is a placeholder
|
||||
// with a stable id so re-seeding is idempotent.
|
||||
func (s *Seeder) upsertOrganization(ctx context.Context, tx pgx.Tx) (string, error) {
|
||||
id := DeterministicUUID("org:" + orgctx.DevOrgSlug)
|
||||
_, err := tx.Exec(ctx,
|
||||
`INSERT INTO organizations (id, name, slug) VALUES ($1::uuid, $2, $3::citext)
|
||||
ON CONFLICT (id) DO UPDATE SET name = EXCLUDED.name, updated_date = now()`,
|
||||
id, orgctx.DevOrgName, orgctx.DevOrgSlug)
|
||||
return id, err
|
||||
}
|
||||
|
||||
// upsertUser writes the demo user and splits its preferences into their own
|
||||
// table, as api-contract.md §9 describes.
|
||||
func (s *Seeder) upsertUser(ctx context.Context, tx pgx.Tx, orgID string) (int, error) {
|
||||
u := s.fixture.DemoUser
|
||||
if u == nil {
|
||||
return 0, nil
|
||||
}
|
||||
legacy, _ := u["id"].(string)
|
||||
id := DeterministicUUID("User:" + legacy)
|
||||
created := stringOr(u["created_date"], iso(s.now))
|
||||
|
||||
_, err := tx.Exec(ctx,
|
||||
`INSERT INTO users (id, legacy_id, org_id, email, full_name, role, account_type, created_date, updated_date)
|
||||
VALUES ($1::uuid, $2::text, $3::uuid, $4::citext, $5::text, $6::text, $7::text, $8::timestamptz, $8::timestamptz)
|
||||
ON CONFLICT (id) DO UPDATE SET
|
||||
email = EXCLUDED.email, full_name = EXCLUDED.full_name,
|
||||
role = EXCLUDED.role, account_type = EXCLUDED.account_type,
|
||||
created_date = EXCLUDED.created_date, updated_date = now()`,
|
||||
id, legacy, orgID,
|
||||
stringOr(u["email"], ""), stringOr(u["full_name"], ""),
|
||||
stringOr(u["role"], "admin"), stringOr(u["account_type"], "employer"), created)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
prefs, _ := u["preferences"].(map[string]any)
|
||||
if prefs == nil {
|
||||
prefs = map[string]any{}
|
||||
}
|
||||
extra := map[string]any{}
|
||||
for k, v := range prefs {
|
||||
switch k {
|
||||
case "owliverDefault", "compactDensity", "emailDigest":
|
||||
default:
|
||||
extra[k] = v
|
||||
}
|
||||
}
|
||||
extraJSON, err := json.Marshal(extra)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
_, err = tx.Exec(ctx,
|
||||
`INSERT INTO user_preferences (user_id, owliver_default, compact_density, email_digest, extra)
|
||||
VALUES ($1::uuid, $2::boolean, $3::boolean, $4::boolean, $5::jsonb)
|
||||
ON CONFLICT (user_id) DO UPDATE SET
|
||||
owliver_default = EXCLUDED.owliver_default,
|
||||
compact_density = EXCLUDED.compact_density,
|
||||
email_digest = EXCLUDED.email_digest,
|
||||
extra = EXCLUDED.extra,
|
||||
updated_date = now()`,
|
||||
id, boolOr(prefs["owliverDefault"], true), boolOr(prefs["compactDensity"], false),
|
||||
boolOr(prefs["emailDigest"], true), extraJSON)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return 1, nil
|
||||
}
|
||||
|
||||
// pruneShiftRecords deletes this organization's shift rows that this run did
|
||||
// not generate.
|
||||
//
|
||||
// WHY THIS EXISTS, and why it is the only place the seeder deletes anything:
|
||||
//
|
||||
// A shift's stable id is `shift_<worker>_<NN>`, where NN counts the shift's
|
||||
// position from the OLDEST end of the rolling 56-day window
|
||||
// (attendanceSeed.js:190 and shifts.go:167 — the port is faithful, the scheme
|
||||
// is the problem). That number is a position, not an identity, so it means a
|
||||
// different date every day the window slides. Measured against a database
|
||||
// seeded one day earlier: all 114 surviving ids had moved to a different date,
|
||||
// and one — `shift_marcus_41` — was orphaned, because Marcus works Mon–Fri and
|
||||
// a Saturday window holds 40 of his shifts rather than 41.
|
||||
//
|
||||
// Upsert can rewrite the rows it still generates. It has no way to remove the
|
||||
// one it no longer generates, so the collection ratchets up to the historical
|
||||
// maximum and never returns to the size the generator actually produces.
|
||||
//
|
||||
// Deleting is safe here in a way it would not be for any other collection:
|
||||
// ShiftRecord is `Ops: OpList` (api-contract.md §2 — there is no
|
||||
// POST /shift-records, and U1 in §11 is exactly the question of where these
|
||||
// records come from), so every row is seeder-owned and no API call can create
|
||||
// one. shift_records is also a leaf table: no foreign key points at it, so
|
||||
// nothing cascades. Between them, this delete cannot reach data the seeder did
|
||||
// not write.
|
||||
//
|
||||
// Note what this does NOT do: it invents no records and changes no generated
|
||||
// value. After it, the collection is exactly what BuildShifts produced for
|
||||
// s.now — which is what a regenerated rolling window means.
|
||||
func (s *Seeder) pruneShiftRecords(ctx context.Context, tx pgx.Tx, orgID string,
|
||||
records []map[string]any) (int, error) {
|
||||
|
||||
// A generation that produced nothing is a bug in BuildShifts, not an
|
||||
// instruction to empty the table: `id <> ALL('{}')` is true for every row.
|
||||
// Refuse rather than wipe.
|
||||
if len(records) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
keep := make([]string, 0, len(records))
|
||||
for _, rec := range records {
|
||||
legacy, _ := rec["id"].(string)
|
||||
if legacy == "" {
|
||||
return 0, fmt.Errorf("generated shift record has no id")
|
||||
}
|
||||
keep = append(keep, DeterministicUUID("ShiftRecord:"+legacy))
|
||||
}
|
||||
|
||||
tag, err := tx.Exec(ctx,
|
||||
`DELETE FROM shift_records WHERE org_id = $1::uuid AND id <> ALL($2::uuid[])`,
|
||||
orgID, keep)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return int(tag.RowsAffected()), nil
|
||||
}
|
||||
|
||||
// upsertEntity writes one collection.
|
||||
func (s *Seeder) upsertEntity(ctx context.Context, tx pgx.Tx, orgID, entity string, records []map[string]any) (int, error) {
|
||||
table := entityTable[entity]
|
||||
res, ok := domain.ResourceByTable[table]
|
||||
if !ok {
|
||||
return 0, fmt.Errorf("no resource descriptor for table %s", table)
|
||||
}
|
||||
for _, rec := range records {
|
||||
if err := s.upsertRecord(ctx, tx, orgID, entity, res, rec); err != nil {
|
||||
id, _ := rec["id"].(string)
|
||||
return 0, fmt.Errorf("record %s: %w", id, err)
|
||||
}
|
||||
}
|
||||
return len(records), nil
|
||||
}
|
||||
|
||||
func (s *Seeder) upsertRecord(ctx context.Context, tx pgx.Tx, orgID, entity string,
|
||||
res *domain.Resource, rec map[string]any) error {
|
||||
|
||||
legacy, _ := rec["id"].(string)
|
||||
if legacy == "" {
|
||||
return fmt.Errorf("record has no id")
|
||||
}
|
||||
|
||||
// user_activity's primary key is a GENERATED ALWAYS AS IDENTITY bigint, not
|
||||
// a uuid, so no explicit id can be supplied for it. Its stable identity is
|
||||
// legacy_id, which is what the upsert conflicts on instead.
|
||||
idCol, _ := res.Column("id")
|
||||
generatedID := idCol.Kind != domain.KindUUID
|
||||
conflictTarget := "id"
|
||||
|
||||
values := map[string]any{
|
||||
"legacy_id": legacy,
|
||||
"org_id": orgID,
|
||||
}
|
||||
if generatedID {
|
||||
conflictTarget = "legacy_id"
|
||||
} else {
|
||||
values["id"] = DeterministicUUID(entity + ":" + legacy)
|
||||
}
|
||||
|
||||
created := stringOr(rec["created_date"], iso(s.now))
|
||||
values["created_date"] = created
|
||||
if _, hasUpdated := res.Column("updated_date"); hasUpdated {
|
||||
// The fixture carries updated_date only on job applications, where the
|
||||
// gap from created_date is what buildHires reads as time-to-hire.
|
||||
// Everywhere else the column is NOT NULL and the record has never been
|
||||
// modified, so it takes the creation instant.
|
||||
values["updated_date"] = stringOr(rec["updated_date"], created)
|
||||
}
|
||||
|
||||
for key, value := range rec {
|
||||
switch key {
|
||||
case "id", "created_date", "updated_date":
|
||||
continue
|
||||
}
|
||||
if droppedFields[key] {
|
||||
continue
|
||||
}
|
||||
col, ok := res.Column(key)
|
||||
if !ok {
|
||||
return fmt.Errorf("field %q has no column on %s", key, res.Table)
|
||||
}
|
||||
if referenceFields[key] && col.Kind == domain.KindUUID {
|
||||
str, isStr := value.(string)
|
||||
if !isStr || str == "" {
|
||||
values[key] = nil
|
||||
continue
|
||||
}
|
||||
values[key] = DeterministicUUID(referencedEntity(key) + ":" + str)
|
||||
continue
|
||||
}
|
||||
values[key] = value
|
||||
}
|
||||
|
||||
// Deterministic column order keeps the generated SQL stable.
|
||||
names := make([]string, 0, len(values))
|
||||
for k := range values {
|
||||
names = append(names, k)
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
cols := make([]string, 0, len(names))
|
||||
placeholders := make([]string, 0, len(names))
|
||||
updates := make([]string, 0, len(names))
|
||||
args := make([]any, 0, len(names))
|
||||
|
||||
for _, name := range names {
|
||||
col, ok := res.Column(name)
|
||||
if !ok {
|
||||
return fmt.Errorf("no column %q on %s", name, res.Table)
|
||||
}
|
||||
bound, err := bindSeedValue(*col, values[name])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
args = append(args, bound)
|
||||
cols = append(cols, name)
|
||||
placeholders = append(placeholders, fmt.Sprintf("$%d::%s", len(args), col.PGType))
|
||||
if name != conflictTarget {
|
||||
updates = append(updates, fmt.Sprintf("%s = EXCLUDED.%s", name, name))
|
||||
}
|
||||
}
|
||||
|
||||
q := fmt.Sprintf(
|
||||
"INSERT INTO %s (%s) VALUES (%s) ON CONFLICT (%s) DO UPDATE SET %s",
|
||||
res.Table, strings.Join(cols, ", "), strings.Join(placeholders, ", "),
|
||||
conflictTarget, strings.Join(updates, ", "))
|
||||
|
||||
_, err := tx.Exec(ctx, q, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
// referencedEntity says which entity a reference column points at, so the
|
||||
// derived UUID is built from the same namespace the target was written with.
|
||||
func referencedEntity(field string) string {
|
||||
switch field {
|
||||
case "job_posting_id", "position_id":
|
||||
return "JobPosting"
|
||||
case "application_id":
|
||||
return "JobApplication"
|
||||
case "course_id":
|
||||
return "Course"
|
||||
case "staff_id":
|
||||
return "Staff"
|
||||
case "assignment_id":
|
||||
return "Assignment"
|
||||
case "worker_profile_id", "candidate_id":
|
||||
return "WorkerProfile"
|
||||
case "interview_id":
|
||||
return "AIInterview"
|
||||
case "user_id", "created_by":
|
||||
return "User"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// bindSeedValue converts a fixture value into something pgx can send. It is
|
||||
// deliberately separate from the repository's binder: the seeder writes
|
||||
// server-owned columns (id, legacy_id, org_id, created_date) that the API never
|
||||
// accepts from a client.
|
||||
func bindSeedValue(col domain.Column, v any) (any, error) {
|
||||
if v == nil {
|
||||
return nil, nil
|
||||
}
|
||||
switch col.Kind {
|
||||
case domain.KindTextArray:
|
||||
switch t := v.(type) {
|
||||
case []any:
|
||||
out := make([]string, 0, len(t))
|
||||
for _, e := range t {
|
||||
s, ok := e.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("%s: expected an array of strings", col.Name)
|
||||
}
|
||||
out = append(out, s)
|
||||
}
|
||||
return out, nil
|
||||
case []string:
|
||||
return t, nil
|
||||
}
|
||||
return nil, fmt.Errorf("%s: expected an array", col.Name)
|
||||
|
||||
case domain.KindJSON:
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", col.Name, err)
|
||||
}
|
||||
return raw, nil
|
||||
|
||||
case domain.KindInt:
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
return int64(t), nil
|
||||
case int:
|
||||
return int64(t), nil
|
||||
case int64:
|
||||
return t, nil
|
||||
}
|
||||
return nil, fmt.Errorf("%s: expected a number, got %T", col.Name, v)
|
||||
|
||||
case domain.KindFloat:
|
||||
switch t := v.(type) {
|
||||
case float64:
|
||||
return t, nil
|
||||
case int:
|
||||
return float64(t), nil
|
||||
}
|
||||
return nil, fmt.Errorf("%s: expected a number, got %T", col.Name, v)
|
||||
|
||||
case domain.KindBool:
|
||||
if b, ok := v.(bool); ok {
|
||||
return b, nil
|
||||
}
|
||||
return nil, fmt.Errorf("%s: expected a boolean, got %T", col.Name, v)
|
||||
|
||||
default:
|
||||
if s, ok := v.(string); ok {
|
||||
return s, nil
|
||||
}
|
||||
return nil, fmt.Errorf("%s: expected a string, got %T", col.Name, v)
|
||||
}
|
||||
}
|
||||
|
||||
func stringOr(v any, fallback string) string {
|
||||
if s, ok := v.(string); ok && s != "" {
|
||||
return s
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func boolOr(v any, fallback bool) bool {
|
||||
if b, ok := v.(bool); ok {
|
||||
return b
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
300
go-api/internal/seeder/seeder_test.go
Normal file
300
go-api/internal/seeder/seeder_test.go
Normal file
@@ -0,0 +1,300 @@
|
||||
package seeder_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/seeder"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// tableFor maps the fixture's entity names onto their tables, so the counts
|
||||
// asserted below come from the frontend's own data rather than from literals.
|
||||
var tableFor = map[string]string{
|
||||
"JobPosting": "job_postings", "JobApplication": "job_applications",
|
||||
"AIInterview": "ai_interviews", "Staff": "staff", "WorkerProfile": "worker_profiles",
|
||||
"Course": "courses", "Badge": "badges", "LearningPath": "learning_paths",
|
||||
"Certification": "certifications", "RoleCategory": "role_categories",
|
||||
"UserActivity": "user_activity", "Evidence": "evidence", "Assignment": "assignments",
|
||||
}
|
||||
|
||||
func count(t *testing.T, h *testutil.Harness, table string) int {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := h.Pool.QueryRow(context.Background(), "SELECT count(*) FROM "+table).Scan(&n); err != nil {
|
||||
t.Fatalf("count %s: %v", table, err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// TestSeedMatchesFixtureCounts checks every entity against the fixture rather
|
||||
// than against a hardcoded headline number.
|
||||
func TestSeedMatchesFixtureCounts(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
|
||||
for entity, table := range tableFor {
|
||||
want := len(fx.Entities[entity])
|
||||
if got := count(t, h, table); got != want {
|
||||
t.Errorf("%s: seeded %d rows, fixture has %d", table, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedRegressionAnchors pins the figures the demo dataset is built to
|
||||
// produce. These are verified against the source, not assumed: the prompt's
|
||||
// "6 postings / 22 applications" is 6 *active* postings and 24 applications.
|
||||
func TestSeedRegressionAnchors(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var active int
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT count(*) FROM job_postings WHERE status = 'active'").Scan(&active); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if active != 6 {
|
||||
t.Errorf("active postings = %d, want 6", active)
|
||||
}
|
||||
if total := count(t, h, "job_postings"); total != 8 {
|
||||
t.Errorf("job postings = %d, want 8 (6 active, 1 paused, 1 closed)", total)
|
||||
}
|
||||
if total := count(t, h, "job_applications"); total != 24 {
|
||||
t.Errorf("applications = %d, want 24", total)
|
||||
}
|
||||
|
||||
var scored int
|
||||
var avgScored float64
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT count(*), coalesce(avg(ai_score), 0) FROM job_applications WHERE ai_score > 0").
|
||||
Scan(&scored, &avgScored); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if scored != 9 {
|
||||
t.Errorf("scored applications = %d, want 9", scored)
|
||||
}
|
||||
if avgScored < 75.95 || avgScored > 76.05 {
|
||||
t.Errorf("average scored ai_score = %.2f, want 76.0", avgScored)
|
||||
}
|
||||
|
||||
var hires int
|
||||
var avgHire float64
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT count(*), coalesce(avg(ai_score), 0) FROM staff").Scan(&hires, &avgHire); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hires != 3 {
|
||||
t.Errorf("hires = %d, want 3", hires)
|
||||
}
|
||||
if avgHire < 94.0 || avgHire > 94.5 {
|
||||
t.Errorf("average hire ai_score = %.2f, want ~94.3", avgHire)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedPreservesSourceValues compares stored rows field-by-field against the
|
||||
// fixture, rather than trusting that the counts lining up means the data did.
|
||||
func TestSeedPreservesSourceValues(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, want := range fx.Entities["JobPosting"] {
|
||||
legacy := want["id"].(string)
|
||||
var title, status, company, roleCategory, createdDate string
|
||||
var payMin, payMax int
|
||||
err := h.Pool.QueryRow(ctx, `
|
||||
SELECT title, status::text, company, role_category, pay_range_min, pay_range_max,
|
||||
to_char(created_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"')
|
||||
FROM job_postings WHERE legacy_id = $1`, legacy).
|
||||
Scan(&title, &status, &company, &roleCategory, &payMin, &payMax, &createdDate)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", legacy, err)
|
||||
}
|
||||
if title != want["title"] {
|
||||
t.Errorf("%s title = %q, want %q", legacy, title, want["title"])
|
||||
}
|
||||
if status != want["status"] {
|
||||
t.Errorf("%s status = %q, want %q", legacy, status, want["status"])
|
||||
}
|
||||
if createdDate != want["created_date"] {
|
||||
t.Errorf("%s created_date = %q, want %q", legacy, createdDate, want["created_date"])
|
||||
}
|
||||
if v, ok := want["pay_range_min"].(float64); ok && payMin != int(v) {
|
||||
t.Errorf("%s pay_range_min = %d, want %d", legacy, payMin, int(v))
|
||||
}
|
||||
if v, ok := want["pay_range_max"].(float64); ok && payMax != int(v) {
|
||||
t.Errorf("%s pay_range_max = %d, want %d", legacy, payMax, int(v))
|
||||
}
|
||||
}
|
||||
|
||||
// Applications carry updated_date in the source, and the gap from
|
||||
// created_date is what buildHires reads as time-to-hire.
|
||||
for _, want := range fx.Entities["JobApplication"] {
|
||||
legacy := want["id"].(string)
|
||||
var name, email, status, created, updated string
|
||||
var score int
|
||||
err := h.Pool.QueryRow(ctx, `
|
||||
SELECT applicant_name, email::text, status::text, ai_score,
|
||||
to_char(created_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'),
|
||||
to_char(updated_date AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"')
|
||||
FROM job_applications WHERE legacy_id = $1`, legacy).
|
||||
Scan(&name, &email, &status, &score, &created, &updated)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", legacy, err)
|
||||
}
|
||||
if name != want["applicant_name"] {
|
||||
t.Errorf("%s applicant_name = %q, want %q", legacy, name, want["applicant_name"])
|
||||
}
|
||||
if status != want["status"] {
|
||||
t.Errorf("%s status = %q, want %q", legacy, status, want["status"])
|
||||
}
|
||||
if created != want["created_date"] {
|
||||
t.Errorf("%s created_date = %q, want %q", legacy, created, want["created_date"])
|
||||
}
|
||||
if updated != want["updated_date"] {
|
||||
t.Errorf("%s updated_date = %q, want %q", legacy, updated, want["updated_date"])
|
||||
}
|
||||
if v, ok := want["ai_score"].(float64); ok && score != int(v) {
|
||||
t.Errorf("%s ai_score = %d, want %d", legacy, score, int(v))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedIsIdempotent runs the seeder a second time over an already-seeded
|
||||
// database and expects every count and every id to be unchanged.
|
||||
func TestSeedIsIdempotent(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
before := map[string]int{}
|
||||
for _, table := range tableFor {
|
||||
before[table] = count(t, h, table)
|
||||
}
|
||||
var idsBefore string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT coalesce(string_agg(id::text, ',' ORDER BY id), '') FROM job_applications").
|
||||
Scan(&idsBefore); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := seeder.New(h.Pool, fx, h.Now).Run(ctx); err != nil {
|
||||
t.Fatalf("second seed: %v", err)
|
||||
}
|
||||
|
||||
for _, table := range tableFor {
|
||||
if got := count(t, h, table); got != before[table] {
|
||||
t.Errorf("%s: %d rows after re-seed, %d before — the seeder duplicated rows",
|
||||
table, got, before[table])
|
||||
}
|
||||
}
|
||||
var idsAfter string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT coalesce(string_agg(id::text, ',' ORDER BY id), '') FROM job_applications").
|
||||
Scan(&idsAfter); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if idsAfter != idsBefore {
|
||||
t.Error("application ids changed across a re-seed; keys are not deterministic")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedRelationships checks that every reference was rewritten to a real row.
|
||||
func TestSeedRelationships(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
|
||||
dangling := []struct{ name, query string }{
|
||||
{"applications without a posting",
|
||||
`SELECT count(*) FROM job_applications a
|
||||
LEFT JOIN job_postings p ON p.id = a.job_posting_id WHERE p.id IS NULL`},
|
||||
{"interviews without an application",
|
||||
`SELECT count(*) FROM ai_interviews i
|
||||
LEFT JOIN job_applications a ON a.id = i.application_id WHERE a.id IS NULL`},
|
||||
{"staff without an application",
|
||||
`SELECT count(*) FROM staff s LEFT JOIN job_applications a ON a.id = s.application_id
|
||||
WHERE s.application_id IS NOT NULL AND a.id IS NULL`},
|
||||
{"shifts without staff",
|
||||
`SELECT count(*) FROM shift_records r LEFT JOIN staff s ON s.id = r.staff_id
|
||||
WHERE r.staff_id IS NOT NULL AND s.id IS NULL`},
|
||||
{"evidence without a course",
|
||||
`SELECT count(*) FROM evidence e LEFT JOIN courses c ON c.id = e.course_id
|
||||
WHERE e.course_id IS NOT NULL AND c.id IS NULL`},
|
||||
}
|
||||
for _, d := range dangling {
|
||||
var n int
|
||||
if err := h.Pool.QueryRow(ctx, d.query).Scan(&n); err != nil {
|
||||
t.Fatalf("%s: %v", d.name, err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Errorf("%s: %d", d.name, n)
|
||||
}
|
||||
}
|
||||
|
||||
// interview_id is a soft reference on purpose: the source contains one
|
||||
// dangling value (app_devon -> int_devon), and preserving it is the point.
|
||||
var set, resolve int
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`SELECT (SELECT count(*) FROM job_applications WHERE interview_id IS NOT NULL),
|
||||
(SELECT count(*) FROM job_applications a JOIN ai_interviews i ON i.id = a.interview_id)`).
|
||||
Scan(&set, &resolve); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if set != 5 {
|
||||
t.Errorf("applications carrying interview_id = %d, want 5", set)
|
||||
}
|
||||
if resolve != 4 {
|
||||
t.Errorf("resolvable interview_id = %d, want 4 (int_devon dangles in the source)", resolve)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedOrganizationScope checks every seeded row belongs to the development
|
||||
// organization, so organization scoping has something real to filter on.
|
||||
func TestSeedOrganizationScope(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
for _, table := range tableFor {
|
||||
var wrong int
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT count(*) FROM "+table+" WHERE org_id IS DISTINCT FROM $1::uuid", h.OrgID).
|
||||
Scan(&wrong); err != nil {
|
||||
t.Fatalf("%s: %v", table, err)
|
||||
}
|
||||
if wrong != 0 {
|
||||
t.Errorf("%s: %d rows outside the development organization", table, wrong)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDeterministicUUID pins the key derivation: the same source id must always
|
||||
// produce the same key, or a re-seed would duplicate every row.
|
||||
func TestDeterministicUUID(t *testing.T) {
|
||||
a := seeder.DeterministicUUID("JobPosting:job_chef")
|
||||
b := seeder.DeterministicUUID("JobPosting:job_chef")
|
||||
if a != b {
|
||||
t.Fatalf("not deterministic: %s != %s", a, b)
|
||||
}
|
||||
if c := seeder.DeterministicUUID("JobPosting:job_security"); c == a {
|
||||
t.Fatal("distinct source ids produced the same key")
|
||||
}
|
||||
if len(a) != 36 || a[14] != '5' {
|
||||
t.Errorf("expected a v5 UUID, got %q", a)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSeedDoesNotDependOnWallClock: seeding twice with the same anchor must
|
||||
// produce identical shift records.
|
||||
func TestSeedShiftsStableForAnchor(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 15, 0, 0, 0, time.Local)
|
||||
a := seeder.BuildShifts(anchor)
|
||||
b := seeder.BuildShifts(anchor)
|
||||
if len(a) != len(b) {
|
||||
t.Fatalf("shift count differs between runs: %d vs %d", len(a), len(b))
|
||||
}
|
||||
for i := range a {
|
||||
if a[i]["id"] != b[i]["id"] || a[i]["created_date"] != b[i]["created_date"] {
|
||||
t.Fatalf("shift %d differs between runs", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
215
go-api/internal/seeder/shifts.go
Normal file
215
go-api/internal/seeder/shifts.go
Normal file
@@ -0,0 +1,215 @@
|
||||
package seeder
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"sort"
|
||||
"time"
|
||||
)
|
||||
|
||||
// A port of src/api/attendanceSeed.js.
|
||||
//
|
||||
// Ported rather than snapshotted because this is the one collection whose dates
|
||||
// are anchored to *now* rather than to a fixed calendar. `dataResolver.inPeriod`
|
||||
// windows every collection on created_date, so "attendance last week" has to
|
||||
// mean last week on the day the seeder runs. A frozen JSON snapshot would read
|
||||
// as permanently empty a fortnight later.
|
||||
//
|
||||
// Nothing here is random. The distribution is deterministic given the date the
|
||||
// seeder runs, so the same day always produces the same figures and the
|
||||
// regression tests can assert against them:
|
||||
//
|
||||
// Marco — the control: reliable, weekend event overtime
|
||||
// Marcus — attendance degrading over the last fortnight (the anomaly)
|
||||
// Antoine — present throughout, overtime climbing week on week (the trend)
|
||||
|
||||
const windowDays = 56
|
||||
|
||||
type rosterEntry struct {
|
||||
staffID, workerName, workerEmail string
|
||||
jobPostingID, role, roleCategory string
|
||||
weekdays []time.Weekday
|
||||
startHour int
|
||||
scheduledHours float64
|
||||
}
|
||||
|
||||
var roster = []rosterEntry{
|
||||
{
|
||||
staffID: "staff_marco", workerName: "Marco Rivera", workerEmail: "marco.rivera@email.com",
|
||||
jobPostingID: "job_bartender_corp", role: "Experienced Bartender – Corporate Events",
|
||||
roleCategory: "Bartender",
|
||||
// Wed–Sat: corporate events run late in the week.
|
||||
weekdays: []time.Weekday{time.Wednesday, time.Thursday, time.Friday, time.Saturday},
|
||||
startHour: 16, scheduledHours: 8,
|
||||
},
|
||||
{
|
||||
staffID: "staff_marcus", workerName: "Marcus Williams", workerEmail: "marcus.w@email.com",
|
||||
jobPostingID: "job_security", role: "Event Security Officer", roleCategory: "Security",
|
||||
// Mon–Fri: a fixed rota, which is what makes the recent absences stand
|
||||
// out rather than read as an irregular schedule.
|
||||
weekdays: []time.Weekday{time.Monday, time.Tuesday, time.Wednesday, time.Thursday, time.Friday},
|
||||
startHour: 14, scheduledHours: 8,
|
||||
},
|
||||
{
|
||||
staffID: "staff_antoine", workerName: "Chef Antoine Dubois", workerEmail: "antoine.dubois@email.com",
|
||||
jobPostingID: "job_chef", role: "Executive Chef – Catering", roleCategory: "Chef",
|
||||
// Tue–Sat: kitchen service.
|
||||
weekdays: []time.Weekday{time.Tuesday, time.Wednesday, time.Thursday, time.Friday, time.Saturday},
|
||||
startHour: 12, scheduledHours: 9,
|
||||
},
|
||||
}
|
||||
|
||||
type behaviour struct {
|
||||
status string
|
||||
minutesLate int
|
||||
overtime float64
|
||||
notes string
|
||||
}
|
||||
|
||||
// behaviourFor mirrors the BEHAVIOUR map. `i` counts back from the most recent
|
||||
// shift, so "the last fortnight" stays a range of small indices as the window
|
||||
// rolls forward.
|
||||
func behaviourFor(staffID string, i int, weekday time.Weekday) behaviour {
|
||||
switch staffID {
|
||||
case "staff_marco":
|
||||
b := behaviour{status: "present"}
|
||||
if i == 14 {
|
||||
b.status, b.minutesLate = "late", 9
|
||||
}
|
||||
// Friday and Saturday events overrun; midweek ones do not.
|
||||
if weekday == time.Friday || weekday == time.Saturday {
|
||||
b.overtime = 1
|
||||
}
|
||||
return b
|
||||
|
||||
case "staff_marcus":
|
||||
switch i {
|
||||
case 2, 7:
|
||||
return behaviour{status: "absent", notes: "Called in sick"}
|
||||
case 4:
|
||||
return behaviour{status: "no_show", notes: "No contact"}
|
||||
case 1:
|
||||
return behaviour{status: "late", minutesLate: 24}
|
||||
case 5:
|
||||
return behaviour{status: "late", minutesLate: 16}
|
||||
case 9:
|
||||
return behaviour{status: "late", minutesLate: 12}
|
||||
case 26:
|
||||
return behaviour{status: "late", minutesLate: 7}
|
||||
}
|
||||
return behaviour{status: "present"}
|
||||
|
||||
case "staff_antoine":
|
||||
weekIndex := i / 5
|
||||
busy := weekday == time.Thursday || weekday == time.Friday || weekday == time.Saturday
|
||||
b := behaviour{status: "present"}
|
||||
if busy {
|
||||
b.overtime = math.Max(0.5, round1(3.5-float64(weekIndex)*0.45))
|
||||
}
|
||||
return b
|
||||
}
|
||||
return behaviour{status: "present"}
|
||||
}
|
||||
|
||||
// round1 and round2 reproduce JavaScript's Math.round, which rounds halves away
|
||||
// from zero — the same rule as Go's math.Round.
|
||||
func round1(n float64) float64 { return math.Round(n*10) / 10 }
|
||||
func round2(n float64) float64 { return math.Round(n*100) / 100 }
|
||||
|
||||
// daysAgo is `n` days back at a given local hour.
|
||||
//
|
||||
// Local rather than UTC because a shift belongs to the day it was worked in the
|
||||
// place it was worked, and periodRange windows on local day boundaries too.
|
||||
func daysAgo(now time.Time, n, hour int) time.Time {
|
||||
d := now.AddDate(0, 0, -n)
|
||||
return time.Date(d.Year(), d.Month(), d.Day(), hour, 0, 0, 0, now.Location())
|
||||
}
|
||||
|
||||
func containsWeekday(set []time.Weekday, w time.Weekday) bool {
|
||||
for _, x := range set {
|
||||
if x == w {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// shiftOffsets is every day offset in the window on which this worker is rostered.
|
||||
func shiftOffsets(now time.Time, weekdays []time.Weekday) []int {
|
||||
var offsets []int
|
||||
for offset := 0; offset <= windowDays; offset++ {
|
||||
if containsWeekday(weekdays, daysAgo(now, offset, 0).Weekday()) {
|
||||
offsets = append(offsets, offset)
|
||||
}
|
||||
}
|
||||
return offsets
|
||||
}
|
||||
|
||||
const isoMillis = "2006-01-02T15:04:05.000Z"
|
||||
|
||||
func iso(t time.Time) string { return t.UTC().Format(isoMillis) }
|
||||
|
||||
// BuildShifts generates the shift records for a given instant, most recent first.
|
||||
func BuildShifts(now time.Time) []map[string]any {
|
||||
records := make([]map[string]any, 0, 128)
|
||||
|
||||
for _, w := range roster {
|
||||
offsets := shiftOffsets(now, w.weekdays)
|
||||
for i, offset := range offsets {
|
||||
scheduledStart := daysAgo(now, offset, w.startHour)
|
||||
weekday := scheduledStart.Weekday()
|
||||
scheduledEnd := scheduledStart.Add(time.Duration(w.scheduledHours * float64(time.Hour)))
|
||||
|
||||
b := behaviourFor(w.staffID, i, weekday)
|
||||
worked := b.status != "absent" && b.status != "no_show"
|
||||
|
||||
var actualStart, actualEnd any
|
||||
endForUpdated := scheduledEnd
|
||||
if worked {
|
||||
actualStart = iso(scheduledStart.Add(time.Duration(b.minutesLate) * time.Minute))
|
||||
ae := scheduledEnd.Add(time.Duration(b.overtime * float64(time.Hour)))
|
||||
actualEnd, endForUpdated = iso(ae), ae
|
||||
}
|
||||
|
||||
actualHours, overtimeHours, minutesLate := 0.0, 0.0, 0
|
||||
if worked {
|
||||
actualHours = round2(w.scheduledHours - float64(b.minutesLate)/60 + b.overtime)
|
||||
overtimeHours = round1(b.overtime)
|
||||
minutesLate = b.minutesLate
|
||||
}
|
||||
|
||||
records = append(records, map[string]any{
|
||||
"id": fmt.Sprintf("shift_%s_%02d", w.staffID[len("staff_"):], len(offsets)-i),
|
||||
"staff_id": w.staffID,
|
||||
"worker_name": w.workerName,
|
||||
"worker_email": w.workerEmail,
|
||||
"job_posting_id": w.jobPostingID,
|
||||
"role": w.role,
|
||||
"role_category": w.roleCategory,
|
||||
"shift_date": scheduledStart.Format("2006-01-02"),
|
||||
"scheduled_start": iso(scheduledStart),
|
||||
"scheduled_end": iso(scheduledEnd),
|
||||
"scheduled_hours": w.scheduledHours,
|
||||
"actual_start": actualStart,
|
||||
"actual_end": actualEnd,
|
||||
"actual_hours": actualHours,
|
||||
"overtime_hours": overtimeHours,
|
||||
"minutes_late": minutesLate,
|
||||
"status": b.status,
|
||||
"notes": b.notes,
|
||||
// Load-bearing: dataResolver windows every collection on
|
||||
// created_date, so a shift's created date IS the instant it
|
||||
// was worked.
|
||||
"created_date": iso(scheduledStart),
|
||||
"updated_date": iso(endForUpdated),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Most recent first, matching the -created_date order every other
|
||||
// collection is listed in.
|
||||
sort.SliceStable(records, func(a, b int) bool {
|
||||
return records[a]["created_date"].(string) > records[b]["created_date"].(string)
|
||||
})
|
||||
return records
|
||||
}
|
||||
178
go-api/internal/seeder/shifts_convergence_test.go
Normal file
178
go-api/internal/seeder/shifts_convergence_test.go
Normal file
@@ -0,0 +1,178 @@
|
||||
package seeder_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/seeder"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The shift collection is a rolling 56-day window, and its stable id encodes a
|
||||
// shift's POSITION in that window rather than its identity. Seed on Friday and
|
||||
// re-seed on Saturday and every id means a different date; one of them —
|
||||
// Marcus's, because he works Mon–Fri and a Saturday window holds one fewer of
|
||||
// his shifts — is no longer generated at all.
|
||||
//
|
||||
// Upsert cannot express that. These tests pin the behaviour that can: after any
|
||||
// seed, shift_records holds exactly what BuildShifts produced for that instant,
|
||||
// and nothing left over from a previous run.
|
||||
|
||||
func shiftCount(t *testing.T, h *testutil.Harness) int {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := h.Pool.QueryRow(context.Background(),
|
||||
"SELECT count(*) FROM shift_records").Scan(&n); err != nil {
|
||||
t.Fatalf("count shift_records: %v", err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// TestReseedOnALaterDayLeavesNoStaleShifts is the regression itself: the
|
||||
// database was seeded on one day, the frontend regenerates on the next, and the
|
||||
// two must still describe the same collection.
|
||||
func TestReseedOnALaterDayLeavesNoStaleShifts(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// A Friday, then the Saturday after it — the exact pair that orphaned
|
||||
// shift_marcus_41 in the live database.
|
||||
friday := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
saturday := friday.AddDate(0, 0, 1)
|
||||
|
||||
if _, err := seeder.New(h.Pool, fx, friday).Run(ctx); err != nil {
|
||||
t.Fatalf("seed on the Friday: %v", err)
|
||||
}
|
||||
fridayRows := shiftCount(t, h)
|
||||
if want := len(seeder.BuildShifts(friday)); fridayRows != want {
|
||||
t.Fatalf("after the Friday seed: %d rows, generator produced %d", fridayRows, want)
|
||||
}
|
||||
|
||||
result, err := seeder.New(h.Pool, fx, saturday).Run(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("re-seed on the Saturday: %v", err)
|
||||
}
|
||||
|
||||
generated := seeder.BuildShifts(saturday)
|
||||
if got := shiftCount(t, h); got != len(generated) {
|
||||
t.Errorf("after re-seeding a day later: %d rows, but the generator produced %d "+
|
||||
"— %d stale record(s) survived the re-seed", got, len(generated), got-len(generated))
|
||||
}
|
||||
if result.Pruned != fridayRows-len(generated) {
|
||||
t.Errorf("Pruned = %d, want %d", result.Pruned, fridayRows-len(generated))
|
||||
}
|
||||
|
||||
// Every surviving row must be one this run generated, holding this run's
|
||||
// date for that id — not the previous run's.
|
||||
want := map[string]string{}
|
||||
for _, rec := range generated {
|
||||
want[rec["id"].(string)] = rec["shift_date"].(string)
|
||||
}
|
||||
rows, err := h.Pool.Query(ctx, "SELECT legacy_id, shift_date::text FROM shift_records")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var legacy, date string
|
||||
if err := rows.Scan(&legacy, &date); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
switch expected, generatedNow := want[legacy]; {
|
||||
case !generatedNow:
|
||||
t.Errorf("%s is in the database but was not generated for this instant", legacy)
|
||||
case expected != date:
|
||||
t.Errorf("%s holds %s, but this run generated it as %s", legacy, date, expected)
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReseedIsConvergentAcrossAWeek walks a whole week, so the assertion does
|
||||
// not depend on the one day pair that happened to expose the bug. Every day of
|
||||
// the week changes the roster composition differently.
|
||||
func TestReseedIsConvergentAcrossAWeek(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
day := time.Date(2026, 8, 17, 9, 0, 0, 0, time.Local) // a Monday
|
||||
for i := 0; i < 7; i++ {
|
||||
now := day.AddDate(0, 0, i)
|
||||
if _, err := seeder.New(h.Pool, fx, now).Run(ctx); err != nil {
|
||||
t.Fatalf("seed on %s: %v", now.Weekday(), err)
|
||||
}
|
||||
want := len(seeder.BuildShifts(now))
|
||||
if got := shiftCount(t, h); got != want {
|
||||
t.Errorf("%s %s: %d rows, generator produced %d",
|
||||
now.Weekday(), now.Format("2006-01-02"), got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Re-seeding the same instant twice must still change nothing — the prune must
|
||||
// not delete rows it just wrote.
|
||||
func TestReseedSameInstantPrunesNothing(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
before := shiftCount(t, h)
|
||||
result, err := seeder.New(h.Pool, fx, h.Now).Run(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("re-seed: %v", err)
|
||||
}
|
||||
if result.Pruned != 0 {
|
||||
t.Errorf("re-seeding the same instant pruned %d record(s), want 0", result.Pruned)
|
||||
}
|
||||
if got := shiftCount(t, h); got != before {
|
||||
t.Errorf("shift_records went from %d to %d rows on an identical re-seed", before, got)
|
||||
}
|
||||
}
|
||||
|
||||
// The prune is scoped to the organization being seeded. Another organization's
|
||||
// shift records are none of its business — and once authentication lands, that
|
||||
// is the difference between a re-seed and an incident.
|
||||
func TestPruneIsScopedToTheSeededOrganization(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
fx := testutil.Fixture(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var otherOrg string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('Other', 'other-org') RETURNING id::text`).
|
||||
Scan(&otherOrg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Copy one of this organization's shifts into the other one, with an id and
|
||||
// legacy_id no generation will ever produce.
|
||||
if _, err := h.Pool.Exec(ctx,
|
||||
`INSERT INTO shift_records (org_id, legacy_id, staff_id, worker_name, worker_email,
|
||||
job_posting_id, role, role_category, shift_date, scheduled_start, scheduled_end,
|
||||
scheduled_hours, actual_start, actual_end, actual_hours, overtime_hours,
|
||||
minutes_late, status, notes, created_date, updated_date)
|
||||
SELECT $1::uuid, 'shift_other_99', staff_id, worker_name, worker_email,
|
||||
job_posting_id, role, role_category, shift_date, scheduled_start, scheduled_end,
|
||||
scheduled_hours, actual_start, actual_end, actual_hours, overtime_hours,
|
||||
minutes_late, status, notes, created_date, updated_date
|
||||
FROM shift_records LIMIT 1`, otherOrg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if _, err := seeder.New(h.Pool, fx, h.Now.AddDate(0, 0, 1)).Run(ctx); err != nil {
|
||||
t.Fatalf("re-seed: %v", err)
|
||||
}
|
||||
|
||||
var survived int
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
"SELECT count(*) FROM shift_records WHERE org_id = $1::uuid", otherOrg).Scan(&survived); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if survived != 1 {
|
||||
t.Errorf("the other organization's shift record was pruned: %d survived, want 1", survived)
|
||||
}
|
||||
}
|
||||
134
go-api/internal/seeder/shifts_test.go
Normal file
134
go-api/internal/seeder/shifts_test.go
Normal file
@@ -0,0 +1,134 @@
|
||||
package seeder_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/seeder"
|
||||
)
|
||||
|
||||
// The shift generator is a port of src/api/attendanceSeed.js. These tests pin
|
||||
// the distribution that module's own documentation describes, so a drift in the
|
||||
// port shows up as a failing assertion rather than as quietly different
|
||||
// attendance figures.
|
||||
|
||||
func shiftsFor(anchor time.Time, email string) []map[string]any {
|
||||
var out []map[string]any
|
||||
for _, r := range seeder.BuildShifts(anchor) {
|
||||
if r["worker_email"] == email {
|
||||
out = append(out, r)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestShiftDistributionMatchesSource(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
all := seeder.BuildShifts(anchor)
|
||||
|
||||
counts := map[string]int{}
|
||||
for _, r := range all {
|
||||
counts[r["status"].(string)]++
|
||||
}
|
||||
|
||||
// Marcus alone supplies the attendance anomaly: two absences, one no-show
|
||||
// and three late arrivals in the recent window, plus one older late.
|
||||
if counts["absent"] != 2 {
|
||||
t.Errorf("absent = %d, want 2", counts["absent"])
|
||||
}
|
||||
if counts["no_show"] != 1 {
|
||||
t.Errorf("no_show = %d, want 1", counts["no_show"])
|
||||
}
|
||||
// Marcus i=1,5,9,26 plus Marco i=14.
|
||||
if counts["late"] != 5 {
|
||||
t.Errorf("late = %d, want 5", counts["late"])
|
||||
}
|
||||
if counts["present"] == 0 {
|
||||
t.Error("no present shifts generated")
|
||||
}
|
||||
}
|
||||
|
||||
func TestShiftRosterIsThreePeople(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
emails := map[string]bool{}
|
||||
for _, r := range seeder.BuildShifts(anchor) {
|
||||
emails[r["worker_email"].(string)] = true
|
||||
}
|
||||
if len(emails) != 3 {
|
||||
t.Fatalf("roster has %d people, want 3 (one per hire)", len(emails))
|
||||
}
|
||||
}
|
||||
|
||||
// A missed shift is zero hours worked, not a short one — and the schema's
|
||||
// shift_records_absence_has_no_hours constraint depends on it.
|
||||
func TestAbsencesHaveNoHours(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
for _, r := range seeder.BuildShifts(anchor) {
|
||||
status := r["status"].(string)
|
||||
if status != "absent" && status != "no_show" {
|
||||
continue
|
||||
}
|
||||
if r["actual_hours"].(float64) != 0 {
|
||||
t.Errorf("%s: %s shift has actual_hours %v", r["id"], status, r["actual_hours"])
|
||||
}
|
||||
if r["minutes_late"].(int) != 0 {
|
||||
t.Errorf("%s: %s shift has minutes_late %v", r["id"], status, r["minutes_late"])
|
||||
}
|
||||
if r["actual_start"] != nil || r["actual_end"] != nil {
|
||||
t.Errorf("%s: %s shift has actual timestamps", r["id"], status)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Antoine's overtime climbs week on week — a trend rather than a spike. It is
|
||||
// the overtime anomaly the analytics are shaped to surface.
|
||||
func TestAntoineOvertimeClimbs(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
shifts := shiftsFor(anchor, "antoine.dubois@email.com")
|
||||
if len(shifts) == 0 {
|
||||
t.Fatal("no shifts generated for Antoine")
|
||||
}
|
||||
|
||||
// BuildShifts returns most-recent-first, so recent overtime should exceed
|
||||
// the overtime from the far end of the window.
|
||||
var recent, older float64
|
||||
for i, r := range shifts {
|
||||
ot := r["overtime_hours"].(float64)
|
||||
if i < 10 {
|
||||
recent += ot
|
||||
}
|
||||
if i >= len(shifts)-10 {
|
||||
older += ot
|
||||
}
|
||||
}
|
||||
if recent <= older {
|
||||
t.Errorf("overtime is not climbing: recent 10 = %.1fh, oldest 10 = %.1fh", recent, older)
|
||||
}
|
||||
}
|
||||
|
||||
// created_date is the instant the shift was worked. dataResolver windows every
|
||||
// collection on it, so a shift dated anywhere else would vanish from every
|
||||
// period reading.
|
||||
func TestShiftCreatedDateIsTheShiftInstant(t *testing.T) {
|
||||
anchor := time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local)
|
||||
for _, r := range seeder.BuildShifts(anchor) {
|
||||
if r["created_date"] != r["scheduled_start"] {
|
||||
t.Fatalf("%s: created_date %v is not the scheduled start %v",
|
||||
r["id"], r["created_date"], r["scheduled_start"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The window rolls forward with the anchor: shifts must stay recent relative to
|
||||
// whenever the seeder runs, which is the whole reason this is a port rather
|
||||
// than a frozen snapshot.
|
||||
func TestShiftWindowFollowsTheAnchor(t *testing.T) {
|
||||
early := seeder.BuildShifts(time.Date(2026, 3, 1, 12, 0, 0, 0, time.Local))
|
||||
late := seeder.BuildShifts(time.Date(2026, 8, 21, 12, 0, 0, 0, time.Local))
|
||||
if early[0]["created_date"] == late[0]["created_date"] {
|
||||
t.Fatal("shift dates did not move with the anchor")
|
||||
}
|
||||
if got := late[0]["created_date"].(string)[:4]; got != "2026" {
|
||||
t.Errorf("most recent shift is dated %s", got)
|
||||
}
|
||||
}
|
||||
451
go-api/internal/service/definitions.go
Normal file
451
go-api/internal/service/definitions.go
Normal file
@@ -0,0 +1,451 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/definition"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
// allowedDefinitionFilters names the accepted query parameters for definition collections.
|
||||
var allowedDefinitionFilters = map[string]bool{
|
||||
"visibility": true,
|
||||
"status": true,
|
||||
"definition_id": true,
|
||||
"sort": true,
|
||||
"limit": true,
|
||||
"offset": true,
|
||||
}
|
||||
|
||||
// DefinitionsService manages authored Agent and Skill definitions.
|
||||
type DefinitionsService struct {
|
||||
repo *repo.DefinitionsRepo
|
||||
}
|
||||
|
||||
// NewDefinitions builds a definitions service over a repository.
|
||||
func NewDefinitions(db repo.Querier) *DefinitionsService {
|
||||
return &DefinitionsService{repo: repo.NewDefinitionsRepo(db)}
|
||||
}
|
||||
|
||||
// ParseListParams validates query parameters for listing definitions.
|
||||
func (s *DefinitionsService) ParseListParams(q url.Values) (repo.DefinitionListParams, error) {
|
||||
p := repo.DefinitionListParams{
|
||||
Limit: 100,
|
||||
Sort: "created_date",
|
||||
Desc: true,
|
||||
}
|
||||
|
||||
for name := range q {
|
||||
if !allowedDefinitionFilters[name] {
|
||||
return p, domain.Invalid(fmt.Sprintf("unknown filter field %q", name))
|
||||
}
|
||||
}
|
||||
|
||||
if raw := q.Get("visibility"); raw != "" {
|
||||
if raw != "personal" && raw != "organization" {
|
||||
return p, domain.Invalid("visibility must be one of: personal, organization")
|
||||
}
|
||||
p.Visibility = raw
|
||||
}
|
||||
|
||||
if raw := q.Get("status"); raw != "" {
|
||||
p.Status = raw
|
||||
}
|
||||
|
||||
if raw := q.Get("definition_id"); raw != "" {
|
||||
p.DefinitionID = raw
|
||||
}
|
||||
|
||||
if raw := q.Get("sort"); raw != "" {
|
||||
field := raw
|
||||
desc := false
|
||||
if strings.HasPrefix(field, "-") {
|
||||
desc = true
|
||||
field = field[1:]
|
||||
}
|
||||
switch field {
|
||||
case "created_date", "updated_date", "name", "definition_id", "status", "version":
|
||||
p.Sort = field
|
||||
p.Desc = desc
|
||||
default:
|
||||
return p, domain.Invalid(fmt.Sprintf("cannot sort by %q", field))
|
||||
}
|
||||
}
|
||||
|
||||
if raw := q.Get("limit"); raw != "" {
|
||||
n, err := strconv.Atoi(raw)
|
||||
if err != nil || n < 0 {
|
||||
return p, domain.Invalid("limit must be a non-negative integer")
|
||||
}
|
||||
if n > MaxLimit {
|
||||
n = MaxLimit
|
||||
}
|
||||
p.Limit = n
|
||||
}
|
||||
|
||||
if raw := q.Get("offset"); raw != "" {
|
||||
n, err := strconv.Atoi(raw)
|
||||
if err != nil || n < 0 {
|
||||
return p, domain.Invalid("offset must be a non-negative integer")
|
||||
}
|
||||
p.Offset = n
|
||||
}
|
||||
|
||||
return p, nil
|
||||
}
|
||||
|
||||
/* ── Agents ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// ListAgents returns a page of authored agent definitions.
|
||||
func (s *DefinitionsService) ListAgents(ctx context.Context, ident authctx.Identity, p repo.DefinitionListParams) (*domain.Page, error) {
|
||||
records, total, err := s.repo.ListAgents(ctx, ident, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if records == nil {
|
||||
records = []domain.Record{}
|
||||
}
|
||||
return &domain.Page{
|
||||
Records: records,
|
||||
Total: total,
|
||||
Limit: p.Limit,
|
||||
Offset: p.Offset,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetAgent returns one agent definition by id within the caller's tenant and ownership scope.
|
||||
func (s *DefinitionsService) GetAgent(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return nil, domain.NotFound("AgentDefinition", id)
|
||||
}
|
||||
rec, err := s.repo.GetAgent(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, domain.NotFound("AgentDefinition", id)
|
||||
}
|
||||
return rec, nil
|
||||
}
|
||||
|
||||
// CreateAgent validates, parses and persists a new authored agent definition.
|
||||
func (s *DefinitionsService) CreateAgent(ctx context.Context, ident authctx.Identity, body domain.Record) (domain.Record, error) {
|
||||
mdRaw, ok := body["markdown"]
|
||||
if !ok || mdRaw == nil {
|
||||
return nil, domain.Validation("Paste or upload a Markdown definition.", nil)
|
||||
}
|
||||
markdown, isStr := mdRaw.(string)
|
||||
if !isStr {
|
||||
return nil, domain.Validation("markdown must be a string", nil)
|
||||
}
|
||||
|
||||
if err := definition.ValidateAgent(markdown); err != nil {
|
||||
return nil, domain.Validation(err.Error(), nil)
|
||||
}
|
||||
|
||||
visibility := "personal"
|
||||
if visRaw, ok := body["visibility"]; ok && visRaw != nil {
|
||||
v, isStr := visRaw.(string)
|
||||
if !isStr || (v != "personal" && v != "organization") {
|
||||
return nil, domain.Validation("visibility must be one of: personal, organization", map[string]string{"visibility": "invalid"})
|
||||
}
|
||||
visibility = v
|
||||
}
|
||||
|
||||
if visibility == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
agent, err := definition.ParseAgent(markdown, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, domain.Validation("That definition could not be parsed. "+err.Error(), nil)
|
||||
}
|
||||
|
||||
input := repo.AgentInsertInput{
|
||||
DefinitionID: agent.ID,
|
||||
OrgID: ident.OrgID,
|
||||
Visibility: visibility,
|
||||
CreatedBy: &ident.UserID,
|
||||
Markdown: markdown,
|
||||
Status: agent.Status,
|
||||
Version: agent.Version,
|
||||
Name: agent.Name,
|
||||
Description: agent.Description,
|
||||
Pages: agent.Pages,
|
||||
}
|
||||
|
||||
if visibility == "personal" {
|
||||
input.OwnerUserID = &ident.UserID
|
||||
}
|
||||
|
||||
return s.repo.InsertAgent(ctx, ident, input)
|
||||
}
|
||||
|
||||
// UpdateAgent validates and applies updates to an authored agent definition.
|
||||
func (s *DefinitionsService) UpdateAgent(ctx context.Context, ident authctx.Identity, id string, patch domain.Record) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return nil, domain.NotFound("AgentDefinition", id)
|
||||
}
|
||||
|
||||
existing, err := s.repo.GetAgent(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
return nil, domain.NotFound("AgentDefinition", id)
|
||||
}
|
||||
|
||||
if existing["visibility"] == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
if visRaw, ok := patch["visibility"]; ok && visRaw != nil {
|
||||
if v, isStr := visRaw.(string); isStr && v != existing["visibility"] {
|
||||
return nil, domain.Validation("visibility cannot be modified after creation", map[string]string{"visibility": "immutable"})
|
||||
}
|
||||
}
|
||||
|
||||
var input repo.AgentUpdateInput
|
||||
|
||||
if mdRaw, ok := patch["markdown"]; ok && mdRaw != nil {
|
||||
markdown, isStr := mdRaw.(string)
|
||||
if !isStr {
|
||||
return nil, domain.Validation("markdown must be a string", nil)
|
||||
}
|
||||
if err := definition.ValidateAgent(markdown); err != nil {
|
||||
return nil, domain.Validation(err.Error(), nil)
|
||||
}
|
||||
agent, err := definition.ParseAgent(markdown, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, domain.Validation("That definition could not be parsed. "+err.Error(), nil)
|
||||
}
|
||||
input.Markdown = &markdown
|
||||
input.DefinitionID = &agent.ID
|
||||
input.Name = &agent.Name
|
||||
input.Description = &agent.Description
|
||||
input.Status = &agent.Status
|
||||
input.Version = &agent.Version
|
||||
input.Pages = agent.Pages
|
||||
} else if statusRaw, ok := patch["status"]; ok && statusRaw != nil {
|
||||
status, isStr := statusRaw.(string)
|
||||
if !isStr || (status != "draft" && status != "published" && status != "archived") {
|
||||
return nil, domain.Validation("status must be one of: draft, published, archived", map[string]string{"status": "invalid"})
|
||||
}
|
||||
input.Status = &status
|
||||
}
|
||||
|
||||
return s.repo.UpdateAgent(ctx, ident, id, input)
|
||||
}
|
||||
|
||||
// DeleteAgent removes an agent definition following idempotent delete semantics.
|
||||
func (s *DefinitionsService) DeleteAgent(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
existing, err := s.repo.GetAgent(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
if existing["visibility"] == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := s.repo.DeleteAgent(ctx, ident, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
/* ── Skills ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// ListSkills returns a page of authored skill definitions.
|
||||
func (s *DefinitionsService) ListSkills(ctx context.Context, ident authctx.Identity, p repo.DefinitionListParams) (*domain.Page, error) {
|
||||
records, total, err := s.repo.ListSkills(ctx, ident, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if records == nil {
|
||||
records = []domain.Record{}
|
||||
}
|
||||
return &domain.Page{
|
||||
Records: records,
|
||||
Total: total,
|
||||
Limit: p.Limit,
|
||||
Offset: p.Offset,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetSkill returns one skill definition by id within the caller's tenant and ownership scope.
|
||||
func (s *DefinitionsService) GetSkill(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return nil, domain.NotFound("SkillDefinition", id)
|
||||
}
|
||||
rec, err := s.repo.GetSkill(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, domain.NotFound("SkillDefinition", id)
|
||||
}
|
||||
return rec, nil
|
||||
}
|
||||
|
||||
// CreateSkill validates, parses and persists a new authored skill definition.
|
||||
func (s *DefinitionsService) CreateSkill(ctx context.Context, ident authctx.Identity, body domain.Record) (domain.Record, error) {
|
||||
mdRaw, ok := body["markdown"]
|
||||
if !ok || mdRaw == nil {
|
||||
return nil, domain.Validation("Paste or upload a Markdown definition.", nil)
|
||||
}
|
||||
markdown, isStr := mdRaw.(string)
|
||||
if !isStr {
|
||||
return nil, domain.Validation("markdown must be a string", nil)
|
||||
}
|
||||
|
||||
if err := definition.ValidateSkill(markdown); err != nil {
|
||||
return nil, domain.Validation(err.Error(), nil)
|
||||
}
|
||||
|
||||
visibility := "personal"
|
||||
if visRaw, ok := body["visibility"]; ok && visRaw != nil {
|
||||
v, isStr := visRaw.(string)
|
||||
if !isStr || (v != "personal" && v != "organization") {
|
||||
return nil, domain.Validation("visibility must be one of: personal, organization", map[string]string{"visibility": "invalid"})
|
||||
}
|
||||
visibility = v
|
||||
}
|
||||
|
||||
if visibility == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
skill, err := definition.ParseSkill(markdown, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, domain.Validation("That definition could not be parsed. "+err.Error(), nil)
|
||||
}
|
||||
|
||||
input := repo.SkillInsertInput{
|
||||
DefinitionID: skill.ID,
|
||||
OrgID: ident.OrgID,
|
||||
Visibility: visibility,
|
||||
CreatedBy: &ident.UserID,
|
||||
Markdown: markdown,
|
||||
Status: skill.Status,
|
||||
Name: skill.Name,
|
||||
Description: skill.Description,
|
||||
Pages: skill.Pages,
|
||||
}
|
||||
|
||||
if visibility == "personal" {
|
||||
input.OwnerUserID = &ident.UserID
|
||||
}
|
||||
|
||||
return s.repo.InsertSkill(ctx, ident, input)
|
||||
}
|
||||
|
||||
// UpdateSkill validates and applies updates to an authored skill definition.
|
||||
func (s *DefinitionsService) UpdateSkill(ctx context.Context, ident authctx.Identity, id string, patch domain.Record) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return nil, domain.NotFound("SkillDefinition", id)
|
||||
}
|
||||
|
||||
existing, err := s.repo.GetSkill(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
return nil, domain.NotFound("SkillDefinition", id)
|
||||
}
|
||||
|
||||
if existing["visibility"] == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
if visRaw, ok := patch["visibility"]; ok && visRaw != nil {
|
||||
if v, isStr := visRaw.(string); isStr && v != existing["visibility"] {
|
||||
return nil, domain.Validation("visibility cannot be modified after creation", map[string]string{"visibility": "immutable"})
|
||||
}
|
||||
}
|
||||
|
||||
var input repo.SkillUpdateInput
|
||||
|
||||
if mdRaw, ok := patch["markdown"]; ok && mdRaw != nil {
|
||||
markdown, isStr := mdRaw.(string)
|
||||
if !isStr {
|
||||
return nil, domain.Validation("markdown must be a string", nil)
|
||||
}
|
||||
if err := definition.ValidateSkill(markdown); err != nil {
|
||||
return nil, domain.Validation(err.Error(), nil)
|
||||
}
|
||||
skill, err := definition.ParseSkill(markdown, definition.Options{})
|
||||
if err != nil {
|
||||
return nil, domain.Validation("That definition could not be parsed. "+err.Error(), nil)
|
||||
}
|
||||
input.Markdown = &markdown
|
||||
input.DefinitionID = &skill.ID
|
||||
input.Name = &skill.Name
|
||||
input.Description = &skill.Description
|
||||
input.Status = &skill.Status
|
||||
input.Pages = skill.Pages
|
||||
} else if statusRaw, ok := patch["status"]; ok && statusRaw != nil {
|
||||
status, isStr := statusRaw.(string)
|
||||
if !isStr || (status != "active" && status != "inactive") {
|
||||
return nil, domain.Validation("status must be one of: active, inactive", map[string]string{"status": "invalid"})
|
||||
}
|
||||
input.Status = &status
|
||||
}
|
||||
|
||||
return s.repo.UpdateSkill(ctx, ident, id, input)
|
||||
}
|
||||
|
||||
// DeleteSkill removes a skill definition following idempotent delete semantics.
|
||||
func (s *DefinitionsService) DeleteSkill(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
existing, err := s.repo.GetSkill(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
if existing["visibility"] == "organization" {
|
||||
role, known := domain.ParseRole(ident.Role)
|
||||
if !known || role == domain.RoleTalent {
|
||||
return nil, domain.Forbidden()
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := s.repo.DeleteSkill(ctx, ident, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
335
go-api/internal/service/service.go
Normal file
335
go-api/internal/service/service.go
Normal file
@@ -0,0 +1,335 @@
|
||||
// Package service sits between the HTTP layer and the repositories.
|
||||
//
|
||||
// It owns request validation, organization scoping, and the three behaviours
|
||||
// the contract is most specific about: what a missing record does on read
|
||||
// (§5.1), what a missing record does on delete (§12.7), and what a PATCH is
|
||||
// allowed to touch (§3.2).
|
||||
//
|
||||
// No business rule lives here that the frontend does not already impose. The
|
||||
// scoring, funnel and matching logic all stay client-side in Phase 2C.
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
// MaxLimit caps how much a single request can ask for. Nothing in the frontend
|
||||
// asks for more than 500; this exists so a hand-written query cannot ask for
|
||||
// everything. See api-contract.md §8.
|
||||
const MaxLimit = 1000
|
||||
|
||||
// Service serves one resource.
|
||||
type Service struct {
|
||||
res *domain.Resource
|
||||
db repo.Querier
|
||||
}
|
||||
|
||||
// New builds a service for a resource.
|
||||
func New(res *domain.Resource, db repo.Querier) *Service {
|
||||
return &Service{res: res, db: db}
|
||||
}
|
||||
|
||||
// Resource is the descriptor this service serves.
|
||||
func (s *Service) Resource() *domain.Resource { return s.res }
|
||||
|
||||
func (s *Service) repo() *repo.Repo { return repo.New(s.res, s.db) }
|
||||
|
||||
/* ── Query parsing ──────────────────────────────────────────────────────── */
|
||||
|
||||
// Reserved query parameters. Every other parameter is a field filter.
|
||||
// No column in any resource collides with these. See api-contract.md §1.
|
||||
var reserved = map[string]bool{"sort": true, "limit": true, "offset": true}
|
||||
|
||||
// ParseList turns a query string into validated list parameters, applying this
|
||||
// resource's own defaults. The defaults are not generic: each one is the
|
||||
// literal argument at the frontend call site (api-contract.md §8.1).
|
||||
func (s *Service) ParseList(q url.Values) (domain.ListParams, error) {
|
||||
p := domain.ListParams{Limit: s.res.DefaultLimit}
|
||||
|
||||
sortSpec := s.res.DefaultSort
|
||||
if raw, ok := q["sort"]; ok && len(raw) > 0 {
|
||||
sortSpec = raw[0] // an explicitly empty ?sort= means "no ordering"
|
||||
}
|
||||
if sortSpec != "" {
|
||||
field := sortSpec
|
||||
if strings.HasPrefix(field, "-") {
|
||||
p.Desc, field = true, field[1:]
|
||||
}
|
||||
if !s.res.Sortable(field) {
|
||||
return p, domain.Invalid(fmt.Sprintf("cannot sort by %q on %s", field, s.res.Name))
|
||||
}
|
||||
p.Sort = field
|
||||
}
|
||||
|
||||
if raw := q.Get("limit"); raw != "" {
|
||||
n, err := strconv.Atoi(raw)
|
||||
if err != nil || n < 0 {
|
||||
return p, domain.Invalid("limit must be a non-negative integer")
|
||||
}
|
||||
if n > MaxLimit {
|
||||
n = MaxLimit
|
||||
}
|
||||
p.Limit = n
|
||||
}
|
||||
if raw := q.Get("offset"); raw != "" {
|
||||
n, err := strconv.Atoi(raw)
|
||||
if err != nil || n < 0 {
|
||||
return p, domain.Invalid("offset must be a non-negative integer")
|
||||
}
|
||||
p.Offset = n
|
||||
}
|
||||
|
||||
// Deterministic filter order keeps generated SQL stable and cacheable.
|
||||
names := make([]string, 0, len(q))
|
||||
for name := range q {
|
||||
if !reserved[name] {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
for _, name := range names {
|
||||
col, ok := s.res.Column(name)
|
||||
if !ok {
|
||||
return p, domain.Invalid(fmt.Sprintf("unknown filter field %q on %s", name, s.res.Name))
|
||||
}
|
||||
if !s.res.Filterable(name) {
|
||||
return p, domain.Invalid(fmt.Sprintf(
|
||||
"%s is not filterable: array and JSON columns cannot be compared for equality", name))
|
||||
}
|
||||
values := q[name]
|
||||
if len(values) == 0 {
|
||||
continue
|
||||
}
|
||||
p.Filters = append(p.Filters, domain.Filter{Column: col, Values: values})
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
/* ── Reads ──────────────────────────────────────────────────────────────── */
|
||||
|
||||
// List returns a page. An empty result is a page with no records, never an error.
|
||||
func (s *Service) List(ctx context.Context, ident authctx.Identity, p domain.ListParams) (*domain.Page, error) {
|
||||
page, err := s.repo().List(ctx, ident, p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if page.Records == nil {
|
||||
page.Records = []domain.Record{}
|
||||
}
|
||||
return page, nil
|
||||
}
|
||||
|
||||
// Get returns one record, or a not_found error carrying store.js's message.
|
||||
func (s *Service) Get(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
// store.js throws "<Entity> <id> not found" for any id it cannot find,
|
||||
// and a malformed id is simply an id it cannot find.
|
||||
return nil, domain.NotFound(s.res.Name, id)
|
||||
}
|
||||
rec, err := s.repo().Get(ctx, ident, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, domain.NotFound(s.res.Name, id)
|
||||
}
|
||||
return rec, nil
|
||||
}
|
||||
|
||||
/* ── Writes ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Create validates and inserts, returning the complete stored record.
|
||||
func (s *Service) Create(ctx context.Context, ident authctx.Identity, body domain.Record) (domain.Record, error) {
|
||||
clean, err := s.validate(body, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.repo().Insert(ctx, ident, clean)
|
||||
}
|
||||
|
||||
// Update shallow-merges the supplied fields. Absent keys are left untouched.
|
||||
func (s *Service) Update(ctx context.Context, ident authctx.Identity, id string, patch domain.Record) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return nil, domain.NotFound(s.res.Name, id)
|
||||
}
|
||||
clean, err := s.validate(patch, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rec, err := s.repo().Update(ctx, ident, id, clean)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil, domain.NotFound(s.res.Name, id)
|
||||
}
|
||||
return rec, nil
|
||||
}
|
||||
|
||||
// Delete removes a record and always reports success.
|
||||
//
|
||||
// store.js filters its array and returns { id } whether or not anything
|
||||
// matched, and both live callers delete inside loops without checking. A 404
|
||||
// here would surface an error toast where none appears today.
|
||||
// See api-contract.md §12.7.
|
||||
func (s *Service) Delete(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
|
||||
if !isUUID(id) {
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
if _, err := s.repo().Delete(ctx, ident, id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return domain.Record{"id": id}, nil
|
||||
}
|
||||
|
||||
/* ── Validation ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// validate checks a request body against the resource's columns and returns a
|
||||
// copy with server-owned fields removed.
|
||||
//
|
||||
// Unknown fields are rejected rather than ignored. Silently dropping them is
|
||||
// exactly how `interview_id`, `training_outline` and `score_breakdown` would
|
||||
// have been lost: the frontend would have written them, the API would have
|
||||
// accepted the request, and the data would never have arrived.
|
||||
func (s *Service) validate(in domain.Record, isCreate bool) (domain.Record, error) {
|
||||
details := map[string]string{}
|
||||
out := make(domain.Record, len(in))
|
||||
|
||||
for name, value := range in {
|
||||
col, ok := s.res.Column(name)
|
||||
if !ok {
|
||||
details[name] = "unknown field"
|
||||
continue
|
||||
}
|
||||
if col.ReadOnly {
|
||||
continue // server-owned: ignored, not rejected (api-contract.md §3.1)
|
||||
}
|
||||
if value == nil {
|
||||
if col.NotNull {
|
||||
details[name] = "must not be null"
|
||||
continue
|
||||
}
|
||||
out[name] = nil
|
||||
continue
|
||||
}
|
||||
if col.Kind == domain.KindEnum {
|
||||
str, isStr := value.(string)
|
||||
if !isStr || !contains(col.Enum, str) {
|
||||
details[name] = fmt.Sprintf("must be one of: %s", strings.Join(col.Enum, ", "))
|
||||
continue
|
||||
}
|
||||
}
|
||||
out[name] = value
|
||||
}
|
||||
|
||||
if isCreate {
|
||||
for _, col := range s.res.Columns {
|
||||
if !col.Required {
|
||||
continue
|
||||
}
|
||||
if s.serverSupplies(col.Name) {
|
||||
// The repository fills this from the session, so demanding it
|
||||
// from the caller would reject a request the server is about to
|
||||
// complete correctly. evidence.worker_email is the live case.
|
||||
continue
|
||||
}
|
||||
v, ok := out[col.Name]
|
||||
if !ok {
|
||||
details[col.Name] = "required"
|
||||
continue
|
||||
}
|
||||
if str, isStr := v.(string); isStr && strings.TrimSpace(str) == "" {
|
||||
details[col.Name] = "must not be blank"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(details) > 0 {
|
||||
return nil, domain.Validation(
|
||||
fmt.Sprintf("%s payload is not valid", s.res.Name), details)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// serverSupplies reports whether a column is filled in from the authenticated
|
||||
// session rather than from the request body.
|
||||
func (s *Service) serverSupplies(name string) bool {
|
||||
if s.res.Policy == nil {
|
||||
return false
|
||||
}
|
||||
for _, d := range s.res.Policy.Derived {
|
||||
if d.Column == name {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func contains(set []string, v string) bool {
|
||||
for _, s := range set {
|
||||
if s == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isUUID reports whether a string is shaped like a canonical UUID. Cheap enough
|
||||
// to run per request and it keeps a malformed id out of the SQL entirely.
|
||||
func isUUID(s string) bool {
|
||||
if len(s) != 36 {
|
||||
return false
|
||||
}
|
||||
for i, c := range s {
|
||||
switch i {
|
||||
case 8, 13, 18, 23:
|
||||
if c != '-' {
|
||||
return false
|
||||
}
|
||||
default:
|
||||
isHex := (c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F')
|
||||
if !isHex {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
/* ── Registry ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// Registry holds one service per resource that has an endpoint.
|
||||
type Registry struct {
|
||||
byPath map[string]*Service
|
||||
order []*Service
|
||||
}
|
||||
|
||||
// NewRegistry builds services for every resource in domain.AllResources.
|
||||
func NewRegistry(db repo.Querier) *Registry {
|
||||
reg := &Registry{byPath: make(map[string]*Service, len(domain.AllResources))}
|
||||
for _, res := range domain.AllResources {
|
||||
svc := New(res, db)
|
||||
reg.byPath[res.Path] = svc
|
||||
reg.order = append(reg.order, svc)
|
||||
}
|
||||
return reg
|
||||
}
|
||||
|
||||
// Get returns the service for a URL path segment.
|
||||
func (r *Registry) Get(path string) (*Service, bool) {
|
||||
s, ok := r.byPath[path]
|
||||
return s, ok
|
||||
}
|
||||
|
||||
// All returns every service, in declaration order.
|
||||
func (r *Registry) All() []*Service { return r.order }
|
||||
217
go-api/internal/service/service_test.go
Normal file
217
go-api/internal/service/service_test.go
Normal file
@@ -0,0 +1,217 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/domain"
|
||||
)
|
||||
|
||||
// These exercise query parsing and validation without a database, so the
|
||||
// contract's defaults are pinned even when PostgreSQL is not available.
|
||||
|
||||
func resource(t *testing.T, path string) *domain.Resource {
|
||||
t.Helper()
|
||||
res, ok := domain.ResourceByPath[path]
|
||||
if !ok {
|
||||
t.Fatalf("no resource for path %q", path)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func parse(t *testing.T, path, query string) (domain.ListParams, error) {
|
||||
t.Helper()
|
||||
values, err := url.ParseQuery(query)
|
||||
if err != nil {
|
||||
t.Fatalf("bad test query %q: %v", query, err)
|
||||
}
|
||||
return New(resource(t, path), nil).ParseList(values)
|
||||
}
|
||||
|
||||
// Each default is the literal argument at the frontend call site
|
||||
// (api-contract.md §8.1), not a generic value.
|
||||
func TestParseListDefaults(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
path string
|
||||
limit int
|
||||
sort string
|
||||
desc bool
|
||||
}{
|
||||
{"job-postings", 100, "created_date", true},
|
||||
{"job-applications", 200, "ai_score", true},
|
||||
{"worker-profiles", 500, "krow_score", true},
|
||||
{"shift-records", 500, "created_date", true},
|
||||
{"user-activity", 500, "created_date", true},
|
||||
{"courses", 200, "created_date", true},
|
||||
} {
|
||||
p, err := parse(t, tc.path, "")
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", tc.path, err)
|
||||
}
|
||||
if p.Limit != tc.limit {
|
||||
t.Errorf("%s limit = %d, want %d", tc.path, p.Limit, tc.limit)
|
||||
}
|
||||
if p.Sort != tc.sort || p.Desc != tc.desc {
|
||||
t.Errorf("%s sort = %q desc=%v, want %q desc=%v", tc.path, p.Sort, p.Desc, tc.sort, tc.desc)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// An explicitly empty ?sort= means no ordering, matching `if (!sort) return records`.
|
||||
func TestParseListEmptySortMeansUnordered(t *testing.T) {
|
||||
p, err := parse(t, "job-postings", "sort=")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if p.Sort != "" {
|
||||
t.Errorf("sort = %q, want empty", p.Sort)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseListLimitClampAndRejection(t *testing.T) {
|
||||
p, err := parse(t, "job-postings", "limit=99999")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if p.Limit != MaxLimit {
|
||||
t.Errorf("limit = %d, want it clamped to %d", p.Limit, MaxLimit)
|
||||
}
|
||||
// A zero limit is legitimate: store.js's slice(0, 0) returns nothing.
|
||||
if p, err := parse(t, "job-postings", "limit=0"); err != nil || p.Limit != 0 {
|
||||
t.Errorf("limit=0 -> %d, %v", p.Limit, err)
|
||||
}
|
||||
for _, bad := range []string{"limit=-1", "limit=abc", "offset=-3", "offset=x"} {
|
||||
if _, err := parse(t, "job-postings", bad); err == nil {
|
||||
t.Errorf("%s was accepted", bad)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseListRejectsUnknownSortAndFilter(t *testing.T) {
|
||||
if _, err := parse(t, "job-postings", "sort=-nope"); err == nil {
|
||||
t.Error("unknown sort field was accepted")
|
||||
}
|
||||
if _, err := parse(t, "job-postings", "nope=1"); err == nil {
|
||||
t.Error("unknown filter field was accepted")
|
||||
}
|
||||
// Arrays and JSON are not comparable for equality, so they are not filterable.
|
||||
if _, err := parse(t, "job-postings", "responsibilities=x"); err == nil {
|
||||
t.Error("an array column was accepted as a filter")
|
||||
}
|
||||
if _, err := parse(t, "job-postings", "vetting_criteria=x"); err == nil {
|
||||
t.Error("a jsonb column was accepted as a filter")
|
||||
}
|
||||
}
|
||||
|
||||
// A repeated parameter is membership, matching `Array.isArray(want)`.
|
||||
func TestParseListRepeatedParameterIsMembership(t *testing.T) {
|
||||
p, err := parse(t, "job-applications", "status=hired&status=interview")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(p.Filters) != 1 {
|
||||
t.Fatalf("filters = %d, want 1", len(p.Filters))
|
||||
}
|
||||
if len(p.Filters[0].Values) != 2 {
|
||||
t.Errorf("values = %v, want two", p.Filters[0].Values)
|
||||
}
|
||||
}
|
||||
|
||||
// sort, limit and offset are reserved; no column collides with them.
|
||||
func TestReservedParametersAreNotFilters(t *testing.T) {
|
||||
p, err := parse(t, "job-applications", "sort=-ai_score&limit=5&offset=2&status=hired")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(p.Filters) != 1 || p.Filters[0].Column.Name != "status" {
|
||||
t.Errorf("filters = %#v, want only status", p.Filters)
|
||||
}
|
||||
if p.Limit != 5 || p.Offset != 2 {
|
||||
t.Errorf("limit/offset = %d/%d, want 5/2", p.Limit, p.Offset)
|
||||
}
|
||||
for _, r := range []string{"sort", "limit", "offset"} {
|
||||
if _, isColumn := resource(t, "job-applications").Column(r); isColumn {
|
||||
t.Errorf("a column named %q collides with a reserved parameter", r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateRequiredAndUnknownAndEnum(t *testing.T) {
|
||||
svc := New(resource(t, "job-postings"), nil)
|
||||
|
||||
if _, err := svc.validate(domain.Record{}, true); err == nil {
|
||||
t.Error("a create with no title was accepted")
|
||||
}
|
||||
if _, err := svc.validate(domain.Record{"title": " "}, true); err == nil {
|
||||
t.Error("a blank title was accepted")
|
||||
}
|
||||
if _, err := svc.validate(domain.Record{"title": "X", "bogus": 1}, true); err == nil {
|
||||
t.Error("an unknown field was accepted")
|
||||
}
|
||||
if _, err := svc.validate(domain.Record{"title": "X", "status": "archived"}, true); err == nil {
|
||||
t.Error("an invalid enum value was accepted")
|
||||
}
|
||||
if _, err := svc.validate(domain.Record{"title": "X", "status": "active"}, true); err != nil {
|
||||
t.Errorf("a valid payload was rejected: %v", err)
|
||||
}
|
||||
|
||||
// Server-owned fields are stripped, not rejected.
|
||||
out, err := svc.validate(domain.Record{"title": "X", "id": "abc", "org_id": "def"}, true)
|
||||
if err != nil {
|
||||
t.Fatalf("server-owned fields caused a rejection: %v", err)
|
||||
}
|
||||
if _, present := out["id"]; present {
|
||||
t.Error("id survived validation")
|
||||
}
|
||||
if _, present := out["org_id"]; present {
|
||||
t.Error("org_id survived validation")
|
||||
}
|
||||
|
||||
// An update needs no required fields — it is a partial by definition.
|
||||
if _, err := svc.validate(domain.Record{"location": "Here"}, false); err != nil {
|
||||
t.Errorf("a partial update was rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsUUID(t *testing.T) {
|
||||
valid := []string{
|
||||
"00000000-0000-0000-0000-000000000000",
|
||||
"9A88DEBC-76E5-572C-A7E7-6EB5F43A6705",
|
||||
}
|
||||
for _, v := range valid {
|
||||
if !isUUID(v) {
|
||||
t.Errorf("%q rejected", v)
|
||||
}
|
||||
}
|
||||
invalid := []string{"", "not-a-uuid", "00000000000000000000000000000000",
|
||||
"00000000-0000-0000-0000-00000000000g", "00000000-0000-0000-0000-0000000000000"}
|
||||
for _, v := range invalid {
|
||||
if isUUID(v) {
|
||||
t.Errorf("%q accepted", v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Every resource must declare a sort column that actually exists, and a limit
|
||||
// in range — a typo in the generated metadata would otherwise only surface as a
|
||||
// 400 at runtime.
|
||||
func TestEveryResourceIsCoherent(t *testing.T) {
|
||||
for _, res := range domain.AllResources {
|
||||
field := res.DefaultSort
|
||||
if len(field) > 0 && field[0] == '-' {
|
||||
field = field[1:]
|
||||
}
|
||||
if !res.Sortable(field) {
|
||||
t.Errorf("%s: default sort %q is not a column", res.Name, res.DefaultSort)
|
||||
}
|
||||
if res.DefaultLimit < 1 || res.DefaultLimit > MaxLimit {
|
||||
t.Errorf("%s: default limit %d is out of range", res.Name, res.DefaultLimit)
|
||||
}
|
||||
if _, ok := res.Column("id"); !ok {
|
||||
t.Errorf("%s: no id column, so the sort tiebreaker cannot apply", res.Name)
|
||||
}
|
||||
if _, ok := res.Column("org_id"); !ok {
|
||||
t.Errorf("%s: no org_id column, so it cannot be scoped", res.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
309
go-api/internal/testutil/db.go
Normal file
309
go-api/internal/testutil/db.go
Normal file
@@ -0,0 +1,309 @@
|
||||
// Package testutil builds a disposable, fully migrated and seeded database for
|
||||
// tests.
|
||||
//
|
||||
// The test database is created from scratch on every run and is named
|
||||
// distinctly from any real one. Nothing here ever connects to, reads or drops
|
||||
// the development database.
|
||||
package testutil
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/seeder"
|
||||
)
|
||||
|
||||
// testDBPrefix names the throwaway databases. The prefix is deliberate: a name
|
||||
// this specific cannot be mistaken for, or collide with, "Krow-force".
|
||||
//
|
||||
// The pid is appended because `go test ./...` runs each package in its own
|
||||
// process, concurrently — a single shared name means one package drops the
|
||||
// database another is still using.
|
||||
const testDBPrefix = "krow_backend_autotest"
|
||||
|
||||
// TestDBName is this process's throwaway database.
|
||||
var TestDBName = fmt.Sprintf("%s_%d", testDBPrefix, os.Getpid())
|
||||
|
||||
// Harness is a ready database plus what was seeded into it.
|
||||
type Harness struct {
|
||||
Pool *pgxpool.Pool
|
||||
OrgID string
|
||||
Seeded *seeder.Result
|
||||
Now time.Time
|
||||
}
|
||||
|
||||
func env(key, fallback string) string {
|
||||
if v := strings.TrimSpace(os.Getenv(key)); v != "" {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func dsn(database string) string {
|
||||
return fmt.Sprintf("postgres://%s:%s@%s:%s/%s?sslmode=disable",
|
||||
env("DATABASE_USER", "postgres"), env("DATABASE_PASSWORD", ""),
|
||||
env("DATABASE_HOST", "127.0.0.1"), env("DATABASE_PORT", "5432"), database)
|
||||
}
|
||||
|
||||
// repoRoot walks up from the test's working directory to the repository root,
|
||||
// found by the migrations directory sitting beside go-api.
|
||||
func repoRoot(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatalf("getwd: %v", err)
|
||||
}
|
||||
for i := 0; i < 6; i++ {
|
||||
if _, err := os.Stat(filepath.Join(dir, "migrations")); err == nil {
|
||||
return dir
|
||||
}
|
||||
dir = filepath.Dir(dir)
|
||||
}
|
||||
t.Fatalf("could not locate the repository root from the test working directory")
|
||||
return ""
|
||||
}
|
||||
|
||||
// New builds a migrated, seeded database, or skips the test when PostgreSQL is
|
||||
// not reachable — so `go test ./...` still runs on a machine without a server.
|
||||
func New(t *testing.T) *Harness {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
|
||||
admin, err := pgxpool.New(ctx, dsn("postgres"))
|
||||
if err != nil {
|
||||
t.Skipf("PostgreSQL unavailable, skipping database tests: %v", err)
|
||||
}
|
||||
if err := admin.Ping(ctx); err != nil {
|
||||
admin.Close()
|
||||
t.Skipf("PostgreSQL unavailable, skipping database tests: %v", err)
|
||||
}
|
||||
|
||||
// Terminate stragglers so DROP cannot block on a leaked connection.
|
||||
_, _ = admin.Exec(ctx,
|
||||
`SELECT pg_terminate_backend(pid) FROM pg_stat_activity
|
||||
WHERE datname = $1 AND pid <> pg_backend_pid()`, TestDBName)
|
||||
if _, err := admin.Exec(ctx, `DROP DATABASE IF EXISTS `+quoteIdent(TestDBName)); err != nil {
|
||||
admin.Close()
|
||||
t.Fatalf("drop test database: %v", err)
|
||||
}
|
||||
if _, err := admin.Exec(ctx, `CREATE DATABASE `+quoteIdent(TestDBName)); err != nil {
|
||||
admin.Close()
|
||||
t.Fatalf("create test database: %v", err)
|
||||
}
|
||||
admin.Close()
|
||||
|
||||
pool, err := pgxpool.New(ctx, dsn(TestDBName))
|
||||
if err != nil {
|
||||
t.Fatalf("connect to test database: %v", err)
|
||||
}
|
||||
|
||||
root := repoRoot(t)
|
||||
applyMigrations(t, ctx, pool, filepath.Join(root, "migrations"))
|
||||
|
||||
fixture, err := seeder.Load(filepath.Join(root, "seed", "fixtures", "seed.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("load fixture: %v", err)
|
||||
}
|
||||
now := time.Now()
|
||||
result, err := seeder.New(pool, fixture, now).Run(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("seed: %v", err)
|
||||
}
|
||||
|
||||
t.Cleanup(func() {
|
||||
pool.Close()
|
||||
dropTestDatabase()
|
||||
})
|
||||
return &Harness{Pool: pool, OrgID: result.OrgID, Seeded: result, Now: now}
|
||||
}
|
||||
|
||||
// dropTestDatabase removes this process's throwaway database. Best effort: a
|
||||
// leftover is harmless because the next run drops it before creating it.
|
||||
func dropTestDatabase() {
|
||||
ctx := context.Background()
|
||||
admin, err := pgxpool.New(ctx, dsn("postgres"))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer admin.Close()
|
||||
_, _ = admin.Exec(ctx,
|
||||
`SELECT pg_terminate_backend(pid) FROM pg_stat_activity
|
||||
WHERE datname = $1 AND pid <> pg_backend_pid()`, TestDBName)
|
||||
_, _ = admin.Exec(ctx, `DROP DATABASE IF EXISTS `+quoteIdent(TestDBName))
|
||||
}
|
||||
|
||||
// applyMigrations runs every *.up.sql in filename order. This is the same SQL
|
||||
// golang-migrate applies; running it directly keeps the tests independent of
|
||||
// the CLI being installed.
|
||||
func applyMigrations(t *testing.T, ctx context.Context, pool *pgxpool.Pool, dir string) {
|
||||
t.Helper()
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("read migrations: %v", err)
|
||||
}
|
||||
var files []string
|
||||
for _, e := range entries {
|
||||
if strings.HasSuffix(e.Name(), ".up.sql") {
|
||||
files = append(files, e.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(files)
|
||||
if len(files) == 0 {
|
||||
t.Fatal("no migrations found")
|
||||
}
|
||||
for _, name := range files {
|
||||
sqlBytes, err := os.ReadFile(filepath.Join(dir, name))
|
||||
if err != nil {
|
||||
t.Fatalf("read %s: %v", name, err)
|
||||
}
|
||||
if _, err := pool.Exec(ctx, string(sqlBytes)); err != nil {
|
||||
t.Fatalf("apply %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// quoteIdent renders an identifier safely. The only value passed here is the
|
||||
// package constant above, but building DDL by concatenation without quoting is
|
||||
// a habit worth not having.
|
||||
func quoteIdent(s string) string {
|
||||
return `"` + strings.ReplaceAll(s, `"`, `""`) + `"`
|
||||
}
|
||||
|
||||
// Fixture reloads the raw fixture so tests can assert the database against the
|
||||
// frontend's own data rather than against numbers typed into a test.
|
||||
func Fixture(t *testing.T) *seeder.Fixture {
|
||||
t.Helper()
|
||||
f, err := seeder.Load(filepath.Join(repoRoot(t), "seed", "fixtures", "seed.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("load fixture: %v", err)
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
/* ── Migration sandboxes ────────────────────────────────────────────────────
|
||||
*
|
||||
* The helpers below exist for tests that drive the migration FILES themselves
|
||||
* — applying them, rolling them back, re-applying them — rather than using the
|
||||
* migrated database New() hands out.
|
||||
*
|
||||
* They need a database of their own for two reasons. New()'s database is
|
||||
* dropped by its own t.Cleanup, so sharing it across a test that rolls the
|
||||
* schema back would leave the next test's fixtures on the floor; and a down
|
||||
* migration must run against a database whose contents the test controls,
|
||||
* because 000003's down migration deliberately fails on seeded data.
|
||||
*
|
||||
* Like New(), nothing here can reach a real database: every name is built from
|
||||
* testDBPrefix, which cannot be confused with "Krow-force".
|
||||
*/
|
||||
|
||||
// Sandbox creates an empty throwaway database and returns a pool on it.
|
||||
//
|
||||
// Nothing is migrated and nothing is seeded — that is the point. The label
|
||||
// distinguishes concurrent sandboxes within one package; the pid keeps
|
||||
// packages, which `go test ./...` runs in parallel processes, from colliding.
|
||||
func Sandbox(t *testing.T, label string) *pgxpool.Pool {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
name := fmt.Sprintf("%s_%s_%d", testDBPrefix, label, os.Getpid())
|
||||
|
||||
admin, err := pgxpool.New(ctx, dsn("postgres"))
|
||||
if err != nil {
|
||||
t.Skipf("PostgreSQL unavailable, skipping database tests: %v", err)
|
||||
}
|
||||
if err := admin.Ping(ctx); err != nil {
|
||||
admin.Close()
|
||||
t.Skipf("PostgreSQL unavailable, skipping database tests: %v", err)
|
||||
}
|
||||
dropDatabase(ctx, admin, name)
|
||||
if _, err := admin.Exec(ctx, `CREATE DATABASE `+quoteIdent(name)); err != nil {
|
||||
admin.Close()
|
||||
t.Fatalf("create sandbox database %s: %v", name, err)
|
||||
}
|
||||
admin.Close()
|
||||
|
||||
pool, err := pgxpool.New(ctx, dsn(name))
|
||||
if err != nil {
|
||||
t.Fatalf("connect to sandbox database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
pool.Close()
|
||||
cleanup, err := pgxpool.New(context.Background(), dsn("postgres"))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer cleanup.Close()
|
||||
dropDatabase(context.Background(), cleanup, name)
|
||||
})
|
||||
return pool
|
||||
}
|
||||
|
||||
// dropDatabase terminates stragglers, then drops. Best effort on the drop
|
||||
// itself: a leftover is harmless because the next run drops it before creating.
|
||||
func dropDatabase(ctx context.Context, admin *pgxpool.Pool, name string) {
|
||||
_, _ = admin.Exec(ctx,
|
||||
`SELECT pg_terminate_backend(pid) FROM pg_stat_activity
|
||||
WHERE datname = $1 AND pid <> pg_backend_pid()`, name)
|
||||
_, _ = admin.Exec(ctx, `DROP DATABASE IF EXISTS `+quoteIdent(name))
|
||||
}
|
||||
|
||||
// RepoRoot is the repository root, located from the test's working directory.
|
||||
func RepoRoot(t *testing.T) string {
|
||||
t.Helper()
|
||||
return repoRoot(t)
|
||||
}
|
||||
|
||||
// MigrationsDir is the directory holding the migration files.
|
||||
func MigrationsDir(t *testing.T) string {
|
||||
t.Helper()
|
||||
return filepath.Join(repoRoot(t), "migrations")
|
||||
}
|
||||
|
||||
// MigrationFiles lists the migration files with the given suffix — ".up.sql"
|
||||
// or ".down.sql" — in filename order. Callers wanting to roll back should
|
||||
// reverse the result.
|
||||
func MigrationFiles(t *testing.T, suffix string) []string {
|
||||
t.Helper()
|
||||
entries, err := os.ReadDir(MigrationsDir(t))
|
||||
if err != nil {
|
||||
t.Fatalf("read migrations: %v", err)
|
||||
}
|
||||
var files []string
|
||||
for _, e := range entries {
|
||||
if strings.HasSuffix(e.Name(), suffix) {
|
||||
files = append(files, e.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(files)
|
||||
if len(files) == 0 {
|
||||
t.Fatalf("no %s migrations found", suffix)
|
||||
}
|
||||
return files
|
||||
}
|
||||
|
||||
// ApplyMigration runs one migration file and returns its error rather than
|
||||
// failing the test, so a test can assert that a rollback succeeds — or, for
|
||||
// 000003's down migration, that it does not.
|
||||
func ApplyMigration(ctx context.Context, t *testing.T, pool *pgxpool.Pool, name string) error {
|
||||
t.Helper()
|
||||
sqlBytes, err := os.ReadFile(filepath.Join(MigrationsDir(t), name))
|
||||
if err != nil {
|
||||
t.Fatalf("read %s: %v", name, err)
|
||||
}
|
||||
_, err = pool.Exec(ctx, string(sqlBytes))
|
||||
return err
|
||||
}
|
||||
|
||||
// ApplyAllMigrations applies every *.up.sql in order, failing the test on the
|
||||
// first that does not apply. This is the same SQL golang-migrate would run.
|
||||
func ApplyAllMigrations(ctx context.Context, t *testing.T, pool *pgxpool.Pool) {
|
||||
t.Helper()
|
||||
applyMigrations(t, ctx, pool, MigrationsDir(t))
|
||||
}
|
||||
Reference in New Issue
Block a user