// 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: // The int8/16/32 cases are not JSON shapes — encoding/json only ever // produces float64. They are the widths pgx hands back when a value was // READ from the database and is being written somewhere else: an `int` // column arrives as int32, and a service that copies a field from one // record to another (see service.Hire, which carries ai_score from an // application onto the staff row) would otherwise be told its own // database's value "must be a number". 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 case int32: return int64(t), nil case int16: return int64(t), nil case int8: return int64(t), nil case float32: if t != float32(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 } return nil, domain.Validation(fmt.Sprintf("%s must be a number", c.Name), nil) case domain.KindFloat: // float32 and the integer widths for the same reason as above: a // numeric column is projected as float8 and returns float64, but an // integer read from elsewhere may legitimately be written into one. switch t := v.(type) { case float64: return t, nil case float32: return float64(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 case int64: return float64(t), nil case int: return float64(t), nil case int32: return float64(t), nil case int16: return float64(t), 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 }