302 lines
7.4 KiB
Go
302 lines
7.4 KiB
Go
package crud
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"aijianzhan/platform/internal/meta"
|
|
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
// PostgresEngine dynamic row CRUD with tenant_id and optional org_unit_id scope.
|
|
type PostgresEngine struct {
|
|
DB *sql.DB
|
|
Pool *DBPool
|
|
}
|
|
|
|
func NewPostgresEngine(db *sql.DB, pool *DBPool) *PostgresEngine {
|
|
return &PostgresEngine{DB: db, Pool: pool}
|
|
}
|
|
|
|
func (e *PostgresEngine) conn(ref *meta.ResourceRef) (*sql.DB, error) {
|
|
if e.Pool != nil && ref != nil && ref.App != nil {
|
|
return e.Pool.ForApp(ref.App)
|
|
}
|
|
return e.DB, nil
|
|
}
|
|
|
|
func (e *PostgresEngine) List(ctx context.Context, ref *meta.ResourceRef, tenantID int64, page, pageSize int, filters map[string]string, sortBy string) ([]map[string]any, int, error) {
|
|
if err := ensureOp(ref, "list"); err != nil {
|
|
return nil, 0, err
|
|
}
|
|
if err := validateFilters(ref, filters); err != nil {
|
|
return nil, 0, err
|
|
}
|
|
if err := validateSort(ref, sortBy); err != nil {
|
|
return nil, 0, err
|
|
}
|
|
page, pageSize = pageBounds(ref, page, pageSize)
|
|
db, err := e.conn(ref)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
|
|
where := []string{`tenant_id = $1`}
|
|
args := []any{tenantID}
|
|
argN := 2
|
|
scope := RowScopeFrom(ctx)
|
|
where, args, argN = appendOrgFilter(where, args, argN, scope.OrgUnitIDs)
|
|
for col, val := range filters {
|
|
where = append(where, fmt.Sprintf("%s = $%d", quoteIdent(col), argN))
|
|
args = append(args, val)
|
|
argN++
|
|
}
|
|
whereSQL := strings.Join(where, " AND ")
|
|
table := qualifiedTable(ref)
|
|
|
|
var total int
|
|
countSQL := fmt.Sprintf("SELECT COUNT(1) FROM %s WHERE %s", table, whereSQL)
|
|
if err := db.QueryRowContext(ctx, countSQL, args...).Scan(&total); err != nil {
|
|
return nil, 0, fmt.Errorf("count: %w", err)
|
|
}
|
|
|
|
orderSQL := ""
|
|
if sortBy != "" {
|
|
desc := strings.HasPrefix(sortBy, "-")
|
|
field := strings.TrimPrefix(sortBy, "-")
|
|
dir := "ASC"
|
|
if desc {
|
|
dir = "DESC"
|
|
}
|
|
orderSQL = " ORDER BY " + quoteIdent(field) + " " + dir
|
|
}
|
|
|
|
offset := (page - 1) * pageSize
|
|
listArgs := append(append([]any{}, args...), pageSize, offset)
|
|
listSQL := fmt.Sprintf(
|
|
"SELECT * FROM %s WHERE %s%s LIMIT $%d OFFSET $%d",
|
|
table, whereSQL, orderSQL, argN, argN+1,
|
|
)
|
|
rows, err := db.QueryContext(ctx, listSQL, listArgs...)
|
|
if err != nil {
|
|
return nil, 0, fmt.Errorf("list: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
|
|
items, err := scanRows(rows)
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
return items, total, nil
|
|
}
|
|
|
|
func (e *PostgresEngine) Get(ctx context.Context, ref *meta.ResourceRef, tenantID int64, id string) (map[string]any, error) {
|
|
if err := ensureOp(ref, "get"); err != nil {
|
|
return nil, err
|
|
}
|
|
db, err := e.conn(ref)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
pk := ref.Entity.PrimaryKey
|
|
where := []string{`tenant_id = $1`, fmt.Sprintf("%s = $2", quoteIdent(pk))}
|
|
args := []any{tenantID, id}
|
|
argN := 3
|
|
scope := RowScopeFrom(ctx)
|
|
where, args, argN = appendOrgFilter(where, args, argN, scope.OrgUnitIDs)
|
|
_ = argN
|
|
q := fmt.Sprintf(
|
|
"SELECT * FROM %s WHERE %s LIMIT 1",
|
|
qualifiedTable(ref), strings.Join(where, " AND "),
|
|
)
|
|
rows, err := db.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
items, err := scanRows(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(items) == 0 {
|
|
return nil, fmt.Errorf("not found")
|
|
}
|
|
return items[0], nil
|
|
}
|
|
|
|
func (e *PostgresEngine) Create(ctx context.Context, ref *meta.ResourceRef, tenantID, userID int64, body map[string]any) (map[string]any, error) {
|
|
if err := ensureOp(ref, "create"); err != nil {
|
|
return nil, err
|
|
}
|
|
db, err := e.conn(ref)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
row, err := sanitizeBody(ref.Entity, body, false)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
pk := ref.Entity.PrimaryKey
|
|
autoPK := isAutoPK(ref.Entity)
|
|
if !autoPK {
|
|
if _, ok := row[pk]; !ok {
|
|
row[pk] = uuid.NewString()
|
|
}
|
|
} else {
|
|
delete(row, pk)
|
|
}
|
|
|
|
now := time.Now().UTC()
|
|
row["tenant_id"] = tenantID
|
|
row["created_by"] = userID
|
|
row["created_at"] = now
|
|
row["updated_at"] = now
|
|
if scope := RowScopeFrom(ctx); scope.WriteOrgUnit > 0 {
|
|
row["org_unit_id"] = scope.WriteOrgUnit
|
|
}
|
|
|
|
cols := make([]string, 0, len(row))
|
|
placeholders := make([]string, 0, len(row))
|
|
args := make([]any, 0, len(row))
|
|
i := 1
|
|
for k, v := range row {
|
|
cols = append(cols, quoteIdent(k))
|
|
placeholders = append(placeholders, fmt.Sprintf("$%d", i))
|
|
args = append(args, v)
|
|
i++
|
|
}
|
|
|
|
q := fmt.Sprintf(
|
|
"INSERT INTO %s (%s) VALUES (%s) RETURNING *",
|
|
qualifiedTable(ref),
|
|
strings.Join(cols, ", "),
|
|
strings.Join(placeholders, ", "),
|
|
)
|
|
rows, err := db.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("insert: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
items, err := scanRows(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(items) == 0 {
|
|
return nil, fmt.Errorf("insert returned no row")
|
|
}
|
|
return items[0], nil
|
|
}
|
|
|
|
func (e *PostgresEngine) Update(ctx context.Context, ref *meta.ResourceRef, tenantID int64, id string, body map[string]any) (map[string]any, error) {
|
|
if err := ensureOp(ref, "update"); err != nil {
|
|
return nil, err
|
|
}
|
|
db, err := e.conn(ref)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
patch, err := sanitizeBody(ref.Entity, body, true)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
delete(patch, ref.Entity.PrimaryKey)
|
|
delete(patch, "tenant_id")
|
|
delete(patch, "org_unit_id")
|
|
if len(patch) == 0 {
|
|
return e.Get(ctx, ref, tenantID, id)
|
|
}
|
|
patch["updated_at"] = time.Now().UTC()
|
|
|
|
sets := make([]string, 0, len(patch))
|
|
args := make([]any, 0, len(patch)+2)
|
|
i := 1
|
|
for k, v := range patch {
|
|
sets = append(sets, fmt.Sprintf("%s = $%d", quoteIdent(k), i))
|
|
args = append(args, v)
|
|
i++
|
|
}
|
|
where := []string{fmt.Sprintf("tenant_id = $%d", i), fmt.Sprintf("%s = $%d", quoteIdent(ref.Entity.PrimaryKey), i+1)}
|
|
args = append(args, tenantID, id)
|
|
argN := i + 2
|
|
scope := RowScopeFrom(ctx)
|
|
where, args, argN = appendOrgFilter(where, args, argN, scope.OrgUnitIDs)
|
|
_ = argN
|
|
q := fmt.Sprintf(
|
|
"UPDATE %s SET %s WHERE %s RETURNING *",
|
|
qualifiedTable(ref),
|
|
strings.Join(sets, ", "),
|
|
strings.Join(where, " AND "),
|
|
)
|
|
rows, err := db.QueryContext(ctx, q, args...)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("update: %w", err)
|
|
}
|
|
defer rows.Close()
|
|
items, err := scanRows(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(items) == 0 {
|
|
return nil, fmt.Errorf("not found")
|
|
}
|
|
return items[0], nil
|
|
}
|
|
|
|
func (e *PostgresEngine) Delete(ctx context.Context, ref *meta.ResourceRef, tenantID int64, id string) error {
|
|
if err := ensureOp(ref, "delete"); err != nil {
|
|
return err
|
|
}
|
|
db, err := e.conn(ref)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
where := []string{`tenant_id = $1`, fmt.Sprintf("%s = $2", quoteIdent(ref.Entity.PrimaryKey))}
|
|
args := []any{tenantID, id}
|
|
argN := 3
|
|
scope := RowScopeFrom(ctx)
|
|
where, args, argN = appendOrgFilter(where, args, argN, scope.OrgUnitIDs)
|
|
_ = argN
|
|
q := fmt.Sprintf(
|
|
"DELETE FROM %s WHERE %s",
|
|
qualifiedTable(ref), strings.Join(where, " AND "),
|
|
)
|
|
res, err := db.ExecContext(ctx, q, args...)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
n, _ := res.RowsAffected()
|
|
if n == 0 {
|
|
return fmt.Errorf("not found")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func scanRows(rows *sql.Rows) ([]map[string]any, error) {
|
|
cols, err := rows.Columns()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
out := make([]map[string]any, 0)
|
|
for rows.Next() {
|
|
vals := make([]any, len(cols))
|
|
ptrs := make([]any, len(cols))
|
|
for i := range vals {
|
|
ptrs[i] = &vals[i]
|
|
}
|
|
if err := rows.Scan(ptrs...); err != nil {
|
|
return nil, err
|
|
}
|
|
row := make(map[string]any, len(cols))
|
|
for i, c := range cols {
|
|
row[c] = normalizeDBValue(vals[i])
|
|
}
|
|
out = append(out, row)
|
|
}
|
|
return out, rows.Err()
|
|
}
|