first commit
This commit is contained in:
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)
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user