Files
krow_backend/go-api/internal/repo/repo.go
2026-08-24 13:06:29 +05:30

669 lines
21 KiB
Go

// Package repo is the PostgreSQL access layer.
//
// Every statement is built from a *domain.Resource: column lists are explicit
// and come from the generated descriptors, and every value reaches the database
// as a bind parameter cast to its declared type. No identifier is ever taken
// from user input — a filter or sort name is resolved to a *domain.Column
// first, and an unresolved name is rejected before any SQL is assembled.
package repo
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/krow/krow-backend/go-api/internal/authctx"
"github.com/krow/krow-backend/go-api/internal/domain"
)
// Querier is satisfied by both *pgxpool.Pool and pgx.Tx, so every method here
// works inside or outside a transaction.
type Querier interface {
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
}
// Repo reads and writes one resource.
type Repo struct {
res *domain.Resource
db Querier
}
// New builds a repository for a resource over a pool or transaction.
func New(res *domain.Resource, db Querier) *Repo { return &Repo{res: res, db: db} }
// Resource is the descriptor this repository serves.
func (r *Repo) Resource() *domain.Resource { return r.res }
/* ── Projection ─────────────────────────────────────────────────────────── */
func (r *Repo) selectList() string {
parts := make([]string, 0, len(r.res.Columns))
for _, c := range r.res.Columns {
parts = append(parts, c.SelectExpr())
}
return strings.Join(parts, ", ")
}
/* ── Predicates ─────────────────────────────────────────────────────────── */
type builder struct {
args []any
where []string
}
func (b *builder) add(v any) string {
b.args = append(b.args, v)
return "$" + strconv.Itoa(len(b.args))
}
// scope applies organization scoping. A NULL org_id on an OrgNullable table
// means the shared platform library, which every organization can read.
func (b *builder) scope(res *domain.Resource, orgID string) {
p := b.add(orgID)
if res.OrgNullable {
b.where = append(b.where, fmt.Sprintf("(org_id = %s::uuid OR org_id IS NULL)", p))
return
}
b.where = append(b.where, fmt.Sprintf("org_id = %s::uuid", p))
}
// ownership narrows the rows a caller may touch, beyond their organization.
//
// This is the second half of authorization and it lives HERE, in the WHERE
// clause, rather than in a loop over fetched rows. The difference is not
// stylistic: List runs `count(*)` over the same predicate, so a filtered-in-Go
// approach would return the right page and the wrong total, and would fetch
// rows the caller may not see in order to discard them. A row outside the
// predicate is never read.
//
// Only talent is scoped — Policy.ScopeFor decides that, not this function.
// Admin and employer see the whole organization, which is what an operator
// console is for.
//
// A role the API does not recognise matches nothing. That should be
// unreachable — the handler answers 403 before any query runs — but a
// scope-narrowing function whose failure mode is "see everything" is the wrong
// shape to leave lying around.
func (b *builder) ownership(res *domain.Resource, ident authctx.Identity) {
role, known := domain.ParseRole(ident.Role)
if !known {
b.where = append(b.where, "false")
return
}
scope := res.Policy.ScopeFor(role)
switch scope.Kind {
case domain.ScopeNone:
return
case domain.ScopeUserID:
col, ok := res.Column(scope.Column)
if !ok {
b.where = append(b.where, "false")
return
}
b.where = append(b.where,
fmt.Sprintf("%s = %s::%s", col.Name, b.add(ident.UserID), col.PGType))
case domain.ScopeEmail:
col, ok := res.Column(scope.Column)
if !ok {
b.where = append(b.where, "false")
return
}
// citext, so the comparison is case-insensitive — the same equality the
// rest of the domain uses for email.
b.where = append(b.where,
fmt.Sprintf("%s = %s::%s", col.Name, b.add(ident.Email), col.PGType))
case domain.ScopeActivePostings:
col, ok := res.Column(scope.Column)
if !ok {
b.where = append(b.where, "false")
return
}
b.where = append(b.where,
fmt.Sprintf("%s = %s::%s", col.Name, b.add("active"), col.PGType))
case domain.ScopeOwnApplications:
// Ownership by reference: the row names an application, the application
// names a person. The subquery is scoped to the organization as well as
// the email, so it cannot reach across tenants even if an id from
// another one were supplied.
apps, ok := domain.ResourceByPath["job-applications"]
if !ok {
b.where = append(b.where, "false")
return
}
col, colOK := res.Column(scope.Column)
if !colOK {
b.where = append(b.where, "false")
return
}
b.where = append(b.where, fmt.Sprintf(
"%s IN (SELECT id FROM %s WHERE org_id = %s::uuid AND email = %s::citext)",
col.Name, apps.Table, b.add(ident.OrgID), b.add(ident.Email)))
default:
b.where = append(b.where, "false")
}
}
// filters renders the contract's two operators and nothing else: equality for a
// single value, membership for several. See api-contract.md §6.
func (b *builder) filters(fs []domain.Filter) error {
for _, f := range fs {
if len(f.Values) == 1 {
v, err := bindValue(*f.Column, f.Values[0])
if err != nil {
return err
}
b.where = append(b.where,
fmt.Sprintf("%s = %s::%s", f.Column.Name, b.add(v), f.Column.PGType))
continue
}
vals := make([]string, 0, len(f.Values))
vals = append(vals, f.Values...)
b.where = append(b.where,
fmt.Sprintf("%s = ANY(%s::%s[])", f.Column.Name, b.add(vals), arrayElem(f.Column)))
}
return nil
}
func arrayElem(c *domain.Column) string {
if c.Kind == domain.KindEnum {
return c.PGType
}
switch c.PGType {
case "citext":
return "citext"
case "uuid":
return "uuid"
case "int", "bigint":
return c.PGType
default:
return "text"
}
}
func (b *builder) clause() string {
if len(b.where) == 0 {
return ""
}
return " WHERE " + strings.Join(b.where, " AND ")
}
/* ── Ordering ───────────────────────────────────────────────────────────── */
// orderBy renders the two ordering rules the contract calls out.
//
// NULLS LAST in *both* directions, because store.js's comparator returns before
// the descending negation is applied — PostgreSQL's default would put nulls
// first on DESC. And `, id` as a final tiebreaker, because JavaScript's sort is
// stable and PostgreSQL's is not. See api-contract.md §7.1 and §7.3.
//
// Names are qualified with the table so they bind to the real column rather
// than to a same-named output alias from the projection.
func (r *Repo) orderBy(p domain.ListParams) string {
if p.Sort == "" {
return fmt.Sprintf(" ORDER BY %s.id", r.res.Table)
}
dir := "ASC"
if p.Desc {
dir = "DESC"
}
return fmt.Sprintf(" ORDER BY %s.%s %s NULLS LAST, %s.id",
r.res.Table, p.Sort, dir, r.res.Table)
}
/* ── Reads ──────────────────────────────────────────────────────────────── */
// List returns one page plus the total matching count.
func (r *Repo) List(ctx context.Context, ident authctx.Identity, p domain.ListParams) (*domain.Page, error) {
b := &builder{}
b.scope(r.res, ident.OrgID)
b.ownership(r.res, ident)
if err := b.filters(p.Filters); err != nil {
return nil, err
}
where := b.clause()
var total int
countSQL := "SELECT count(*) FROM " + r.res.Table + where
if err := r.db.QueryRow(ctx, countSQL, b.args...).Scan(&total); err != nil {
return nil, fmt.Errorf("count %s: %w", r.res.Table, err)
}
limitP := b.add(p.Limit)
offsetP := b.add(p.Offset)
q := "SELECT " + r.selectList() + " FROM " + r.res.Table + where +
r.orderBy(p) + " LIMIT " + limitP + " OFFSET " + offsetP
rows, err := r.db.Query(ctx, q, b.args...)
if err != nil {
return nil, fmt.Errorf("list %s: %w", r.res.Table, err)
}
defer rows.Close()
records, err := collect(rows)
if err != nil {
return nil, err
}
return &domain.Page{Records: records, Total: total, Limit: p.Limit, Offset: p.Offset}, nil
}
// Get returns one record, or a nil record when nothing matches.
func (r *Repo) Get(ctx context.Context, ident authctx.Identity, id string) (domain.Record, error) {
b := &builder{}
b.scope(r.res, ident.OrgID)
b.ownership(r.res, ident)
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
q := "SELECT " + r.selectList() + " FROM " + r.res.Table + b.clause() + " LIMIT 1"
rows, err := r.db.Query(ctx, q, b.args...)
if err != nil {
return nil, fmt.Errorf("get %s: %w", r.res.Table, err)
}
defer rows.Close()
records, err := collect(rows)
if err != nil || len(records) == 0 {
return nil, err
}
return records[0], nil
}
/* ── Writes ─────────────────────────────────────────────────────────────── */
// derivedValues are the columns this caller does not get to choose.
//
// Every column here is also ReadOnly in the descriptors, so a value in the
// request body has already been dropped by service.validate. This supplies what
// goes in instead: a fact about the authenticated session.
//
// Two shapes share this mechanism and it is worth keeping them apart. `created_by`
// and the user_activity identity columns record WHO ACTED — always the session
// user, whatever their role. The TalentOnly ones record WHO THE ROW IS ABOUT,
// and only a talent caller is necessarily writing about themselves; when an
// admin creates a candidate's worker profile, the subject is the candidate, so
// nothing is derived and the column is left to the database's default (NULL).
func (r *Repo) derivedValues(ident authctx.Identity) map[string]any {
if r.res.Policy == nil || len(r.res.Policy.Derived) == 0 {
return nil
}
isTalent := ident.Role == string(domain.RoleTalent)
out := make(map[string]any, len(r.res.Policy.Derived))
for _, d := range r.res.Policy.Derived {
if d.TalentOnly && !isTalent {
continue
}
switch d.Source {
case domain.DeriveUserID:
out[d.Column] = ident.UserID
case domain.DeriveEmail:
out[d.Column] = ident.Email
case domain.DeriveFullName:
out[d.Column] = ident.FullName
case domain.DeriveAccountType:
out[d.Column] = ident.AccountType
}
}
return out
}
// guardInsert refuses a create whose ownership is expressed by reference rather
// than by a column of its own.
//
// Only ai_interviews needs this: the row names an application, and a talent
// caller may only interview for an application of theirs. There is no column on
// the interview to derive, so the reference itself has to be checked.
//
// The refusal is the same 404 an application that does not exist would produce.
// Answering "that application is not yours" would confirm it exists.
func (r *Repo) guardInsert(ctx context.Context, ident authctx.Identity, in domain.Record) error {
role, known := domain.ParseRole(ident.Role)
if !known {
return domain.Forbidden()
}
scope := r.res.Policy.ScopeFor(role)
if scope.Kind != domain.ScopeOwnApplications {
return nil
}
apps, ok := domain.ResourceByPath["job-applications"]
if !ok {
return domain.Internal(errors.New("repo: job-applications resource is missing"))
}
ref, _ := in[scope.Column].(string)
if ref == "" {
// A required reference that is absent is a validation problem, and the
// service's Required check has already reported it.
return nil
}
var owns bool
q := fmt.Sprintf(
`SELECT EXISTS (SELECT 1 FROM %s WHERE org_id = $1::uuid AND id = $2::uuid AND email = $3::citext)`,
apps.Table)
if err := r.db.QueryRow(ctx, q, ident.OrgID, ref, ident.Email).Scan(&owns); err != nil {
return fmt.Errorf("check application ownership: %w", err)
}
if !owns {
return domain.NotFound(apps.Name, ref)
}
return nil
}
// Insert writes a record and returns it as the API represents it.
func (r *Repo) Insert(ctx context.Context, ident authctx.Identity, in domain.Record) (domain.Record, error) {
if err := r.guardInsert(ctx, ident, in); err != nil {
return nil, err
}
cols := []string{"org_id"}
b := &builder{}
// The organization is the session's, never the body's. org_id is ReadOnly
// on every resource, so this is the only way a value reaches the column.
vals := []string{b.add(ident.OrgID) + "::uuid"}
derived := r.derivedValues(ident)
for _, c := range r.res.Columns {
// A derived column is written whether or not it is ReadOnly, and it
// overrides anything the caller sent — which is what closes the
// attribution holes: a talent caller cannot file an application, a
// profile or an audit entry under somebody else's name.
if v, ok := derived[c.Name]; ok {
cols = append(cols, c.Name)
vals = append(vals, b.add(v)+"::"+c.PGType)
continue
}
if c.ReadOnly {
continue
}
v, present := in[c.Name]
if !present {
continue // let the column default apply
}
bound, err := bindValue(c, v)
if err != nil {
return nil, err
}
cols = append(cols, c.Name)
vals = append(vals, b.add(bound)+"::"+c.PGType)
}
q := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s) RETURNING %s",
r.res.Table, strings.Join(cols, ", "), strings.Join(vals, ", "), r.selectList())
rows, err := r.db.Query(ctx, q, b.args...)
if err != nil {
return nil, translate(err)
}
defer rows.Close()
// pgx does not execute until the rows are read, so a constraint violation
// arrives here rather than from Query above. Both paths must translate.
records, err := collect(rows)
if err != nil {
return nil, translate(err)
}
if len(records) == 0 {
return nil, fmt.Errorf("insert %s returned no row", r.res.Table)
}
return records[0], nil
}
// Update applies a shallow merge: only the supplied columns are written.
//
// A key that is absent is left alone; a key present with null sets NULL. There
// is no deep merge — store.js spreads one level, and useSubmitChallenge relies
// on whole arrays being replaced rather than appended to. See api-contract.md §3.2.
func (r *Repo) Update(ctx context.Context, ident authctx.Identity, id string, patch domain.Record) (domain.Record, error) {
b := &builder{}
sets := []string{}
for _, c := range r.res.Columns {
if c.ReadOnly {
continue
}
v, present := patch[c.Name]
if !present {
continue
}
bound, err := bindValue(c, v)
if err != nil {
return nil, err
}
sets = append(sets, fmt.Sprintf("%s = %s::%s", c.Name, b.add(bound), c.PGType))
}
// updated_date is always server-owned. user_activity is append-only and has
// no such column, so this is conditional rather than assumed.
if _, ok := r.res.Column("updated_date"); ok {
sets = append(sets, "updated_date = now()")
}
if len(sets) == 0 {
return r.Get(ctx, ident, id)
}
b.scope(r.res, ident.OrgID)
b.ownership(r.res, ident)
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
q := fmt.Sprintf("UPDATE %s SET %s%s RETURNING %s",
r.res.Table, strings.Join(sets, ", "), b.clause(), r.selectList())
rows, err := r.db.Query(ctx, q, b.args...)
if err != nil {
return nil, translate(err)
}
defer rows.Close()
records, err := collect(rows)
if err != nil {
return nil, translate(err)
}
if len(records) == 0 {
return nil, nil
}
return records[0], nil
}
// Delete removes a record and reports how many rows went.
//
// A miss is not an error: store.js filters its array and returns { id }
// whether or not anything matched. See api-contract.md §12.7.
func (r *Repo) Delete(ctx context.Context, ident authctx.Identity, id string) (int64, error) {
b := &builder{}
b.scope(r.res, ident.OrgID)
b.ownership(r.res, ident)
b.where = append(b.where, "id = "+b.add(id)+"::uuid")
tag, err := r.db.Exec(ctx, "DELETE FROM "+r.res.Table+b.clause(), b.args...)
if err != nil {
return 0, translate(err)
}
return tag.RowsAffected(), nil
}
/* ── Scanning ───────────────────────────────────────────────────────────── */
func collect(rows pgx.Rows) ([]domain.Record, error) {
fields := rows.FieldDescriptions()
out := make([]domain.Record, 0, 16)
for rows.Next() {
vals, err := rows.Values()
if err != nil {
return nil, err
}
rec := make(domain.Record, len(fields))
for i, f := range fields {
rec[string(f.Name)] = normalise(vals[i])
}
out = append(out, rec)
}
return out, rows.Err()
}
// normalise flattens the few pgx representations that would not JSON-encode the
// way the frontend expects. Most values arrive ready to use because the
// projection casts them (see Column.SelectExpr).
func normalise(v any) any {
switch t := v.(type) {
case nil:
return nil
case [16]byte: // a uuid that slipped through without a ::text cast
return fmt.Sprintf("%x-%x-%x-%x-%x", t[0:4], t[4:6], t[6:8], t[8:10], t[10:16])
case []any:
out := make([]any, len(t))
for i, e := range t {
out[i] = normalise(e)
}
return out
default:
return v
}
}
/* ── Binding ────────────────────────────────────────────────────────────── */
// bindValue converts a decoded-JSON value into something pgx can send for a
// column of this type. Everything is explicitly cast in the SQL, so the job
// here is only to pick a Go representation PostgreSQL will accept.
func bindValue(c domain.Column, v any) (any, error) {
if v == nil {
return nil, nil
}
switch c.Kind {
case domain.KindTextArray:
switch t := v.(type) {
case []any:
out := make([]string, 0, len(t))
for _, e := range t {
s, ok := e.(string)
if !ok {
return nil, domain.Validation(
fmt.Sprintf("%s must be an array of strings", c.Name),
map[string]string{c.Name: "expected string elements"})
}
out = append(out, s)
}
return out, nil
case []string:
return t, nil
default:
return nil, domain.Validation(
fmt.Sprintf("%s must be an array", c.Name),
map[string]string{c.Name: "expected an array"})
}
case domain.KindJSON:
raw, err := json.Marshal(v)
if err != nil {
return nil, domain.Validation(fmt.Sprintf("%s is not encodable as JSON", c.Name), nil)
}
return raw, nil
case domain.KindInt:
switch t := v.(type) {
case float64:
if t != float64(int64(t)) {
return nil, domain.Validation(
fmt.Sprintf("%s must be a whole number", c.Name),
map[string]string{c.Name: "expected an integer"})
}
return int64(t), nil
case string:
n, err := strconv.ParseInt(t, 10, 64)
if err != nil {
return nil, domain.Validation(
fmt.Sprintf("%s must be a whole number", c.Name),
map[string]string{c.Name: "expected an integer"})
}
return n, nil
case int64:
return t, nil
case int:
return int64(t), nil
}
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
case domain.KindFloat:
switch t := v.(type) {
case float64:
return t, nil
case string:
f, err := strconv.ParseFloat(t, 64)
if err != nil {
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
}
return f, nil
}
return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil)
case domain.KindBool:
switch t := v.(type) {
case bool:
return t, nil
case string:
b, err := strconv.ParseBool(t)
if err != nil {
return nil, domain.Validation(fmt.Sprintf("%s must be a boolean", c.Name), nil)
}
return b, nil
}
return nil, domain.Validation(fmt.Sprintf("%s must be a boolean", c.Name), nil)
default: // strings, enums, uuids, dates, timestamps
s, ok := v.(string)
if !ok {
return nil, domain.Validation(
fmt.Sprintf("%s must be a string", c.Name),
map[string]string{c.Name: "expected a string"})
}
return s, nil
}
}
/* ── Error translation ──────────────────────────────────────────────────── */
// translate maps PostgreSQL's SQLSTATE codes onto the contract's error codes,
// so a constraint the database enforces surfaces as the documented API error
// rather than as a 500.
func translate(err error) error {
var pg *pgconn.PgError
if !errors.As(err, &pg) {
return err
}
switch pg.Code {
case "23505": // unique_violation
return domain.Conflict(pg.Detail)
case "23503": // foreign_key_violation
return domain.Validation("referenced record does not exist",
map[string]string{constraintField(pg): "no such record"})
case "23514": // check_violation
return domain.Validation("value violates constraint "+pg.ConstraintName,
map[string]string{constraintField(pg): "constraint " + pg.ConstraintName})
case "23502": // not_null_violation
return domain.Validation(pg.ColumnName+" must not be null",
map[string]string{pg.ColumnName: "required"})
case "22P02", "22007", "22008": // invalid text representation / datetime
return domain.Validation("value is not valid for its column type", nil)
}
return err
}
func constraintField(pg *pgconn.PgError) string {
if pg.ColumnName != "" {
return pg.ColumnName
}
return pg.ConstraintName
}