refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
10
modules/lagoon/attach/app_test.go
Normal file
10
modules/lagoon/attach/app_test.go
Normal file
@@ -0,0 +1,10 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
)
|
||||
|
||||
func backpackApp(cfg *compass.Config) *backpack.App {
|
||||
return backpack.New(cfg)
|
||||
}
|
||||
95
modules/lagoon/attach/bucket.go
Normal file
95
modules/lagoon/attach/bucket.go
Normal file
@@ -0,0 +1,95 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
"gocloud.dev/blob"
|
||||
_ "gocloud.dev/blob/fileblob"
|
||||
_ "gocloud.dev/blob/memblob"
|
||||
)
|
||||
|
||||
const defaultPublicPathPrefix = "/storage/uploads"
|
||||
|
||||
var (
|
||||
prefixMu sync.RWMutex
|
||||
publicPathPrefix = defaultPublicPathPrefix
|
||||
)
|
||||
|
||||
func setPublicPathPrefix(prefix string) {
|
||||
prefixMu.Lock()
|
||||
publicPathPrefix = prefix
|
||||
prefixMu.Unlock()
|
||||
}
|
||||
|
||||
// PublicPathPrefix is the URL prefix prepended to partition+filename.
|
||||
func PublicPathPrefix() string {
|
||||
prefixMu.RLock()
|
||||
defer prefixMu.RUnlock()
|
||||
return publicPathPrefix
|
||||
}
|
||||
|
||||
func normalizeBucketURL(raw string) (string, error) {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("attach: bucket_url: %w", err)
|
||||
}
|
||||
switch strings.ToLower(u.Scheme) {
|
||||
case "mem", "memory":
|
||||
return "mem://", nil
|
||||
case "fileblob":
|
||||
u.Scheme = "file"
|
||||
fallthrough
|
||||
case "file":
|
||||
q := u.Query()
|
||||
if q.Get("create_dir") == "" {
|
||||
q.Set("create_dir", "true")
|
||||
u.RawQuery = q.Encode()
|
||||
}
|
||||
return u.String(), nil
|
||||
default:
|
||||
return raw, nil
|
||||
}
|
||||
}
|
||||
|
||||
// OpenBucket opens storage.uploads.bucket_url (file:// or mem://) and
|
||||
// records storage.uploads.public_path_prefix. An empty bucket_url fails boot.
|
||||
func OpenBucket(ctx context.Context, cfg *compass.Config) (*blob.Bucket, error) {
|
||||
if cfg == nil {
|
||||
return nil, fmt.Errorf("attach: config is nil")
|
||||
}
|
||||
raw := strings.TrimSpace(cfg.String("storage.uploads.bucket_url"))
|
||||
if raw == "" {
|
||||
return nil, fmt.Errorf("attach: storage.uploads.bucket_url is empty (set SUMMER_STORAGE__UPLOADS__BUCKET_URL)")
|
||||
}
|
||||
bucketURL, err := normalizeBucketURL(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bucket, err := blob.OpenBucket(ctx, bucketURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("attach: open bucket %q: %w", raw, err)
|
||||
}
|
||||
prefix := strings.TrimSpace(cfg.String("storage.uploads.public_path_prefix"))
|
||||
if prefix == "" {
|
||||
prefix = defaultPublicPathPrefix
|
||||
}
|
||||
setPublicPathPrefix(prefix)
|
||||
return bucket, nil
|
||||
}
|
||||
|
||||
// Publish stores the opened bucket once on the app, matching lagoon.Publish.
|
||||
func Publish(app *backpack.App, bucket *blob.Bucket) error {
|
||||
if app == nil {
|
||||
return fmt.Errorf("attach: app is nil")
|
||||
}
|
||||
if bucket == nil {
|
||||
return fmt.Errorf("attach: bucket is nil")
|
||||
}
|
||||
return app.Publish(bucket)
|
||||
}
|
||||
91
modules/lagoon/attach/bucket_test.go
Normal file
91
modules/lagoon/attach/bucket_test.go
Normal file
@@ -0,0 +1,91 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
"gocloud.dev/blob"
|
||||
"gocloud.dev/blob/memblob"
|
||||
)
|
||||
|
||||
func TestOpenBucketRequiresURL(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "storage.yaml"), []byte("uploads:\n bucket_url: \"\"\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{
|
||||
Dir: dir,
|
||||
Env: "development",
|
||||
Environ: []string{"SUMMER_ENV=development"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = OpenBucket(t.Context(), cfg)
|
||||
if err == nil || !strings.Contains(err.Error(), "bucket_url") {
|
||||
t.Fatalf("got %v, want bucket_url error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenBucketOpensMem(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
body := "uploads:\n bucket_url: \"mem://\"\n public_path_prefix: \"/storage/uploads\"\n"
|
||||
if err := os.WriteFile(filepath.Join(dir, "storage.yaml"), []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{
|
||||
Dir: dir,
|
||||
Env: "development",
|
||||
Environ: []string{"SUMMER_ENV=development"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bucket, err := OpenBucket(t.Context(), cfg)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
if PublicPathPrefix() != "/storage/uploads" {
|
||||
t.Fatalf("prefix = %q", PublicPathPrefix())
|
||||
}
|
||||
ctx := t.Context()
|
||||
if err := bucket.WriteAll(ctx, "abc/123/xyz/probe.jpg", []byte("hi"), &blob.WriterOptions{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := bucket.ReadAll(ctx, "abc/123/xyz/probe.jpg")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != "hi" {
|
||||
t.Fatalf("read = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishStoresBucket(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: attach-test\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{
|
||||
Dir: dir,
|
||||
Env: "development",
|
||||
Environ: []string{"SUMMER_ENV=development"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
app := backpackApp(cfg)
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
if err := Publish(app, bucket); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, ok := app.Lookup[*blob.Bucket]()
|
||||
if !ok || got != bucket {
|
||||
t.Fatal("Publish must store the same *blob.Bucket")
|
||||
}
|
||||
}
|
||||
191
modules/lagoon/attach/file.go
Normal file
191
modules/lagoon/attach/file.go
Normal file
@@ -0,0 +1,191 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gocloud.dev/blob"
|
||||
"gocloud.dev/gcerrors"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Owner is implemented by models that own system_files rows. MorphName
|
||||
// must return the PHP class string so cutover-copied attachment_type
|
||||
// values keep matching.
|
||||
type Owner interface {
|
||||
MorphName() string
|
||||
}
|
||||
|
||||
// File is Winter's system_files row. AttachmentID is a string because
|
||||
// Winter stores the morph FK as a string, never an integer.
|
||||
type File struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
DiskName string `gorm:"column:disk_name"`
|
||||
FileName string `gorm:"column:file_name"`
|
||||
FileSize int64 `gorm:"column:file_size"`
|
||||
ContentType string `gorm:"column:content_type"`
|
||||
Title *string `gorm:"column:title"`
|
||||
Description *string `gorm:"column:description"`
|
||||
Field string `gorm:"column:field"`
|
||||
AttachmentID string `gorm:"column:attachment_id"`
|
||||
AttachmentType string `gorm:"column:attachment_type"`
|
||||
// Pointer so GORM can distinguish unset (nil → DEFAULT TRUE) from
|
||||
// explicit false. A non-pointer bool with default:true cannot persist
|
||||
// false because false is the zero value GORM replaces with the default.
|
||||
IsPublic *bool `gorm:"column:is_public;not null;default:true"`
|
||||
SortOrder int `gorm:"column:sort_order"`
|
||||
Metadata *string `gorm:"column:metadata"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
UpdatedAt time.Time `gorm:"column:updated_at"`
|
||||
}
|
||||
|
||||
func (File) TableName() string { return "system_files" }
|
||||
|
||||
// Public reports Winter's is_public flag. A nil pointer is treated as true,
|
||||
// matching the SQL DEFAULT TRUE and File::create() behaviour.
|
||||
func (f File) Public() bool {
|
||||
if f.IsPublic == nil {
|
||||
return true
|
||||
}
|
||||
return *f.IsPublic
|
||||
}
|
||||
|
||||
func fileIsPublic(ctx context.Context, db *gorm.DB, filename string) (bool, error) {
|
||||
if db == nil || filename == "" {
|
||||
return false, nil
|
||||
}
|
||||
var f File
|
||||
q := db.WithContext(ctx).Select("is_public")
|
||||
var err error
|
||||
if id, ok := thumbFileID(filename); ok {
|
||||
err = q.First(&f, id).Error
|
||||
} else {
|
||||
err = q.Where("disk_name = ?", filename).First(&f).Error
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return f.Public(), nil
|
||||
}
|
||||
|
||||
func thumbFileID(name string) (uint, bool) {
|
||||
if !strings.HasPrefix(name, "thumb_") {
|
||||
return 0, false
|
||||
}
|
||||
idStr, _, ok := strings.Cut(strings.TrimPrefix(name, "thumb_"), "_")
|
||||
if !ok || idStr == "" {
|
||||
return 0, false
|
||||
}
|
||||
n, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil || n == 0 {
|
||||
return 0, false
|
||||
}
|
||||
return uint(n), true
|
||||
}
|
||||
|
||||
var all []any
|
||||
|
||||
// Register appends models so schema tooling can see File without a plugin registry.
|
||||
func Register(models ...any) {
|
||||
all = append(all, models...)
|
||||
}
|
||||
|
||||
// All returns every model registered in this package.
|
||||
func All() []any {
|
||||
return all
|
||||
}
|
||||
|
||||
func init() {
|
||||
Register(&File{})
|
||||
}
|
||||
|
||||
func blobKeysFor(f File) []string {
|
||||
part := PartitionDirectory(f.DiskName)
|
||||
return []string{
|
||||
part + f.DiskName,
|
||||
part + fmt.Sprintf("thumb_%d_", f.ID),
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteForOwner removes system_files rows for owner inside tx.
|
||||
//
|
||||
// Two-phase contract (GORM has no post-commit hook):
|
||||
// 1. Inside tx this function SELECTs disk_name values, DELETEs the rows,
|
||||
// and invokes afterCommit with the blob keys so the caller can record
|
||||
// them. afterCommit must not delete blobs — a rollback cannot restore
|
||||
// bytes.
|
||||
// 2. After the top-level Unscoped().Delete(...) returns without error
|
||||
// (the transaction has committed), the caller passes those keys to
|
||||
// DeleteKeys to remove originals and thumbs from the bucket.
|
||||
//
|
||||
// Soft-deleting an owner must not call this helper: rows and blobs stay.
|
||||
func DeleteForOwner(tx *gorm.DB, owner Owner, ownerID string, afterCommit func(blobKeys []string) error) error {
|
||||
if tx == nil {
|
||||
return fmt.Errorf("attach: delete tx is nil")
|
||||
}
|
||||
if owner == nil {
|
||||
return fmt.Errorf("attach: delete owner is nil")
|
||||
}
|
||||
morph := owner.MorphName()
|
||||
var files []File
|
||||
if err := tx.Where("attachment_type = ? AND attachment_id = ?", morph, ownerID).Find(&files).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var keys []string
|
||||
for _, f := range files {
|
||||
keys = append(keys, blobKeysFor(f)...)
|
||||
}
|
||||
if err := tx.Where("attachment_type = ? AND attachment_id = ?", morph, ownerID).Delete(&File{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if afterCommit != nil {
|
||||
return afterCommit(keys)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteKeys removes blob objects after a committed force-delete.
|
||||
// Keys that end in '_' are treated as List prefixes (thumb_<id>_).
|
||||
func DeleteKeys(ctx context.Context, bucket *blob.Bucket, keys []string) error {
|
||||
if bucket == nil {
|
||||
return fmt.Errorf("attach: bucket is nil")
|
||||
}
|
||||
for _, key := range keys {
|
||||
if strings.HasSuffix(key, "_") || strings.HasSuffix(key, "/") {
|
||||
iter := bucket.List(&blob.ListOptions{Prefix: key})
|
||||
for {
|
||||
obj, err := iter.Next(ctx)
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("attach: list %q: %w", key, err)
|
||||
}
|
||||
if err := deleteKey(ctx, bucket, obj.Key); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := deleteKey(ctx, bucket, key); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteKey(ctx context.Context, bucket *blob.Bucket, key string) error {
|
||||
err := bucket.Delete(ctx, key)
|
||||
if err == nil || gcerrors.Code(err) == gcerrors.NotFound {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("attach: delete %q: %w", key, err)
|
||||
}
|
||||
102
modules/lagoon/attach/file_test.go
Normal file
102
modules/lagoon/attach/file_test.go
Normal file
@@ -0,0 +1,102 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gocloud.dev/blob/memblob"
|
||||
)
|
||||
|
||||
func TestFileTableName(t *testing.T) {
|
||||
if got := (File{}).TableName(); got != "system_files" {
|
||||
t.Fatalf("TableName = %q, want system_files", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileIsPublicDefaultsTrue(t *testing.T) {
|
||||
unset := File{}
|
||||
if !unset.Public() {
|
||||
t.Fatal("unset IsPublic must be public")
|
||||
}
|
||||
f, ok := reflect.TypeOf(File{}).FieldByName("IsPublic")
|
||||
if !ok {
|
||||
t.Fatal("missing IsPublic")
|
||||
}
|
||||
tag := f.Tag.Get("gorm")
|
||||
if !strings.Contains(tag, "default:true") {
|
||||
t.Fatalf("gorm tag = %q, want default:true", tag)
|
||||
}
|
||||
priv := false
|
||||
if (File{IsPublic: &priv}).Public() {
|
||||
t.Fatal("explicit false must not be public")
|
||||
}
|
||||
pub := true
|
||||
if !(File{IsPublic: &pub}).Public() {
|
||||
t.Fatal("explicit true must be public")
|
||||
}
|
||||
}
|
||||
|
||||
func TestThumbFileID(t *testing.T) {
|
||||
id, ok := thumbFileID("thumb_42_200_200_0_0_crop.jpg")
|
||||
if !ok || id != 42 {
|
||||
t.Fatalf("thumb id = %d ok=%v, want 42 true", id, ok)
|
||||
}
|
||||
if _, ok := thumbFileID("abc123xyz.jpg"); ok {
|
||||
t.Fatal("original disk_name must not parse as a thumb id")
|
||||
}
|
||||
if _, ok := thumbFileID("thumb_x_200_200_0_0_crop.jpg"); ok {
|
||||
t.Fatal("non-numeric thumb id must be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileSelfRegisters(t *testing.T) {
|
||||
for _, m := range All() {
|
||||
switch m.(type) {
|
||||
case *File, File:
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("File must self-register via attach.Register")
|
||||
}
|
||||
|
||||
// WR-05 pin: the force-delete prefix is thumb_<id>_ with its trailing
|
||||
// underscore, so ID 4 cannot match thumbs of ID 40/41/42 that share a
|
||||
// partition. "thumb_4_" is not a string prefix of "thumb_42_".
|
||||
func TestDeleteKeysThumbPrefixIsIDDelimited(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
defer bucket.Close()
|
||||
const disk = "abc123xyz000.jpg"
|
||||
part := PartitionDirectory(disk)
|
||||
doomed := []string{
|
||||
part + disk,
|
||||
part + ThumbFilename(4, 200, 200, 0, 0, "auto", "jpg"),
|
||||
part + ThumbFilename(4, 50, 50, 0, 0, "crop", "jpg"),
|
||||
}
|
||||
survivors := []string{
|
||||
part + "abc123xyz111.jpg",
|
||||
part + ThumbFilename(40, 200, 200, 0, 0, "auto", "jpg"),
|
||||
part + ThumbFilename(41, 200, 200, 0, 0, "auto", "jpg"),
|
||||
part + ThumbFilename(42, 4, 4, 0, 0, "auto", "jpg"),
|
||||
part + ThumbFilename(14, 200, 200, 0, 0, "auto", "jpg"),
|
||||
}
|
||||
for _, key := range append(append([]string{}, doomed...), survivors...) {
|
||||
if err := bucket.WriteAll(ctx, key, []byte("x"), nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := DeleteKeys(ctx, bucket, blobKeysFor(File{ID: 4, DiskName: disk})); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, key := range doomed {
|
||||
if ok, err := bucket.Exists(ctx, key); err != nil || ok {
|
||||
t.Fatalf("%s must be deleted (exists=%v err=%v)", key, ok, err)
|
||||
}
|
||||
}
|
||||
for _, key := range survivors {
|
||||
if ok, err := bucket.Exists(ctx, key); err != nil || !ok {
|
||||
t.Fatalf("%s must survive deleting file 4 (exists=%v err=%v)", key, ok, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
303
modules/lagoon/attach/lifecycle_test.go
Normal file
303
modules/lagoon/attach/lifecycle_test.go
Normal file
@@ -0,0 +1,303 @@
|
||||
package attach_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/lagoon"
|
||||
"git.golem15.com/golem15/summercms/modules/lagoon/attach"
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"github.com/testcontainers/testcontainers-go"
|
||||
"github.com/testcontainers/testcontainers-go/modules/postgres"
|
||||
"gocloud.dev/blob"
|
||||
"gocloud.dev/blob/memblob"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type lifecycleOwner struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"`
|
||||
}
|
||||
|
||||
func (lifecycleOwner) TableName() string { return "attach_lifecycle_owners" }
|
||||
|
||||
func (lifecycleOwner) MorphName() string {
|
||||
return `Golem15\Fonoteka\Models\Album`
|
||||
}
|
||||
|
||||
func TestFileLifecycle(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("requires testcontainers postgres")
|
||||
}
|
||||
ctx := t.Context()
|
||||
gdb := attachGorm(t)
|
||||
if err := lagoon.Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`
|
||||
CREATE TABLE attach_lifecycle_owners (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
deleted_at TIMESTAMPTZ
|
||||
)`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
owner := lifecycleOwner{Name: "album"}
|
||||
if err := gdb.Create(&owner).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
isPublic := true
|
||||
file := attach.File{
|
||||
DiskName: "abc123xyz.jpg",
|
||||
FileName: "cover.jpg",
|
||||
FileSize: 12,
|
||||
ContentType: "image/jpeg",
|
||||
Field: "photos",
|
||||
AttachmentID: strconv.FormatUint(uint64(owner.ID), 10),
|
||||
AttachmentType: owner.MorphName(),
|
||||
IsPublic: &isPublic,
|
||||
}
|
||||
if err := gdb.Create(&file).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
origKey := attach.BlobKey(file.DiskName)
|
||||
thumbKey := attach.PartitionDirectory(file.DiskName) + attach.ThumbFilename(file.ID, 200, 200, 0, 0, "crop", "jpg")
|
||||
if err := bucket.WriteAll(ctx, origKey, []byte("original"), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := bucket.WriteAll(ctx, thumbKey, []byte("thumb"), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := gdb.Delete(&owner).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var softOwner lifecycleOwner
|
||||
if err := gdb.Unscoped().First(&softOwner, owner.ID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !softOwner.DeletedAt.Valid {
|
||||
t.Fatal("soft-delete must set deleted_at")
|
||||
}
|
||||
assertFileRow(t, gdb, file.ID, true)
|
||||
assertBlob(t, ctx, bucket, origKey, true)
|
||||
assertBlob(t, ctx, bucket, thumbKey, true)
|
||||
|
||||
if err := gdb.Transaction(func(tx *gorm.DB) error {
|
||||
return attach.DeleteForOwner(tx, owner, file.AttachmentID, func(keys []string) error {
|
||||
assertBlob(t, ctx, bucket, origKey, true)
|
||||
assertBlob(t, ctx, bucket, thumbKey, true)
|
||||
return fmt.Errorf("rollback after collecting keys")
|
||||
})
|
||||
}); err == nil {
|
||||
t.Fatal("expected rollback")
|
||||
}
|
||||
assertFileRow(t, gdb, file.ID, true)
|
||||
assertBlob(t, ctx, bucket, origKey, true)
|
||||
assertBlob(t, ctx, bucket, thumbKey, true)
|
||||
|
||||
var pending []string
|
||||
if err := gdb.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Unscoped().Delete(&owner).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return attach.DeleteForOwner(tx, owner, file.AttachmentID, func(keys []string) error {
|
||||
pending = append([]string(nil), keys...)
|
||||
var n int64
|
||||
if err := tx.Model(&attach.File{}).Where("id = ?", file.ID).Count(&n).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if n != 0 {
|
||||
return fmt.Errorf("system_files row must be gone inside the force-delete transaction")
|
||||
}
|
||||
exists, err := bucket.Exists(ctx, origKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !exists {
|
||||
return fmt.Errorf("blob must still exist before commit")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertFileRow(t, gdb, file.ID, false)
|
||||
assertBlob(t, ctx, bucket, origKey, true)
|
||||
if err := attach.DeleteKeys(ctx, bucket, pending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertBlob(t, ctx, bucket, origKey, false)
|
||||
assertBlob(t, ctx, bucket, thumbKey, false)
|
||||
}
|
||||
|
||||
func TestFileCreateDefaultsIsPublic(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("requires testcontainers postgres")
|
||||
}
|
||||
gdb := attachGorm(t)
|
||||
if err := lagoon.Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
file := attach.File{
|
||||
DiskName: "abc123xyz.jpg",
|
||||
FileName: "cover.jpg",
|
||||
FileSize: 1,
|
||||
ContentType: "image/jpeg",
|
||||
}
|
||||
if err := gdb.Create(&file).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var got attach.File
|
||||
if err := gdb.First(&got, file.ID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !got.Public() {
|
||||
t.Fatalf("Create without IsPublic stored %v, want true", got.IsPublic)
|
||||
}
|
||||
|
||||
priv := false
|
||||
hidden := attach.File{
|
||||
DiskName: "def456uvw.jpg",
|
||||
FileName: "secret.jpg",
|
||||
FileSize: 1,
|
||||
ContentType: "image/jpeg",
|
||||
IsPublic: &priv,
|
||||
}
|
||||
if err := gdb.Create(&hidden).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var gotHidden attach.File
|
||||
if err := gdb.First(&gotHidden, hidden.ID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if gotHidden.Public() {
|
||||
t.Fatal("explicit is_public=false must persist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStaticHandlerPublicLooksUpIsPublic(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("requires testcontainers postgres")
|
||||
}
|
||||
ctx := t.Context()
|
||||
gdb := attachGorm(t)
|
||||
if err := lagoon.Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
pub := true
|
||||
priv := false
|
||||
publicFile := attach.File{DiskName: "abc123xyz.jpg", FileName: "cover.jpg", FileSize: 1, ContentType: "image/jpeg", IsPublic: &pub}
|
||||
privateFile := attach.File{DiskName: "def456uvw.jpg", FileName: "secret.jpg", FileSize: 1, ContentType: "image/jpeg", IsPublic: &priv}
|
||||
if err := gdb.Create(&publicFile).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Create(&privateFile).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, f := range []attach.File{publicFile, privateFile} {
|
||||
if err := bucket.WriteAll(ctx, attach.BlobKey(f.DiskName), []byte(f.FileName), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
thumbKey := attach.PartitionDirectory(f.DiskName) + attach.ThumbFilename(f.ID, 50, 50, 0, 0, "crop", "jpg")
|
||||
if err := bucket.WriteAll(ctx, thumbKey, []byte("thumb"), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
h := attach.StaticHandlerPublic(bucket, "/storage/uploads", gdb)
|
||||
get := func(path string) int {
|
||||
t.Helper()
|
||||
rr := httptest.NewRecorder()
|
||||
h.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, path, nil))
|
||||
return rr.Code
|
||||
}
|
||||
if code := get("/storage/uploads/abc/123/xyz/abc123xyz.jpg"); code != http.StatusOK {
|
||||
t.Fatalf("public original status = %d, want 200", code)
|
||||
}
|
||||
if code := get("/storage/uploads/abc/123/xyz/" + attach.ThumbFilename(publicFile.ID, 50, 50, 0, 0, "crop", "jpg")); code != http.StatusOK {
|
||||
t.Fatalf("public thumb status = %d, want 200", code)
|
||||
}
|
||||
if code := get("/storage/uploads/def/456/uvw/def456uvw.jpg"); code != http.StatusNotFound {
|
||||
t.Fatalf("private original status = %d, want 404", code)
|
||||
}
|
||||
if code := get("/storage/uploads/def/456/uvw/" + attach.ThumbFilename(privateFile.ID, 50, 50, 0, 0, "crop", "jpg")); code != http.StatusNotFound {
|
||||
t.Fatalf("private thumb status = %d, want 404", code)
|
||||
}
|
||||
if code := get("/storage/uploads/abc/123/xyz/missing.jpg"); code != http.StatusNotFound {
|
||||
t.Fatalf("unknown disk_name status = %d, want 404", code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertFileRow(t *testing.T, gdb *gorm.DB, id uint, want bool) {
|
||||
t.Helper()
|
||||
var n int64
|
||||
if err := gdb.Model(&attach.File{}).Where("id = ?", id).Count(&n).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := n > 0; got != want {
|
||||
t.Fatalf("system_files id %d exists=%v, want %v", id, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func assertBlob(t *testing.T, ctx context.Context, bucket *blob.Bucket, key string, want bool) {
|
||||
t.Helper()
|
||||
got, err := bucket.Exists(ctx, key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("blob %q exists=%v, want %v", key, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func attachGorm(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
t.Cleanup(cancel)
|
||||
ctr, err := postgres.Run(ctx,
|
||||
"postgres:16-alpine",
|
||||
postgres.WithDatabase("attach"),
|
||||
postgres.WithUsername("attach"),
|
||||
postgres.WithPassword("attach"),
|
||||
postgres.BasicWaitStrategies(),
|
||||
testcontainers.WithEnv(map[string]string{
|
||||
"POSTGRES_INITDB_ARGS": "--locale-provider=icu --icu-locale=pl-PL --encoding=UTF8",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("postgres: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = testcontainers.TerminateContainer(ctr) })
|
||||
dsn, err := ctr.ConnectionString(ctx, "sslmode=disable")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
gdb, err := lagoon.Use(ctx, db)
|
||||
if err != nil {
|
||||
t.Fatalf("lagoon.Use: %v", err)
|
||||
}
|
||||
return gdb
|
||||
}
|
||||
49
modules/lagoon/attach/migrations.go
Normal file
49
modules/lagoon/attach/migrations.go
Normal file
@@ -0,0 +1,49 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"github.com/go-gormigrate/gormigrate/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Migrations creates Winter's system_files table. Folded from
|
||||
// modules/system/database/migrations/2013_10_01_000002_Db_System_Files.php
|
||||
// and 2025_04_10_000031_Db_Add_System_Files_Metadata.php.
|
||||
var Migrations = []*gormigrate.Migration{
|
||||
{
|
||||
ID: "202609180001_create_system_files",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
stmts := []string{
|
||||
`CREATE TABLE system_files (
|
||||
id SERIAL PRIMARY KEY,
|
||||
disk_name TEXT NOT NULL,
|
||||
file_name TEXT NOT NULL,
|
||||
file_size INTEGER NOT NULL,
|
||||
content_type TEXT NOT NULL,
|
||||
title TEXT,
|
||||
description TEXT,
|
||||
field TEXT,
|
||||
attachment_id TEXT,
|
||||
attachment_type TEXT,
|
||||
is_public BOOLEAN NOT NULL DEFAULT TRUE,
|
||||
sort_order INTEGER,
|
||||
metadata TEXT,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
)`,
|
||||
`CREATE INDEX system_files_field_index ON system_files (field)`,
|
||||
`CREATE INDEX system_files_attachment_id_index ON system_files (attachment_id)`,
|
||||
`CREATE INDEX system_files_attachment_type_index ON system_files (attachment_type)`,
|
||||
`CREATE INDEX system_files_attachment_lookup_index ON system_files (attachment_type, attachment_id, field)`,
|
||||
}
|
||||
for _, stmt := range stmts {
|
||||
if err := tx.Exec(stmt).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec("DROP TABLE IF EXISTS system_files").Error
|
||||
},
|
||||
},
|
||||
}
|
||||
125
modules/lagoon/attach/static.go
Normal file
125
modules/lagoon/attach/static.go
Normal file
@@ -0,0 +1,125 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"gocloud.dev/blob"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const defaultStaticContentType = "application/octet-stream"
|
||||
|
||||
// StaticHandler serves GET prefix/<partition>/<filename> from bucket.
|
||||
// filename is the original disk_name or a thumb_* sibling stored in the
|
||||
// original's partition. The blob key is the validated 4-segment path;
|
||||
// unvalidated request segments never reach NewReader (T-05-13).
|
||||
//
|
||||
// This is the public-disk handler: it does not consult system_files.is_public.
|
||||
// Protected files must not be stored in this bucket (Winter uses a second
|
||||
// disk). To 404 is_public=false rows, mount StaticHandlerPublic instead.
|
||||
// Do not mount the ungated handler on the app origin — same-origin
|
||||
// Content-Type from uploads is XSS-relevant.
|
||||
func StaticHandler(bucket *blob.Bucket, prefix string) http.Handler {
|
||||
return servePublicBlobs(bucket, prefix, nil)
|
||||
}
|
||||
|
||||
// StaticHandlerPublic is StaticHandler plus an is_public gate. Missing
|
||||
// rows and is_public=false both 404. The lookup runs before NewReader.
|
||||
func StaticHandlerPublic(bucket *blob.Bucket, prefix string, db *gorm.DB) http.Handler {
|
||||
return servePublicBlobs(bucket, prefix, func(ctx context.Context, filename string) (bool, error) {
|
||||
return fileIsPublic(ctx, db, filename)
|
||||
})
|
||||
}
|
||||
|
||||
func servePublicBlobs(bucket *blob.Bucket, prefix string, allow func(context.Context, string) (bool, error)) http.Handler {
|
||||
prefix = strings.TrimSuffix(prefix, "/")
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodHead {
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
if bucket == nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
rel, ok := stripStaticPrefix(r.URL.Path, prefix)
|
||||
if !ok {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
key, ok := parsePublicBlobPath(rel)
|
||||
if !ok {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
if allow != nil {
|
||||
ok, err := allow(r.Context(), path.Base(key))
|
||||
if err != nil || !ok {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
}
|
||||
reader, err := bucket.NewReader(r.Context(), key, nil)
|
||||
if err != nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
defer reader.Close()
|
||||
ct := reader.ContentType()
|
||||
if ct == "" {
|
||||
ct = defaultStaticContentType
|
||||
}
|
||||
w.Header().Set("Content-Type", ct)
|
||||
if r.Method == http.MethodHead {
|
||||
return
|
||||
}
|
||||
_, _ = io.Copy(w, reader)
|
||||
})
|
||||
}
|
||||
|
||||
func stripStaticPrefix(path, prefix string) (string, bool) {
|
||||
if prefix == "" {
|
||||
return strings.TrimPrefix(path, "/"), true
|
||||
}
|
||||
if path == prefix {
|
||||
return "", false
|
||||
}
|
||||
if strings.HasPrefix(path, prefix+"/") {
|
||||
return path[len(prefix)+1:], true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// parsePublicBlobPath accepts exactly 3 partition groups plus a filename
|
||||
// and returns that path as the blob key. Originals must live in
|
||||
// PartitionDirectory(filename); thumb_* names are stored beside the
|
||||
// original, so their partition is not derived from the thumb filename.
|
||||
// Rejects "..", empty segments, extra slashes, and mismatched original
|
||||
// partitions.
|
||||
func parsePublicBlobPath(p string) (key string, ok bool) {
|
||||
if p == "" || strings.Contains(p, "\\") || strings.Contains(p, "..") || strings.Contains(p, "//") {
|
||||
return "", false
|
||||
}
|
||||
parts := strings.Split(p, "/")
|
||||
if len(parts) != 4 {
|
||||
return "", false
|
||||
}
|
||||
for _, part := range parts {
|
||||
if part == "" || part == "." || part == ".." {
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
filename := parts[3]
|
||||
got := strings.Join(parts[:3], "/")
|
||||
if !strings.HasPrefix(filename, "thumb_") {
|
||||
want := strings.TrimSuffix(PartitionDirectory(filename), "/")
|
||||
if got != want {
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
return got + "/" + filename, true
|
||||
}
|
||||
174
modules/lagoon/attach/static_test.go
Normal file
174
modules/lagoon/attach/static_test.go
Normal file
@@ -0,0 +1,174 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"gocloud.dev/blob"
|
||||
"gocloud.dev/blob/memblob"
|
||||
)
|
||||
|
||||
func TestStaticHandler(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
diskName := "abc123xyz.jpg"
|
||||
key := BlobKey(diskName)
|
||||
body := []byte("cover-bytes")
|
||||
if err := bucket.WriteAll(ctx, key, body, &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
h := StaticHandler(bucket, "/storage/uploads")
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
res, err := http.Get(srv.URL + "/storage/uploads/abc/123/xyz/abc123xyz.jpg")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
t.Fatalf("status = %d", res.StatusCode)
|
||||
}
|
||||
if ct := res.Header.Get("Content-Type"); ct != "image/jpeg" {
|
||||
t.Fatalf("Content-Type = %q, want image/jpeg", ct)
|
||||
}
|
||||
got, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != string(body) {
|
||||
t.Fatalf("body = %q", got)
|
||||
}
|
||||
|
||||
missing, err := http.Get(srv.URL + "/storage/uploads/mis/sin/g.j/missing.jpg")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer missing.Body.Close()
|
||||
if missing.StatusCode != http.StatusNotFound {
|
||||
t.Fatalf("missing status = %d, want 404", missing.StatusCode)
|
||||
}
|
||||
|
||||
for _, path := range []string{
|
||||
"/storage/uploads/abc/123/xyz/../abc123xyz.jpg",
|
||||
"/storage/uploads/abc/123/xyz//abc123xyz.jpg",
|
||||
"/storage/uploads/foo/bar/baz/abc123xyz.jpg",
|
||||
"/storage/uploads/abc/123/xyz/abc123xyz.jpg/extra",
|
||||
"/storage/uploads",
|
||||
} {
|
||||
rr := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, path, nil)
|
||||
h.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusNotFound {
|
||||
t.Fatalf("path %q status = %d, want 404", path, rr.Code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStaticHandlerServesThumbURL(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
f := &File{ID: 42, DiskName: "abc123xyz.jpg"}
|
||||
if err := bucket.WriteAll(ctx, BlobKey(f.DiskName), testJPEG(t), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
thumbURL, err := f.Thumb(ctx, bucket, 200, 200, "crop")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
const wantURL = "/storage/uploads/abc/123/xyz/thumb_42_200_200_0_0_crop.jpg"
|
||||
if thumbURL != wantURL {
|
||||
t.Fatalf("Thumb URL = %q, want %q", thumbURL, wantURL)
|
||||
}
|
||||
|
||||
h := StaticHandler(bucket, "/storage/uploads")
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
res, err := http.Get(srv.URL + thumbURL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
t.Fatalf("thumb GET status = %d, want 200", res.StatusCode)
|
||||
}
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(body) == 0 {
|
||||
t.Fatal("thumb body is empty")
|
||||
}
|
||||
|
||||
mismatch := httptest.NewRecorder()
|
||||
h.ServeHTTP(mismatch, httptest.NewRequest(http.MethodGet, "/storage/uploads/foo/bar/baz/abc123xyz.jpg", nil))
|
||||
if mismatch.Code != http.StatusNotFound {
|
||||
t.Fatalf("mismatched original partition status = %d, want 404", mismatch.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStaticHandlerPublicGate(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
diskName := "abc123xyz.jpg"
|
||||
body := []byte("cover-bytes")
|
||||
if err := bucket.WriteAll(ctx, BlobKey(diskName), body, &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
thumbName := ThumbFilename(42, 200, 200, 0, 0, "crop", "jpg")
|
||||
thumbKey := PartitionDirectory(diskName) + thumbName
|
||||
if err := bucket.WriteAll(ctx, thumbKey, []byte("thumb-bytes"), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var seen []string
|
||||
deny := servePublicBlobs(bucket, "/storage/uploads", func(_ context.Context, name string) (bool, error) {
|
||||
seen = append(seen, name)
|
||||
return false, nil
|
||||
})
|
||||
for _, path := range []string{
|
||||
"/storage/uploads/abc/123/xyz/abc123xyz.jpg",
|
||||
"/storage/uploads/abc/123/xyz/" + thumbName,
|
||||
} {
|
||||
rr := httptest.NewRecorder()
|
||||
deny.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, path, nil))
|
||||
if rr.Code != http.StatusNotFound {
|
||||
t.Fatalf("private %s status = %d, want 404", path, rr.Code)
|
||||
}
|
||||
}
|
||||
if len(seen) != 2 || seen[0] != diskName || seen[1] != thumbName {
|
||||
t.Fatalf("allow saw %v, want [%s %s] before NewReader", seen, diskName, thumbName)
|
||||
}
|
||||
|
||||
allow := servePublicBlobs(bucket, "/storage/uploads", func(_ context.Context, name string) (bool, error) {
|
||||
return name == diskName, nil
|
||||
})
|
||||
ok := httptest.NewRecorder()
|
||||
allow.ServeHTTP(ok, httptest.NewRequest(http.MethodGet, "/storage/uploads/abc/123/xyz/abc123xyz.jpg", nil))
|
||||
if ok.Code != http.StatusOK {
|
||||
t.Fatalf("public original status = %d, want 200", ok.Code)
|
||||
}
|
||||
blocked := httptest.NewRecorder()
|
||||
allow.ServeHTTP(blocked, httptest.NewRequest(http.MethodGet, "/storage/uploads/abc/123/xyz/"+thumbName, nil))
|
||||
if blocked.Code != http.StatusNotFound {
|
||||
t.Fatalf("denied thumb status = %d, want 404", blocked.Code)
|
||||
}
|
||||
|
||||
nilDB := StaticHandlerPublic(bucket, "/storage/uploads", nil)
|
||||
missing := httptest.NewRecorder()
|
||||
nilDB.ServeHTTP(missing, httptest.NewRequest(http.MethodGet, "/storage/uploads/abc/123/xyz/abc123xyz.jpg", nil))
|
||||
if missing.Code != http.StatusNotFound {
|
||||
t.Fatalf("nil db status = %d, want 404", missing.Code)
|
||||
}
|
||||
}
|
||||
177
modules/lagoon/attach/thumb.go
Normal file
177
modules/lagoon/attach/thumb.go
Normal file
@@ -0,0 +1,177 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"image"
|
||||
_ "image/gif"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
"io"
|
||||
"path"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/disintegration/imaging"
|
||||
"gocloud.dev/blob"
|
||||
)
|
||||
|
||||
const (
|
||||
maxThumbEdge = 4096
|
||||
maxThumbSourceBytes = 32 << 20
|
||||
maxThumbSourcePixels = 4096 * 4096
|
||||
)
|
||||
|
||||
// thumbToken is the alphabet allowed for the mode and extension segments of a
|
||||
// thumb filename. Both are interpolated into a blob key, which fileblob maps
|
||||
// to a filesystem path, so separators and dots must never reach it.
|
||||
var thumbToken = regexp.MustCompile(`^[a-z0-9]+$`)
|
||||
|
||||
// ThumbFilename is Winter File::getThumbFilename: thumb_<id>_<w>_<h>_<ox>_<oy>_<mode>.<ext>.
|
||||
// A mode or ext outside [a-z0-9]+ is coerced to "auto" / "jpg" so the result
|
||||
// is always a single safe path element; File.Thumb rejects such input instead.
|
||||
func ThumbFilename(id uint, w, h int, offsetX, offsetY int, mode, ext string) string {
|
||||
if !thumbToken.MatchString(mode) {
|
||||
mode = "auto"
|
||||
}
|
||||
if !thumbToken.MatchString(ext) {
|
||||
ext = "jpg"
|
||||
}
|
||||
return fmt.Sprintf("thumb_%d_%d_%d_%d_%d_%s.%s", id, w, h, offsetX, offsetY, mode, ext)
|
||||
}
|
||||
|
||||
// PartitionDirectory is Winter File::getPartitionDirectory: first 9 chars of
|
||||
// disk_name split into 3 groups of 3, joined by '/', with a trailing slash.
|
||||
func PartitionDirectory(diskName string) string {
|
||||
var groups []string
|
||||
for i := 0; i < len(diskName) && len(groups) < 3; i += 3 {
|
||||
end := i + 3
|
||||
if end > len(diskName) {
|
||||
end = len(diskName)
|
||||
}
|
||||
groups = append(groups, diskName[i:end])
|
||||
}
|
||||
return strings.Join(groups, "/") + "/"
|
||||
}
|
||||
|
||||
// BlobKey is the Winter on-disk key for an original file: partition + disk_name.
|
||||
func BlobKey(diskName string) string {
|
||||
return PartitionDirectory(diskName) + diskName
|
||||
}
|
||||
|
||||
func fileExt(diskName string) string {
|
||||
ext := strings.TrimPrefix(path.Ext(diskName), ".")
|
||||
if ext == "" {
|
||||
return "jpg"
|
||||
}
|
||||
return strings.ToLower(ext)
|
||||
}
|
||||
|
||||
func publicURL(key string) string {
|
||||
prefix := strings.TrimRight(PublicPathPrefix(), "/")
|
||||
key = strings.TrimLeft(key, "/")
|
||||
if prefix == "" {
|
||||
return "/" + key
|
||||
}
|
||||
return prefix + "/" + key
|
||||
}
|
||||
|
||||
func defaultResizeImage(src image.Image, w, h int, mode string) image.Image {
|
||||
switch strings.ToLower(mode) {
|
||||
case "crop":
|
||||
return imaging.Fill(src, w, h, imaging.Center, imaging.Lanczos)
|
||||
case "exact":
|
||||
return imaging.Resize(src, w, h, imaging.Lanczos)
|
||||
default:
|
||||
return imaging.Fit(src, w, h, imaging.Lanczos)
|
||||
}
|
||||
}
|
||||
|
||||
var resizeImage = defaultResizeImage
|
||||
|
||||
func defaultEncodeImage(w io.Writer, img image.Image, ext string) error {
|
||||
format := imaging.JPEG
|
||||
switch strings.ToLower(ext) {
|
||||
case "png":
|
||||
format = imaging.PNG
|
||||
case "gif":
|
||||
format = imaging.GIF
|
||||
}
|
||||
return imaging.Encode(w, img, format)
|
||||
}
|
||||
|
||||
var encodeImage = defaultEncodeImage
|
||||
|
||||
// Thumb returns the public URL of a lazily generated thumbnail. The second
|
||||
// call for the same dimensions hits the existing blob and does not resize.
|
||||
func (f *File) Thumb(ctx context.Context, bucket *blob.Bucket, w, h int, mode string) (string, error) {
|
||||
if f == nil {
|
||||
return "", fmt.Errorf("attach: file is nil")
|
||||
}
|
||||
if bucket == nil {
|
||||
return "", fmt.Errorf("attach: bucket is nil")
|
||||
}
|
||||
if mode == "" {
|
||||
mode = "auto"
|
||||
}
|
||||
mode = strings.ToLower(mode)
|
||||
if !thumbToken.MatchString(mode) {
|
||||
return "", fmt.Errorf("attach: invalid thumb mode %q", mode)
|
||||
}
|
||||
if w <= 0 || h <= 0 || w > maxThumbEdge || h > maxThumbEdge {
|
||||
return "", fmt.Errorf("attach: thumb size %dx%d is out of range", w, h)
|
||||
}
|
||||
ext := fileExt(f.DiskName)
|
||||
if !thumbToken.MatchString(ext) {
|
||||
return "", fmt.Errorf("attach: invalid thumb extension %q", ext)
|
||||
}
|
||||
thumbName := ThumbFilename(f.ID, w, h, 0, 0, mode, ext)
|
||||
part := PartitionDirectory(f.DiskName)
|
||||
thumbKey := part + thumbName
|
||||
exists, err := bucket.Exists(ctx, thumbKey)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("attach: thumb exists: %w", err)
|
||||
}
|
||||
if exists {
|
||||
return publicURL(thumbKey), nil
|
||||
}
|
||||
origKey := part + f.DiskName
|
||||
r, err := bucket.NewReader(ctx, origKey, nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("attach: read original: %w", err)
|
||||
}
|
||||
src, _, err := image.Decode(io.LimitReader(r, maxThumbSourceBytes))
|
||||
closeErr := r.Close()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("attach: decode original: %w", err)
|
||||
}
|
||||
if closeErr != nil {
|
||||
return "", closeErr
|
||||
}
|
||||
bounds := src.Bounds()
|
||||
if int64(bounds.Dx())*int64(bounds.Dy()) > maxThumbSourcePixels {
|
||||
return "", fmt.Errorf("attach: original image is too large")
|
||||
}
|
||||
resized := resizeImage(src, w, h, mode)
|
||||
contentType := "image/jpeg"
|
||||
switch ext {
|
||||
case "png":
|
||||
contentType = "image/png"
|
||||
case "gif":
|
||||
contentType = "image/gif"
|
||||
}
|
||||
wr, err := bucket.NewWriter(ctx, thumbKey, &blob.WriterOptions{ContentType: contentType})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("attach: thumb writer: %w", err)
|
||||
}
|
||||
encErr := encodeImage(wr, resized, ext)
|
||||
closeErr = wr.Close()
|
||||
if encErr != nil || closeErr != nil {
|
||||
_ = deleteKey(ctx, bucket, thumbKey)
|
||||
if encErr != nil {
|
||||
return "", encErr
|
||||
}
|
||||
return "", closeErr
|
||||
}
|
||||
return publicURL(thumbKey), nil
|
||||
}
|
||||
186
modules/lagoon/attach/thumb_test.go
Normal file
186
modules/lagoon/attach/thumb_test.go
Normal file
@@ -0,0 +1,186 @@
|
||||
package attach
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/jpeg"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gocloud.dev/blob"
|
||||
"gocloud.dev/blob/memblob"
|
||||
)
|
||||
|
||||
func TestThumbFilename(t *testing.T) {
|
||||
got := ThumbFilename(42, 200, 200, 0, 0, "crop", "jpg")
|
||||
const want = "thumb_42_200_200_0_0_crop.jpg"
|
||||
if got != want {
|
||||
t.Fatalf("ThumbFilename = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestThumbFilenameRejectsUnsafeTokens(t *testing.T) {
|
||||
for _, tc := range []struct{ mode, ext string }{
|
||||
{"../../secret", "jpg"},
|
||||
{"auto", "../jpg"},
|
||||
{"a/b", "jpg"},
|
||||
{"auto", "j.pg"},
|
||||
{"", ""},
|
||||
{"Crop", "JPG"},
|
||||
} {
|
||||
got := ThumbFilename(42, 200, 200, 0, 0, tc.mode, tc.ext)
|
||||
if strings.ContainsAny(got, `/\`) || strings.Contains(got, "..") {
|
||||
t.Fatalf("mode=%q ext=%q produced unsafe name %q", tc.mode, tc.ext, got)
|
||||
}
|
||||
if !strings.HasPrefix(got, "thumb_42_200_200_0_0_") {
|
||||
t.Fatalf("mode=%q ext=%q name %q lost its prefix", tc.mode, tc.ext, got)
|
||||
}
|
||||
}
|
||||
if got := ThumbFilename(42, 200, 200, 0, 0, "../../secret", "jpg"); got != "thumb_42_200_200_0_0_auto.jpg" {
|
||||
t.Fatalf("unsafe mode coerced to %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileThumbRejectsOutOfRangeSize(t *testing.T) {
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
defer bucket.Close()
|
||||
f := &File{ID: 1, DiskName: "abc123xyz.jpg"}
|
||||
for _, tc := range []struct{ w, h int }{
|
||||
{0, 200},
|
||||
{200, 0},
|
||||
{-1, -1},
|
||||
{maxThumbEdge + 1, 200},
|
||||
{200, maxThumbEdge + 1},
|
||||
{100000, 100000},
|
||||
} {
|
||||
if _, err := f.Thumb(t.Context(), bucket, tc.w, tc.h, "crop"); err == nil {
|
||||
t.Fatalf("size %dx%d must be rejected", tc.w, tc.h)
|
||||
}
|
||||
}
|
||||
iter := bucket.List(nil)
|
||||
if obj, err := iter.Next(t.Context()); err != io.EOF {
|
||||
t.Fatalf("rejected sizes must not touch the bucket, found %v (err %v)", obj, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileThumbRejectsTraversalMode(t *testing.T) {
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
defer bucket.Close()
|
||||
f := &File{ID: 42, DiskName: "abc123xyz789.jpg"}
|
||||
for _, mode := range []string{"../../secret", "a/b", "auto.jpg", "cr op"} {
|
||||
if _, err := f.Thumb(t.Context(), bucket, 200, 200, mode); err == nil {
|
||||
t.Fatalf("mode %q must be rejected", mode)
|
||||
}
|
||||
}
|
||||
iter := bucket.List(nil)
|
||||
if obj, err := iter.Next(t.Context()); err != io.EOF {
|
||||
t.Fatalf("rejected modes must write nothing, found %v (err %v)", obj, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPartitionDirectory(t *testing.T) {
|
||||
got := PartitionDirectory("abc123xyz.jpg")
|
||||
const want = "abc/123/xyz/"
|
||||
if got != want {
|
||||
t.Fatalf("PartitionDirectory = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileThumbResizesOnce(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
f := &File{ID: 42, DiskName: "abc123xyz.jpg", FileName: "cover.jpg"}
|
||||
origKey := PartitionDirectory(f.DiskName) + f.DiskName
|
||||
if origKey == "abc123xyz.jpg" || origKey == "" {
|
||||
// PartitionDirectory still stubbed — write under the Winter key the
|
||||
// GREEN implementation will look up so this test stays the behavior spec.
|
||||
origKey = "abc/123/xyz/" + f.DiskName
|
||||
}
|
||||
if err := bucket.WriteAll(ctx, origKey, testJPEG(t), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var n int
|
||||
orig := resizeImage
|
||||
resizeImage = func(src image.Image, w, h int, mode string) image.Image {
|
||||
n++
|
||||
return orig(src, w, h, mode)
|
||||
}
|
||||
t.Cleanup(func() { resizeImage = orig })
|
||||
|
||||
url1, err := f.Thumb(ctx, bucket, 200, 200, "crop")
|
||||
if err != nil {
|
||||
t.Fatalf("first Thumb: %v", err)
|
||||
}
|
||||
url2, err := f.Thumb(ctx, bucket, 200, 200, "crop")
|
||||
if err != nil {
|
||||
t.Fatalf("second Thumb: %v", err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("resize calls = %d, want 1", n)
|
||||
}
|
||||
want := "/storage/uploads/abc/123/xyz/thumb_42_200_200_0_0_crop.jpg"
|
||||
if url1 != want || url2 != want {
|
||||
t.Fatalf("urls = %q %q, want %q", url1, url2, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileThumbEncodeFailureDoesNotCache(t *testing.T) {
|
||||
ctx := t.Context()
|
||||
bucket := memblob.OpenBucket(nil)
|
||||
t.Cleanup(func() { _ = bucket.Close() })
|
||||
|
||||
f := &File{ID: 7, DiskName: "abc123xyz.jpg"}
|
||||
if err := bucket.WriteAll(ctx, BlobKey(f.DiskName), testJPEG(t), &blob.WriterOptions{ContentType: "image/jpeg"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
orig := encodeImage
|
||||
encodeImage = func(w io.Writer, img image.Image, ext string) error {
|
||||
_, _ = w.Write([]byte("partial"))
|
||||
return errors.New("encode boom")
|
||||
}
|
||||
t.Cleanup(func() { encodeImage = orig })
|
||||
|
||||
if _, err := f.Thumb(ctx, bucket, 200, 200, "crop"); err == nil {
|
||||
t.Fatal("encode failure must surface")
|
||||
}
|
||||
thumbKey := PartitionDirectory(f.DiskName) + ThumbFilename(f.ID, 200, 200, 0, 0, "crop", "jpg")
|
||||
exists, err := bucket.Exists(ctx, thumbKey)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if exists {
|
||||
t.Fatal("failed encode must not leave a cached thumb blob")
|
||||
}
|
||||
|
||||
encodeImage = orig
|
||||
url, err := f.Thumb(ctx, bucket, 200, 200, "crop")
|
||||
if err != nil {
|
||||
t.Fatalf("retry after failed encode: %v", err)
|
||||
}
|
||||
want := "/storage/uploads/abc/123/xyz/thumb_7_200_200_0_0_crop.jpg"
|
||||
if url != want {
|
||||
t.Fatalf("retry url = %q, want %q", url, want)
|
||||
}
|
||||
}
|
||||
|
||||
func testJPEG(t *testing.T) []byte {
|
||||
t.Helper()
|
||||
img := image.NewRGBA(image.Rect(0, 0, 8, 8))
|
||||
for y := 0; y < 8; y++ {
|
||||
for x := 0; x < 8; x++ {
|
||||
img.Set(x, y, color.RGBA{R: 200, A: 255})
|
||||
}
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return buf.Bytes()
|
||||
}
|
||||
88
modules/lagoon/backend_admin_migrations.go
Normal file
88
modules/lagoon/backend_admin_migrations.go
Normal file
@@ -0,0 +1,88 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"github.com/go-gormigrate/gormigrate/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// BackendJWTBlacklistTable is the framework-owned jti table for admin tokens.
|
||||
// It is distinct from the frontend jwt_blacklist table.
|
||||
const BackendJWTBlacklistTable = "backend_jwt_blacklist"
|
||||
|
||||
// BackendAdminMigrations creates Winter-shaped backend identity tables and
|
||||
// seeds the developer and publisher system roles. History is isolated under
|
||||
// the summercms.cabana plugin id. DDL is re-runnable so the system-role seed
|
||||
// stays idempotent if the history row is removed.
|
||||
var BackendAdminMigrations = []*gormigrate.Migration{
|
||||
{
|
||||
ID: "202609240001_backend_admin_identity",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
stmts := []string{
|
||||
`CREATE TABLE IF NOT EXISTS backend_user_roles (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
code TEXT,
|
||||
description TEXT,
|
||||
permissions TEXT,
|
||||
is_system BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
)`,
|
||||
`CREATE TABLE IF NOT EXISTS backend_users (
|
||||
id SERIAL PRIMARY KEY,
|
||||
first_name TEXT,
|
||||
last_name TEXT,
|
||||
login TEXT NOT NULL UNIQUE,
|
||||
email TEXT NOT NULL UNIQUE,
|
||||
password TEXT NOT NULL,
|
||||
activation_code TEXT,
|
||||
persist_code TEXT,
|
||||
reset_password_code TEXT,
|
||||
permissions TEXT,
|
||||
is_activated BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
is_superuser BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
role_id INTEGER REFERENCES backend_user_roles(id),
|
||||
activated_at TIMESTAMPTZ,
|
||||
last_login TIMESTAMPTZ,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
deleted_at TIMESTAMPTZ,
|
||||
tokens_valid_after TIMESTAMPTZ
|
||||
)`,
|
||||
`CREATE TABLE IF NOT EXISTS backend_jwt_blacklist (
|
||||
jti TEXT PRIMARY KEY,
|
||||
expires_at TIMESTAMPTZ NOT NULL,
|
||||
valid_until TIMESTAMPTZ NOT NULL
|
||||
)`,
|
||||
`CREATE INDEX IF NOT EXISTS backend_users_role_id_index ON backend_users (role_id)`,
|
||||
`CREATE INDEX IF NOT EXISTS backend_users_deleted_at_index ON backend_users (deleted_at)`,
|
||||
`CREATE INDEX IF NOT EXISTS backend_users_activation_code_index ON backend_users (activation_code)`,
|
||||
`CREATE INDEX IF NOT EXISTS backend_users_reset_password_code_index ON backend_users (reset_password_code)`,
|
||||
`CREATE INDEX IF NOT EXISTS backend_user_roles_code_index ON backend_user_roles (code)`,
|
||||
`INSERT INTO backend_user_roles (name, code, description, permissions, is_system)
|
||||
VALUES
|
||||
('Developer', 'developer', 'Site administrator with access to developer tools.', '{}', TRUE),
|
||||
('Publisher', 'publisher', 'Site editor with access to publishing tools.', '{}', TRUE)
|
||||
ON CONFLICT (name) DO NOTHING`,
|
||||
}
|
||||
for _, stmt := range stmts {
|
||||
if err := tx.Exec(stmt).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
for _, stmt := range []string{
|
||||
`DROP TABLE IF EXISTS backend_jwt_blacklist`,
|
||||
`DROP TABLE IF EXISTS backend_users`,
|
||||
`DROP TABLE IF EXISTS backend_user_roles`,
|
||||
} {
|
||||
if err := tx.Exec(stmt).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
},
|
||||
}
|
||||
445
modules/lagoon/backend_admin_migrations_test.go
Normal file
445
modules/lagoon/backend_admin_migrations_test.go
Normal file
@@ -0,0 +1,445 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/party"
|
||||
"github.com/go-gormigrate/gormigrate/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestPhase09MigrationsFreshRollback(t *testing.T) {
|
||||
assertNoAutoMigrate(t)
|
||||
db, _ := dedicatedDB(t, "phase09_fresh")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
plugin := migPlugin{
|
||||
id: "demo.keep",
|
||||
migrations: []*gormigrate.Migration{{
|
||||
ID: "202609240010_keep",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`CREATE TABLE phase09_keep (id BIGSERIAL PRIMARY KEY)`).Error
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`DROP TABLE IF EXISTS phase09_keep`).Error
|
||||
},
|
||||
}},
|
||||
}
|
||||
if err := Migrate(gdb, []party.Plugin{plugin}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("system_files") || !gdb.Migrator().HasTable("backend_users") || !gdb.Migrator().HasTable("phase09_keep") {
|
||||
t.Fatal("fresh migrate did not keep framework, admin, and plugin tables together")
|
||||
}
|
||||
before := systemRoles(t, gdb)
|
||||
if len(before) != 2 {
|
||||
t.Fatalf("seeded roles = %v", before)
|
||||
}
|
||||
if err := Migrate(gdb, []party.Plugin{plugin}); err != nil {
|
||||
t.Fatalf("repeated migrate: %v", err)
|
||||
}
|
||||
if after := systemRoles(t, gdb); len(after) != 2 || after["developer"].id != before["developer"].id || after["publisher"].id != before["publisher"].id {
|
||||
t.Fatalf("roles changed on repeat: before=%v after=%v", before, after)
|
||||
}
|
||||
admin, err := migrator(gdb, "summercms.cabana", BackendAdminMigrations)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := admin.RollbackLast(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, name := range []string{"backend_users", "backend_user_roles", "backend_jwt_blacklist"} {
|
||||
if gdb.Migrator().HasTable(name) {
|
||||
t.Fatalf("%s survived admin rollback", name)
|
||||
}
|
||||
}
|
||||
if !gdb.Migrator().HasTable("system_files") || !gdb.Migrator().HasTable("phase09_keep") {
|
||||
t.Fatal("admin rollback removed a framework or plugin table")
|
||||
}
|
||||
var attachIDs, pluginIDs []string
|
||||
if err := gdb.Table("summer_migrations_summercms_attach").Pluck("id", &attachIDs).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Table("summer_migrations_demo_keep").Pluck("id", &pluginIDs).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Join(attachIDs, ",") == "" || strings.Join(pluginIDs, ",") != "202609240010_keep" {
|
||||
t.Fatalf("histories attach=%v plugin=%v", attachIDs, pluginIDs)
|
||||
}
|
||||
if err := Migrate(gdb, []party.Plugin{plugin}); err != nil {
|
||||
t.Fatalf("migrate after rollback: %v", err)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("backend_users") {
|
||||
t.Fatal("admin tables were not recreated")
|
||||
}
|
||||
if again := systemRoles(t, gdb); len(again) != 2 {
|
||||
t.Fatalf("roles after recreate = %v", again)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackendAdminMigration(t *testing.T) {
|
||||
assertNoAutoMigrate(t)
|
||||
db, _ := dedicatedDB(t, "lagoon_admin_mig")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("system_files") {
|
||||
t.Fatal("framework attachment migration must still run before backend admin tables")
|
||||
}
|
||||
users := columnNullability(t, gdb, "backend_users")
|
||||
wantUsers := map[string]bool{
|
||||
"id": false, "first_name": true, "last_name": true, "login": false, "email": false,
|
||||
"password": false, "activation_code": true, "persist_code": true, "reset_password_code": true,
|
||||
"permissions": true, "is_activated": false, "is_superuser": false, "role_id": true,
|
||||
"activated_at": true, "last_login": true, "created_at": false, "updated_at": false,
|
||||
"deleted_at": true, "tokens_valid_after": true,
|
||||
}
|
||||
assertColumns(t, "backend_users", users, wantUsers)
|
||||
roles := columnNullability(t, gdb, "backend_user_roles")
|
||||
wantRoles := map[string]bool{
|
||||
"id": false, "name": false, "code": true, "description": true, "permissions": true,
|
||||
"is_system": false, "created_at": false, "updated_at": false,
|
||||
}
|
||||
assertColumns(t, "backend_user_roles", roles, wantRoles)
|
||||
bl := columnNullability(t, gdb, "backend_jwt_blacklist")
|
||||
assertColumns(t, "backend_jwt_blacklist", bl, map[string]bool{
|
||||
"jti": false, "expires_at": false, "valid_until": false,
|
||||
})
|
||||
for _, name := range []string{"backend_user_groups", "backend_users_groups", "backend_user_preferences", "backend_access_log"} {
|
||||
if gdb.Migrator().HasTable(name) {
|
||||
t.Fatalf("prohibited table %s exists", name)
|
||||
}
|
||||
}
|
||||
for _, col := range []string{"activation_code", "reset_password_code", "role_id", "deleted_at", "login", "email"} {
|
||||
if !indexOn(indexDefs(t, gdb, "backend_users"), col) {
|
||||
t.Fatalf("backend_users missing index on %s", col)
|
||||
}
|
||||
}
|
||||
if !indexOn(indexDefs(t, gdb, "backend_user_roles"), "code") {
|
||||
t.Fatal("backend_user_roles missing index on code")
|
||||
}
|
||||
var fk string
|
||||
if err := gdb.Raw(`SELECT pg_get_constraintdef(oid) FROM pg_constraint WHERE conrelid = 'backend_users'::regclass AND contype = 'f'`).Scan(&fk).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(fk, "backend_user_roles") {
|
||||
t.Fatalf("role relationship = %q", fk)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_users (login, email, password, created_at, updated_at) VALUES ('ada', 'ada@example.test', 'x', NOW(), NOW())`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_users (login, email, password, created_at, updated_at) VALUES ('ada', 'other@example.test', 'x', NOW(), NOW())`).Error; err == nil {
|
||||
t.Fatal("duplicate login was accepted")
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_users (login, email, password, created_at, updated_at) VALUES ('ada-2', 'ada@example.test', 'x', NOW(), NOW())`).Error; err == nil {
|
||||
t.Fatal("duplicate email was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackendAdminSeed(t *testing.T) {
|
||||
db, _ := dedicatedDB(t, "lagoon_admin_seed")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatalf("repeated migrate: %v", err)
|
||||
}
|
||||
before := systemRoles(t, gdb)
|
||||
if err := gdb.Exec(`DELETE FROM summer_migrations_summercms_cabana`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatalf("reapplying the admin migration must leave the developer and publisher seeds idempotent: %v", err)
|
||||
}
|
||||
after := systemRoles(t, gdb)
|
||||
if len(before) != 2 || len(after) != 2 {
|
||||
t.Fatalf("roles before=%v after=%v", before, after)
|
||||
}
|
||||
for code, row := range before {
|
||||
next, ok := after[code]
|
||||
if !ok || next.id != row.id || !next.system {
|
||||
t.Fatalf("role %s changed: before=%+v after=%+v", code, row, next)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackendAdminRollback(t *testing.T) {
|
||||
db, _ := dedicatedDB(t, "lagoon_admin_rollback")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
plugin := migPlugin{
|
||||
id: "demo.keep",
|
||||
migrations: []*gormigrate.Migration{{
|
||||
ID: "202609240010_keep",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`CREATE TABLE lagoon_keep (id BIGSERIAL PRIMARY KEY)`).Error
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`DROP TABLE IF EXISTS lagoon_keep`).Error
|
||||
},
|
||||
}},
|
||||
}
|
||||
if err := Migrate(gdb, []party.Plugin{plugin}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("backend_jwt_blacklist") {
|
||||
t.Fatal("blacklist table missing before rollback")
|
||||
}
|
||||
m, err := migrator(gdb, "summercms.cabana", BackendAdminMigrations)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := m.RollbackLast(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, name := range []string{"backend_users", "backend_user_roles", "backend_jwt_blacklist"} {
|
||||
if gdb.Migrator().HasTable(name) {
|
||||
t.Fatalf("%s survived admin rollback", name)
|
||||
}
|
||||
}
|
||||
if !gdb.Migrator().HasTable("system_files") || !gdb.Migrator().HasTable("lagoon_keep") {
|
||||
t.Fatal("admin rollback removed another framework or plugin table")
|
||||
}
|
||||
var attachIDs, pluginIDs []string
|
||||
if err := gdb.Table("summer_migrations_summercms_attach").Pluck("id", &attachIDs).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Table("summer_migrations_demo_keep").Pluck("id", &pluginIDs).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Join(attachIDs, ",") != "202609180001_create_system_files" {
|
||||
t.Fatalf("attach history = %v", attachIDs)
|
||||
}
|
||||
if strings.Join(pluginIDs, ",") != "202609240010_keep" {
|
||||
t.Fatalf("plugin history = %v", pluginIDs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackendAdminWinterRow(t *testing.T) {
|
||||
db, _ := dedicatedDB(t, "lagoon_admin_winter")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var role winterRole
|
||||
if err := gdb.Where("code = ?", "developer").First(&role).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if role.Name != "Developer" || !role.IsSystem {
|
||||
t.Fatalf("developer role = %+v", role)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_users
|
||||
(first_name, last_name, login, email, password, activation_code, persist_code, reset_password_code, permissions, is_activated, is_superuser, role_id, activated_at, last_login, created_at, updated_at)
|
||||
VALUES ('Ada', 'Lovelace', 'ada', 'Ada@Example.Test', 'winter-hash', 'act', 'persist', 'reset', '{"backend.manage_access":1}', TRUE, FALSE, NULL, NULL, NULL, NOW(), NOW())`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var user winterUser
|
||||
if err := gdb.Preload("Role").Where("login = ?", "ada").First(&user).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if user.FirstName != "Ada" || user.LastName != "Lovelace" || user.Email != "Ada@Example.Test" || user.Password != "winter-hash" || user.RoleID != nil || user.LastLogin != nil || !user.IsActivated || user.IsSuperuser {
|
||||
t.Fatalf("winter row = %+v", user)
|
||||
}
|
||||
user.FirstName = "Augusta"
|
||||
if err := gdb.Save(&user).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var perms string
|
||||
if err := gdb.Raw(`SELECT permissions FROM backend_users WHERE login = 'ada'`).Scan(&perms).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if perms != `{"backend.manage_access":1}` {
|
||||
t.Fatalf("permissions rewritten to %q", perms)
|
||||
}
|
||||
if err := gdb.Delete(&user).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Where("login = ?", "ada").First(&winterUser{}).Error; err == nil {
|
||||
t.Fatal("soft-deleted backend user remained visible")
|
||||
}
|
||||
if err := gdb.Unscoped().Where("login = ?", "ada").First(&winterUser{}).Error; err != nil {
|
||||
t.Fatalf("unscoped load: %v", err)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_user_roles (name, code, is_system, created_at, updated_at) VALUES ('Editor A', 'shared', FALSE, NOW(), NOW())`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO backend_user_roles (name, code, is_system, created_at, updated_at) VALUES ('Editor B', 'shared', FALSE, NOW(), NOW())`).Error; err != nil {
|
||||
t.Fatalf("winter-shaped roles must allow a repeated code so cutover rows load without a schema transform: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// winterUser mirrors cabana.BackendUser's column contract without importing
|
||||
// cabana. Cabana commands call lagoon.OpenFromApp, so a lagoon test cannot
|
||||
// import cabana.
|
||||
type winterRole struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
Code string `gorm:"column:code"`
|
||||
IsSystem bool `gorm:"column:is_system"`
|
||||
}
|
||||
|
||||
func (winterRole) TableName() string { return "backend_user_roles" }
|
||||
|
||||
type winterUser struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
FirstName string `gorm:"column:first_name"`
|
||||
LastName string `gorm:"column:last_name"`
|
||||
Login string `gorm:"column:login"`
|
||||
Email string `gorm:"column:email"`
|
||||
Password string `gorm:"column:password"`
|
||||
Permissions string `gorm:"column:permissions"`
|
||||
IsActivated bool `gorm:"column:is_activated"`
|
||||
IsSuperuser bool `gorm:"column:is_superuser"`
|
||||
RoleID *uint `gorm:"column:role_id"`
|
||||
LastLogin *time.Time `gorm:"column:last_login"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"`
|
||||
TokensValidAfter *time.Time `gorm:"column:tokens_valid_after"`
|
||||
Role winterRole
|
||||
}
|
||||
|
||||
func (winterUser) TableName() string { return "backend_users" }
|
||||
|
||||
type systemRole struct {
|
||||
id int
|
||||
system bool
|
||||
}
|
||||
|
||||
func systemRoles(t *testing.T, gdb *gorm.DB) map[string]systemRole {
|
||||
t.Helper()
|
||||
type row struct {
|
||||
ID int
|
||||
Code string
|
||||
System bool
|
||||
Name string
|
||||
Describe string
|
||||
}
|
||||
var rows []row
|
||||
if err := gdb.Raw(`SELECT id, code, is_system AS system, name, description AS describe FROM backend_user_roles WHERE code IN ('developer', 'publisher') ORDER BY code`).Scan(&rows).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
out := map[string]systemRole{}
|
||||
for _, row := range rows {
|
||||
if !row.System {
|
||||
t.Fatalf("%s is not a system role", row.Code)
|
||||
}
|
||||
switch row.Code {
|
||||
case "developer":
|
||||
if row.Name != "Developer" || row.Describe != "Site administrator with access to developer tools." {
|
||||
t.Fatalf("developer seed = %+v", row)
|
||||
}
|
||||
case "publisher":
|
||||
if row.Name != "Publisher" || row.Describe != "Site editor with access to publishing tools." {
|
||||
t.Fatalf("publisher seed = %+v", row)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("unexpected role %s", row.Code)
|
||||
}
|
||||
out[row.Code] = systemRole{id: row.ID, system: row.System}
|
||||
}
|
||||
if _, ok := out["developer"]; !ok {
|
||||
t.Fatal("developer seed missing")
|
||||
}
|
||||
if _, ok := out["publisher"]; !ok {
|
||||
t.Fatal("publisher seed missing")
|
||||
}
|
||||
var n int64
|
||||
if err := gdb.Raw(`SELECT COUNT(*) FROM backend_user_roles WHERE is_system`).Scan(&n).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("system roles = %d", n)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func columnNullability(t *testing.T, gdb *gorm.DB, table string) map[string]bool {
|
||||
t.Helper()
|
||||
type col struct {
|
||||
Name string
|
||||
Nullable string
|
||||
}
|
||||
var cols []col
|
||||
if err := gdb.Raw(`SELECT column_name AS name, is_nullable AS nullable FROM information_schema.columns WHERE table_schema = 'public' AND table_name = ?`, table).Scan(&cols).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(cols) == 0 {
|
||||
t.Fatalf("table %s has no columns", table)
|
||||
}
|
||||
out := make(map[string]bool, len(cols))
|
||||
for _, c := range cols {
|
||||
out[c.Name] = c.Nullable == "YES"
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func assertColumns(t *testing.T, table string, got, want map[string]bool) {
|
||||
t.Helper()
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("%s columns = %v, want %v", table, got, want)
|
||||
}
|
||||
for name, nullable := range want {
|
||||
gotNullable, ok := got[name]
|
||||
if !ok || gotNullable != nullable {
|
||||
t.Fatalf("%s.%s nullable=%v present=%t, want nullable=%t", table, name, gotNullable, ok, nullable)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func indexDefs(t *testing.T, gdb *gorm.DB, table string) []string {
|
||||
t.Helper()
|
||||
var defs []string
|
||||
if err := gdb.Raw(`SELECT indexdef FROM pg_indexes WHERE schemaname = 'public' AND tablename = ?`, table).Scan(&defs).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return defs
|
||||
}
|
||||
|
||||
func indexOn(defs []string, column string) bool {
|
||||
needle := "(" + column + ")"
|
||||
for _, def := range defs {
|
||||
if strings.Contains(def, needle) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func assertNoAutoMigrate(t *testing.T) {
|
||||
t.Helper()
|
||||
entries, err := os.ReadDir(".")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, entry := range entries {
|
||||
name := entry.Name()
|
||||
if entry.IsDir() || !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") {
|
||||
continue
|
||||
}
|
||||
body, err := os.ReadFile(filepath.Join(".", name))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(string(body), "AutoMigrate") {
|
||||
t.Fatalf("%s uses AutoMigrate", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
88
modules/lagoon/commands.go
Normal file
88
modules/lagoon/commands.go
Normal file
@@ -0,0 +1,88 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/bonfire"
|
||||
"git.golem15.com/golem15/summercms/modules/party"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// RuntimeCommands returns migrate, migrate:rollback, migrate:status, and key:generate.
|
||||
// Serve is registered separately via surf.ServeCommand.
|
||||
func RuntimeCommands(app *backpack.App, plugins []party.Plugin) []bonfire.Command {
|
||||
return []bonfire.Command{
|
||||
{
|
||||
Name: "migrate",
|
||||
Description: "Run plugin migrations in dependency order",
|
||||
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
|
||||
return withDB(ctx, app, func(gdb *gorm.DB) error {
|
||||
if err := Migrate(gdb, plugins); err != nil {
|
||||
return err
|
||||
}
|
||||
out.Success("migrations applied")
|
||||
return nil
|
||||
})
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "migrate:rollback",
|
||||
Description: "Roll back the last migration of a plugin",
|
||||
Flags: []bonfire.Flag{{
|
||||
Name: "plugin",
|
||||
Description: "Plugin ID whose last migration to roll back",
|
||||
}},
|
||||
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
|
||||
plugin, _ := in.Flag("plugin")
|
||||
return withDB(ctx, app, func(gdb *gorm.DB) error {
|
||||
if err := RollbackLast(gdb, plugins, plugin); err != nil {
|
||||
return err
|
||||
}
|
||||
if plugin == "" {
|
||||
plugin = lastMigrationPlugin(plugins)
|
||||
}
|
||||
out.Success(fmt.Sprintf("rolled back last migration of %s", plugin))
|
||||
return nil
|
||||
})
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "migrate:status",
|
||||
Description: "Show per-plugin migration history",
|
||||
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
|
||||
return withDB(ctx, app, func(gdb *gorm.DB) error {
|
||||
rows, err := Status(gdb, plugins)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tableRows := make([][]string, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
ids := "(none)"
|
||||
if len(row.IDs) > 0 {
|
||||
ids = strings.Join(row.IDs, ", ")
|
||||
}
|
||||
tableRows = append(tableRows, []string{row.Plugin, row.Table, ids})
|
||||
}
|
||||
out.Table([]string{"plugin", "table", "applied"}, tableRows)
|
||||
return nil
|
||||
})
|
||||
},
|
||||
},
|
||||
KeyGenerateCommand(),
|
||||
}
|
||||
}
|
||||
|
||||
func withDB(ctx context.Context, app *backpack.App, fn func(*gorm.DB) error) error {
|
||||
sqlDB, gdb, err := OpenFromApp(ctx, app)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer sqlDB.Close()
|
||||
if err := Publish(app, sqlDB, gdb); err != nil {
|
||||
return err
|
||||
}
|
||||
return fn(gdb)
|
||||
}
|
||||
146
modules/lagoon/connection.go
Normal file
146
modules/lagoon/connection.go
Normal file
@@ -0,0 +1,146 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
requiredLocaleProvider = "i"
|
||||
requiredICULocale = "pl-PL"
|
||||
)
|
||||
|
||||
// Open pings dsn through pgx stdlib, requires Postgres 16 ICU pl-PL, and
|
||||
// returns that exact *sql.DB plus a GORM handle opened on it.
|
||||
//
|
||||
// Phase 11 owns a separate pgxpool.Pool for River LISTEN/NOTIFY. Do not
|
||||
// create that listener pool here; application queries share this *sql.DB.
|
||||
func Open(ctx context.Context, dsn string) (*sql.DB, *gorm.DB, error) {
|
||||
dsn = strings.TrimSpace(dsn)
|
||||
if dsn == "" {
|
||||
return nil, nil, fmt.Errorf("lagoon: database.dsn is empty (set SUMMER_DATABASE__DSN)")
|
||||
}
|
||||
sqlDB, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("lagoon: open postgres: %w", err)
|
||||
}
|
||||
if err := sqlDB.PingContext(ctx); err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, nil, fmt.Errorf("lagoon: ping postgres: %w", err)
|
||||
}
|
||||
if err := CheckLocale(ctx, sqlDB); err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
gdb, err := gormFromSQL(sqlDB)
|
||||
if err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
return sqlDB, gdb, nil
|
||||
}
|
||||
|
||||
// Use pings an existing pool, requires ICU pl-PL, and returns a GORM handle
|
||||
// opened on that exact *sql.DB. Callers that already hold a pool (tests,
|
||||
// the app boot seam) must not open a second connection.
|
||||
func Use(ctx context.Context, sqlDB *sql.DB) (*gorm.DB, error) {
|
||||
if sqlDB == nil {
|
||||
return nil, fmt.Errorf("lagoon: sql db is nil")
|
||||
}
|
||||
if err := sqlDB.PingContext(ctx); err != nil {
|
||||
return nil, fmt.Errorf("lagoon: ping postgres: %w", err)
|
||||
}
|
||||
if err := CheckLocale(ctx, sqlDB); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return gormFromSQL(sqlDB)
|
||||
}
|
||||
|
||||
func gormFromSQL(sqlDB *sql.DB) (*gorm.DB, error) {
|
||||
gdb, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: gorm open: %w", err)
|
||||
}
|
||||
got, err := gdb.DB()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: gorm sql handle: %w", err)
|
||||
}
|
||||
if got != sqlDB {
|
||||
return nil, fmt.Errorf("lagoon: GORM is not using the shared *sql.DB")
|
||||
}
|
||||
return gdb, nil
|
||||
}
|
||||
|
||||
// OpenFromApp reads database.dsn from app config and opens the shared pool.
|
||||
// It also loads app.key via LoadAppKey and PublishEncryptionKeys so Encrypted
|
||||
// columns do not re-read config on every row.
|
||||
func OpenFromApp(ctx context.Context, app *backpack.App) (*sql.DB, *gorm.DB, error) {
|
||||
if app == nil || app.Config == nil {
|
||||
return nil, nil, fmt.Errorf("lagoon: app config is missing")
|
||||
}
|
||||
sqlDB, gdb, err := Open(ctx, DSN(app.Config))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if err := loadEncryptionKeysFromApp(app); err != nil {
|
||||
_ = sqlDB.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
return sqlDB, gdb, nil
|
||||
}
|
||||
|
||||
// DSN returns database.dsn from layered config (env SUMMER_DATABASE__DSN).
|
||||
func DSN(cfg *compass.Config) string {
|
||||
if cfg == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(cfg.String("database.dsn"))
|
||||
}
|
||||
|
||||
// Publish stores the shared SQL pool and GORM handle on the app. Both
|
||||
// handles refer to the same *sql.DB.
|
||||
func Publish(app *backpack.App, sqlDB *sql.DB, gdb *gorm.DB) error {
|
||||
if app == nil {
|
||||
return fmt.Errorf("lagoon: app is nil")
|
||||
}
|
||||
if sqlDB == nil || gdb == nil {
|
||||
return fmt.Errorf("lagoon: database handles are nil")
|
||||
}
|
||||
if err := app.Publish(sqlDB); err != nil {
|
||||
return err
|
||||
}
|
||||
return app.Publish(gdb)
|
||||
}
|
||||
|
||||
// CheckLocale fails unless the connected database uses ICU locale pl-PL.
|
||||
func CheckLocale(ctx context.Context, db *sql.DB) error {
|
||||
if db == nil {
|
||||
return fmt.Errorf("lagoon: sql db is nil")
|
||||
}
|
||||
var provider, icu string
|
||||
err := db.QueryRowContext(ctx, `
|
||||
SELECT datlocprovider::text, COALESCE(daticulocale, '')
|
||||
FROM pg_database
|
||||
WHERE datname = current_database()`).Scan(&provider, &icu)
|
||||
if err != nil {
|
||||
return fmt.Errorf("lagoon: read database locale: %w", err)
|
||||
}
|
||||
return checkLocale(provider, icu)
|
||||
}
|
||||
|
||||
func checkLocale(provider, icu string) error {
|
||||
provider = strings.TrimSpace(provider)
|
||||
icu = strings.TrimSpace(icu)
|
||||
if provider == requiredLocaleProvider && icu == requiredICULocale {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("lagoon: database locale must be ICU pl-PL (datlocprovider=%q, daticulocale=%q); got provider %q locale %q. Create the database with: CREATE DATABASE ... TEMPLATE template0 ENCODING 'UTF8' LOCALE_PROVIDER icu ICU_LOCALE 'pl-PL'", requiredLocaleProvider, requiredICULocale, provider, icu)
|
||||
}
|
||||
165
modules/lagoon/connection_test.go
Normal file
165
modules/lagoon/connection_test.go
Normal file
@@ -0,0 +1,165 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestUseRejectsNilSQL(t *testing.T) {
|
||||
_, err := Use(t.Context(), nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "sql db is nil") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckLocaleRejectsNilSQL(t *testing.T) {
|
||||
err := CheckLocale(t.Context(), nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "sql db is nil") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublishRejectsNilHandles(t *testing.T) {
|
||||
app := backpack.New(nil)
|
||||
if err := Publish(nil, nil, nil); err == nil {
|
||||
t.Fatal("want nil app error")
|
||||
}
|
||||
if err := Publish(app, nil, &gorm.DB{}); err == nil || !strings.Contains(err.Error(), "nil") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDSNEmptyWithoutConfig(t *testing.T) {
|
||||
if DSN(nil) != "" {
|
||||
t.Fatal("nil config must yield empty DSN")
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenFromAppMissingConfig(t *testing.T) {
|
||||
_, _, err := OpenFromApp(t.Context(), nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "app config is missing") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
_, _, err = OpenFromApp(t.Context(), backpack.New(nil))
|
||||
if err == nil || !strings.Contains(err.Error(), "app config is missing") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSharedSQLPoolUsedByGORMAndClosed(t *testing.T) {
|
||||
_, dsn := dedicatedDB(t, "lagoon_shared")
|
||||
sqlDB, gdb, err := Open(t.Context(), dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, err := gdb.DB()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if raw != sqlDB {
|
||||
t.Fatal("GORM is not using the shared *sql.DB")
|
||||
}
|
||||
|
||||
app := backpack.New(nil)
|
||||
if err := Publish(app, sqlDB, gdb); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
gotSQL, ok := app.Lookup[*sql.DB]()
|
||||
if !ok || gotSQL != sqlDB {
|
||||
t.Fatal("published *sql.DB is not the shared pool")
|
||||
}
|
||||
gotGORM, ok := app.Lookup[*gorm.DB]()
|
||||
if !ok || gotGORM != gdb {
|
||||
t.Fatal("published *gorm.DB is missing")
|
||||
}
|
||||
|
||||
viaUse, err := Use(t.Context(), sqlDB)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
useRaw, err := viaUse.DB()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if useRaw != sqlDB {
|
||||
t.Fatal("Use must open GORM on the caller pool, not a second connection")
|
||||
}
|
||||
|
||||
if err := sqlDB.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := sqlDB.PingContext(t.Context()); err == nil {
|
||||
t.Fatal("closed pool must reject Ping")
|
||||
}
|
||||
if _, err := Use(t.Context(), sqlDB); err == nil {
|
||||
t.Fatal("Use must fail after the shared pool is closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrongICULocaleFailsOpen(t *testing.T) {
|
||||
admin := lagoonDB(t)
|
||||
ctx := t.Context()
|
||||
if _, err := admin.ExecContext(ctx, `CREATE DATABASE lagoon_locale_fail TEMPLATE template0 ENCODING 'UTF8' LOCALE_PROVIDER libc LOCALE 'C'`); err != nil && !strings.Contains(err.Error(), "already exists") {
|
||||
t.Fatalf("create libc database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = admin.ExecContext(ctx, `DROP DATABASE IF EXISTS lagoon_locale_fail WITH (FORCE)`)
|
||||
})
|
||||
failDSN, err := dsnWithDB(lagoonDSN, "lagoon_locale_fail")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, _, err = Open(ctx, failDSN)
|
||||
if err == nil {
|
||||
t.Fatal("wrong locale must fail boot")
|
||||
}
|
||||
msg := err.Error()
|
||||
for _, want := range []string{"ICU", "pl-PL", "CREATE DATABASE", "LOCALE_PROVIDER icu"} {
|
||||
if !strings.Contains(msg, want) {
|
||||
t.Fatalf("missing %q in %s", want, msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenFromAppReadsDSN(t *testing.T) {
|
||||
_, dsn := dedicatedDB(t, "lagoon_from_app")
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: lagoon-test\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{
|
||||
Dir: dir,
|
||||
Environ: []string{
|
||||
"SUMMER_ENV=development",
|
||||
"SUMMER_DATABASE__DSN=" + dsn,
|
||||
"SUMMER_APP__KEY=" + base64.StdEncoding.EncodeToString(bytes.Repeat([]byte("T"), 32)),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := DSN(cfg); got != dsn {
|
||||
t.Fatalf("DSN = %q want %q", got, dsn)
|
||||
}
|
||||
sqlDB, gdb, err := OpenFromApp(t.Context(), backpack.New(cfg))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = sqlDB.Close() })
|
||||
raw, err := gdb.DB()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if raw != sqlDB {
|
||||
t.Fatal("OpenFromApp must share the same *sql.DB with GORM")
|
||||
}
|
||||
}
|
||||
349
modules/lagoon/encrypted.go
Normal file
349
modules/lagoon/encrypted.go
Normal file
@@ -0,0 +1,349 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/hkdf"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql/driver"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
)
|
||||
|
||||
const (
|
||||
encryptedFormatV1 = byte(0x10)
|
||||
encryptedNonceSize = 12
|
||||
encryptedKeySize = 32
|
||||
columnKeyInfo = "summercms.lagoon.encrypted.v1"
|
||||
redactedLiteral = "[redacted]"
|
||||
appKeyErr = "lagoon: app.key is empty or invalid (set SUMMER_APP__KEY to a 32-byte base64 value)"
|
||||
)
|
||||
|
||||
// Encrypted is an AES-256-GCM at-rest cast. Plaintext is unexported; the only
|
||||
// greppable accessor is Reveal. MarshalJSON, String, and GoString always redact.
|
||||
type Encrypted struct {
|
||||
plaintext []byte
|
||||
set bool
|
||||
}
|
||||
|
||||
// NewEncrypted holds plaintext for a subsequent Value() write.
|
||||
func NewEncrypted(plaintext string) Encrypted {
|
||||
return Encrypted{plaintext: []byte(plaintext), set: true}
|
||||
}
|
||||
|
||||
// Reveal returns the plaintext, or "" if the value is unset.
|
||||
func (e Encrypted) Reveal() string {
|
||||
if !e.set {
|
||||
return ""
|
||||
}
|
||||
return string(e.plaintext)
|
||||
}
|
||||
|
||||
// MarshalJSON always emits a redaction, never plaintext.
|
||||
func (e Encrypted) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(redactedLiteral)
|
||||
}
|
||||
|
||||
// String returns a fixed redaction literal.
|
||||
func (e Encrypted) String() string { return redactedLiteral }
|
||||
|
||||
// GoString returns a fixed redaction so %#v cannot leak plaintext.
|
||||
func (e Encrypted) GoString() string { return "lagoon.Encrypted{[redacted]}" }
|
||||
|
||||
// Scan decrypts versioned ciphertext using the published primary key, then
|
||||
// each previous key. src may be string, []byte, or nil (SQL NULL).
|
||||
func (e *Encrypted) Scan(src any) error {
|
||||
if e == nil {
|
||||
return fmt.Errorf("lagoon: encrypted scan on nil receiver")
|
||||
}
|
||||
if src == nil {
|
||||
e.plaintext = nil
|
||||
e.set = false
|
||||
return nil
|
||||
}
|
||||
var raw string
|
||||
switch v := src.(type) {
|
||||
case string:
|
||||
raw = v
|
||||
case []byte:
|
||||
raw = string(v)
|
||||
default:
|
||||
return fmt.Errorf("lagoon: encrypted scan unsupported type %T", src)
|
||||
}
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
e.plaintext = nil
|
||||
e.set = false
|
||||
return nil
|
||||
}
|
||||
plain, err := decryptCiphertext(raw)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
e.plaintext = plain
|
||||
e.set = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Value re-encrypts the current plaintext under the published primary key.
|
||||
// Unset values store SQL NULL.
|
||||
func (e Encrypted) Value() (driver.Value, error) {
|
||||
if !e.set {
|
||||
return nil, nil
|
||||
}
|
||||
return encryptPlaintext(e.plaintext)
|
||||
}
|
||||
|
||||
type encryptionKeys struct {
|
||||
primary []byte
|
||||
previous [][]byte
|
||||
primaryID byte
|
||||
}
|
||||
|
||||
var (
|
||||
keysMu sync.RWMutex
|
||||
currentKeys *encryptionKeys
|
||||
)
|
||||
|
||||
// PublishEncryptionKeys installs column-encryption keys for Encrypted Scan/Value.
|
||||
// OpenFromApp calls this once per boot so row-level encrypt/decrypt does not
|
||||
// re-read config. Pass a nil app from tests that only need the package-level
|
||||
// key set. The backpack publish is best-effort once-per-boot; package-level
|
||||
// keys are the live source of truth and may be rotated in tests.
|
||||
func PublishEncryptionKeys(app *backpack.App, primaryKey []byte, previousKeys [][]byte) error {
|
||||
k, err := newEncryptionKeys(primaryKey, previousKeys)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
setCurrentKeys(k)
|
||||
if app == nil {
|
||||
return nil
|
||||
}
|
||||
if _, ok := app.Lookup[*encryptionKeys](); ok {
|
||||
return nil
|
||||
}
|
||||
return app.Publish(k)
|
||||
}
|
||||
|
||||
func newEncryptionKeys(primaryKey []byte, previousKeys [][]byte) (*encryptionKeys, error) {
|
||||
if len(primaryKey) != encryptedKeySize {
|
||||
return nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
prev := make([][]byte, 0, len(previousKeys))
|
||||
for i, k := range previousKeys {
|
||||
if len(k) != encryptedKeySize {
|
||||
return nil, fmt.Errorf("lagoon: app.previous_keys[%d] is empty or invalid (set SUMMER_APP__KEY to a 32-byte base64 value)", i)
|
||||
}
|
||||
prev = append(prev, cloneBytes(k))
|
||||
}
|
||||
id := byte(len(prev) + 1)
|
||||
if id > 0x0F {
|
||||
id = 0x0F
|
||||
}
|
||||
return &encryptionKeys{
|
||||
primary: cloneBytes(primaryKey),
|
||||
previous: prev,
|
||||
primaryID: id,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func setCurrentKeys(k *encryptionKeys) {
|
||||
keysMu.Lock()
|
||||
currentKeys = k
|
||||
keysMu.Unlock()
|
||||
}
|
||||
|
||||
func clearEncryptionKeys() {
|
||||
setCurrentKeys(nil)
|
||||
}
|
||||
|
||||
func liveKeys() (*encryptionKeys, error) {
|
||||
keysMu.RLock()
|
||||
k := currentKeys
|
||||
keysMu.RUnlock()
|
||||
if k == nil || len(k.primary) != encryptedKeySize {
|
||||
return nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
return k, nil
|
||||
}
|
||||
|
||||
// LoadAppKey reads app.key and app.previous_keys from cfg. Missing, short, or
|
||||
// undecodable values fail loudly with no default (D-11).
|
||||
func LoadAppKey(cfg *compass.Config) ([]byte, [][]byte, error) {
|
||||
if cfg == nil {
|
||||
return nil, nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
primary, err := decodeAppKey(cfg.String("app.key"))
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var previous [][]byte
|
||||
if v, ok := cfg.Lookup("app.previous_keys"); ok && v != nil {
|
||||
previous, err = decodePreviousKeys(v)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
return primary, previous, nil
|
||||
}
|
||||
|
||||
func decodeAppKey(s string) ([]byte, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
raw, err := base64.StdEncoding.DecodeString(s)
|
||||
if err != nil || len(raw) != encryptedKeySize {
|
||||
return nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
return raw, nil
|
||||
}
|
||||
|
||||
func decodePreviousKeys(v any) ([][]byte, error) {
|
||||
items, ok := asStringList(v)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("lagoon: app.previous_keys is empty or invalid (set SUMMER_APP__KEY to a 32-byte base64 value)")
|
||||
}
|
||||
out := make([][]byte, 0, len(items))
|
||||
for _, item := range items {
|
||||
raw, err := decodeAppKey(item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: app.previous_keys is empty or invalid (set SUMMER_APP__KEY to a 32-byte base64 value)")
|
||||
}
|
||||
out = append(out, raw)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func asStringList(v any) ([]string, bool) {
|
||||
switch list := v.(type) {
|
||||
case []string:
|
||||
return list, true
|
||||
case []any:
|
||||
out := make([]string, 0, len(list))
|
||||
for _, item := range list {
|
||||
s, ok := item.(string)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
out = append(out, s)
|
||||
}
|
||||
return out, true
|
||||
default:
|
||||
return nil, false
|
||||
}
|
||||
}
|
||||
|
||||
func deriveColumnKey(appKey []byte) ([]byte, error) {
|
||||
if len(appKey) != encryptedKeySize {
|
||||
return nil, fmt.Errorf(appKeyErr)
|
||||
}
|
||||
return hkdf.Key(sha256.New, appKey, nil, columnKeyInfo, encryptedKeySize)
|
||||
}
|
||||
|
||||
func encryptPlaintext(plain []byte) (string, error) {
|
||||
keys, err := liveKeys()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
colKey, err := deriveColumnKey(keys.primary)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
block, err := aes.NewCipher(colKey)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lagoon: encrypted aes: %w", err)
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lagoon: encrypted gcm: %w", err)
|
||||
}
|
||||
nonce := make([]byte, encryptedNonceSize)
|
||||
if _, err := rand.Read(nonce); err != nil {
|
||||
return "", fmt.Errorf("lagoon: encrypted nonce: %w", err)
|
||||
}
|
||||
sealed := gcm.Seal(nil, nonce, plain, nil)
|
||||
out := make([]byte, 1+len(nonce)+len(sealed))
|
||||
out[0] = encryptedFormatV1 | (keys.primaryID & 0x0F)
|
||||
copy(out[1:], nonce)
|
||||
copy(out[1+len(nonce):], sealed)
|
||||
return base64.StdEncoding.EncodeToString(out), nil
|
||||
}
|
||||
|
||||
func decryptCiphertext(stored string) ([]byte, error) {
|
||||
keys, err := liveKeys()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
raw, err := base64.StdEncoding.DecodeString(stored)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: encrypted ciphertext is not base64")
|
||||
}
|
||||
if len(raw) < 1+encryptedNonceSize+16 {
|
||||
return nil, fmt.Errorf("lagoon: encrypted ciphertext is truncated")
|
||||
}
|
||||
if raw[0]&0xF0 != encryptedFormatV1 {
|
||||
return nil, fmt.Errorf("lagoon: encrypted ciphertext has unknown format")
|
||||
}
|
||||
nonce := raw[1 : 1+encryptedNonceSize]
|
||||
sealed := raw[1+encryptedNonceSize:]
|
||||
|
||||
candidates := make([][]byte, 0, 1+len(keys.previous))
|
||||
candidates = append(candidates, keys.primary)
|
||||
candidates = append(candidates, keys.previous...)
|
||||
var last error
|
||||
for _, appKey := range candidates {
|
||||
plain, err := gcmOpen(appKey, nonce, sealed)
|
||||
if err == nil {
|
||||
return plain, nil
|
||||
}
|
||||
last = err
|
||||
}
|
||||
if last == nil {
|
||||
last = fmt.Errorf("lagoon: encrypted decrypt failed")
|
||||
}
|
||||
return nil, last
|
||||
}
|
||||
|
||||
func gcmOpen(appKey, nonce, sealed []byte) ([]byte, error) {
|
||||
colKey, err := deriveColumnKey(appKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
block, err := aes.NewCipher(colKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return gcm.Open(nil, nonce, sealed, nil)
|
||||
}
|
||||
|
||||
func cloneBytes(b []byte) []byte {
|
||||
if b == nil {
|
||||
return nil
|
||||
}
|
||||
out := make([]byte, len(b))
|
||||
copy(out, b)
|
||||
return out
|
||||
}
|
||||
|
||||
func loadEncryptionKeysFromApp(app *backpack.App) error {
|
||||
if app == nil || app.Config == nil {
|
||||
return fmt.Errorf(appKeyErr)
|
||||
}
|
||||
primary, previous, err := LoadAppKey(app.Config)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return PublishEncryptionKeys(app, primary, previous)
|
||||
}
|
||||
184
modules/lagoon/encrypted_test.go
Normal file
184
modules/lagoon/encrypted_test.go
Normal file
@@ -0,0 +1,184 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
)
|
||||
|
||||
func TestEncryptedRoundTrip(t *testing.T) {
|
||||
key := bytes.Repeat([]byte("A"), 32)
|
||||
if err := PublishEncryptionKeys(nil, key, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(clearEncryptionKeys)
|
||||
|
||||
original := NewEncrypted("hello-secret")
|
||||
stored, err := original.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ciphertext, ok := stored.(string)
|
||||
if !ok {
|
||||
t.Fatalf("Value() type %T, want string", stored)
|
||||
}
|
||||
if ciphertext == "" || ciphertext == "hello-secret" {
|
||||
t.Fatalf("Value() stored plaintext: %q", ciphertext)
|
||||
}
|
||||
|
||||
var got Encrypted
|
||||
if err := got.Scan(ciphertext); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.Reveal() != "hello-secret" {
|
||||
t.Fatalf("Reveal() = %q, want %q", got.Reveal(), "hello-secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptedFreshNonce(t *testing.T) {
|
||||
key := bytes.Repeat([]byte("B"), 32)
|
||||
if err := PublishEncryptionKeys(nil, key, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(clearEncryptionKeys)
|
||||
|
||||
a, err := NewEncrypted("same-plaintext").Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, err := NewEncrypted("same-plaintext").Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("two encrypts of the same plaintext must not produce identical ciphertext")
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptedRedacts(t *testing.T) {
|
||||
const secret = "super-secret-plaintext"
|
||||
e := NewEncrypted(secret)
|
||||
|
||||
raw, err := json.Marshal(e)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bytes.Contains(raw, []byte(secret)) {
|
||||
t.Fatalf("MarshalJSON leaked plaintext: %s", raw)
|
||||
}
|
||||
if strings.Contains(e.String(), secret) {
|
||||
t.Fatalf("String() leaked plaintext: %q", e.String())
|
||||
}
|
||||
goRepr := fmt.Sprintf("%#v", e)
|
||||
if strings.Contains(goRepr, secret) {
|
||||
t.Fatalf("GoString/%%#v leaked plaintext: %q", goRepr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptedPreviousKeysFallback(t *testing.T) {
|
||||
key1 := bytes.Repeat([]byte{1}, 32)
|
||||
key2 := bytes.Repeat([]byte{2}, 32)
|
||||
cfg1 := appKeyConfig(t, base64.StdEncoding.EncodeToString(key1), nil)
|
||||
primary1, previous1, err := LoadAppKey(cfg1)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := PublishEncryptionKeys(nil, primary1, previous1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
stored, err := NewEncrypted("rotate-me").Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ciphertext, ok := stored.(string)
|
||||
if !ok {
|
||||
t.Fatalf("Value() type %T", stored)
|
||||
}
|
||||
|
||||
cfg2 := appKeyConfig(t, base64.StdEncoding.EncodeToString(key2), []string{base64.StdEncoding.EncodeToString(key1)})
|
||||
primary2, previous2, err := LoadAppKey(cfg2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := PublishEncryptionKeys(nil, primary2, previous2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(clearEncryptionKeys)
|
||||
|
||||
var got Encrypted
|
||||
if err := got.Scan(ciphertext); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.Reveal() != "rotate-me" {
|
||||
t.Fatalf("Reveal() after rotation = %q", got.Reveal())
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncryptedLoadAppKeyRejectsInvalid(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
name string
|
||||
key string
|
||||
}{
|
||||
{name: "empty", key: ""},
|
||||
{name: "short", key: base64.StdEncoding.EncodeToString(bytes.Repeat([]byte("x"), 16))},
|
||||
{name: "undecodable", key: "not-valid-base64!!!"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cfg := appKeyConfig(t, tc.key, nil)
|
||||
_, _, err := LoadAppKey(cfg)
|
||||
if err == nil {
|
||||
t.Fatal("want error")
|
||||
}
|
||||
msg := err.Error()
|
||||
if !strings.Contains(msg, "SUMMER_APP__KEY") {
|
||||
t.Fatalf("error %q must mention SUMMER_APP__KEY", msg)
|
||||
}
|
||||
if !strings.Contains(msg, "lagoon:") {
|
||||
t.Fatalf("error %q must be a named lagoon error", msg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func appKeyConfig(t *testing.T, key string, previous []string) *compass.Config {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
var b strings.Builder
|
||||
b.WriteString("key: ")
|
||||
b.WriteString(quoteYAML(key))
|
||||
b.WriteByte('\n')
|
||||
if previous != nil {
|
||||
b.WriteString("previous_keys:\n")
|
||||
for _, p := range previous {
|
||||
b.WriteString(" - ")
|
||||
b.WriteString(quoteYAML(p))
|
||||
b.WriteByte('\n')
|
||||
}
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte(b.String()), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{
|
||||
Dir: dir,
|
||||
Env: "development",
|
||||
Environ: []string{"SUMMER_ENV=development"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
||||
func quoteYAML(s string) string {
|
||||
return `"` + strings.ReplaceAll(s, `"`, `\"`) + `"`
|
||||
}
|
||||
180
modules/lagoon/fill.go
Normal file
180
modules/lagoon/fill.go
Normal file
@@ -0,0 +1,180 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// HasFillable is the Go form of Eloquent $fillable: the model's backstop
|
||||
// allow-list for mass assignment (D-05).
|
||||
type HasFillable interface {
|
||||
Fillable() []string
|
||||
}
|
||||
|
||||
// HasHidden is the Go form of Eloquent $hidden: column names that must not
|
||||
// appear in an accidental JSON marshal (D-08).
|
||||
type HasHidden interface {
|
||||
Hidden() []string
|
||||
}
|
||||
|
||||
var droppedKeys sync.Map
|
||||
|
||||
// Fill copies requested keys onto model only when they are also in allowed.
|
||||
// Unknown and non-fillable keys are dropped with no error (D-06). In
|
||||
// non-production, each type+key pair is logged once.
|
||||
func Fill(model any, allowed []string, requested map[string]any, production bool) error {
|
||||
if model == nil {
|
||||
return fmt.Errorf("lagoon: fill model is nil")
|
||||
}
|
||||
rv := reflect.ValueOf(model)
|
||||
if rv.Kind() != reflect.Ptr || rv.IsNil() {
|
||||
return fmt.Errorf("lagoon: fill model must be a non-nil pointer")
|
||||
}
|
||||
rv = rv.Elem()
|
||||
if rv.Kind() != reflect.Struct {
|
||||
return fmt.Errorf("lagoon: fill model must point to a struct")
|
||||
}
|
||||
rt := rv.Type()
|
||||
typeName := rt.String()
|
||||
for key, val := range requested {
|
||||
if !allowListed(key, allowed) {
|
||||
logDroppedKeyOnce(production, typeName, key)
|
||||
continue
|
||||
}
|
||||
field, ok := fieldByColumn(rv, rt, key)
|
||||
if !ok {
|
||||
logDroppedKeyOnce(production, typeName, key)
|
||||
continue
|
||||
}
|
||||
if err := setField(field, val); err != nil {
|
||||
return fmt.Errorf("lagoon: fill %s: %w", key, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func logDroppedKeyOnce(production bool, typeName, key string) {
|
||||
if production {
|
||||
return
|
||||
}
|
||||
k := typeName + "." + key
|
||||
if _, loaded := droppedKeys.LoadOrStore(k, struct{}{}); loaded {
|
||||
return
|
||||
}
|
||||
slog.Warn("lagoon: dropped non-fillable key", "type", typeName, "key", key)
|
||||
}
|
||||
|
||||
func fieldByColumn(rv reflect.Value, rt reflect.Type, column string) (reflect.Value, bool) {
|
||||
for i := 0; i < rt.NumField(); i++ {
|
||||
f := rt.Field(i)
|
||||
if !f.IsExported() {
|
||||
continue
|
||||
}
|
||||
if gormColumn(f.Tag.Get("gorm")) == column {
|
||||
return rv.Field(i), true
|
||||
}
|
||||
}
|
||||
return reflect.Value{}, false
|
||||
}
|
||||
|
||||
func gormColumn(tag string) string {
|
||||
for _, part := range strings.Split(tag, ";") {
|
||||
part = strings.TrimSpace(part)
|
||||
if after, ok := strings.CutPrefix(part, "column:"); ok {
|
||||
return after
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func setField(field reflect.Value, val any) error {
|
||||
if !field.CanSet() {
|
||||
return fmt.Errorf("field cannot be set")
|
||||
}
|
||||
if val == nil {
|
||||
field.Set(reflect.Zero(field.Type()))
|
||||
if field.CanAddr() {
|
||||
if scanner, ok := field.Addr().Interface().(sql.Scanner); ok {
|
||||
return scanner.Scan(nil)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
src := reflect.ValueOf(val)
|
||||
if field.Kind() == reflect.Ptr {
|
||||
elemType := field.Type().Elem()
|
||||
converted, err := convertValue(src, elemType)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ptr := reflect.New(elemType)
|
||||
ptr.Elem().Set(converted)
|
||||
field.Set(ptr)
|
||||
return nil
|
||||
}
|
||||
converted, err := convertValue(src, field.Type())
|
||||
if err == nil {
|
||||
field.Set(converted)
|
||||
return nil
|
||||
}
|
||||
if field.Type() == encryptedType {
|
||||
return err
|
||||
}
|
||||
if field.CanAddr() {
|
||||
if scanner, ok := field.Addr().Interface().(sql.Scanner); ok {
|
||||
scanSrc, scanErr := fillScanSource(val)
|
||||
if scanErr != nil {
|
||||
return scanErr
|
||||
}
|
||||
return scanner.Scan(scanSrc)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func fillScanSource(val any) (any, error) {
|
||||
switch v := val.(type) {
|
||||
case string, []byte:
|
||||
return v, nil
|
||||
default:
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return b, nil
|
||||
}
|
||||
}
|
||||
|
||||
var encryptedType = reflect.TypeOf(Encrypted{})
|
||||
|
||||
func convertValue(src reflect.Value, destType reflect.Type) (reflect.Value, error) {
|
||||
if destType == encryptedType {
|
||||
return encryptedFromRequest(src)
|
||||
}
|
||||
if src.Type().AssignableTo(destType) {
|
||||
return src, nil
|
||||
}
|
||||
if src.Type().ConvertibleTo(destType) {
|
||||
return src.Convert(destType), nil
|
||||
}
|
||||
return reflect.Value{}, fmt.Errorf("cannot assign %s to %s", src.Type(), destType)
|
||||
}
|
||||
|
||||
// encryptedFromRequest treats request input for an Encrypted column as
|
||||
// plaintext. It never falls through to Encrypted.Scan: Scan decrypts, so a
|
||||
// write path that scanned request input would reject real secrets and accept
|
||||
// another row's ciphertext, copying that row's secret.
|
||||
func encryptedFromRequest(src reflect.Value) (reflect.Value, error) {
|
||||
if src.Type() == encryptedType {
|
||||
return src, nil
|
||||
}
|
||||
if src.Kind() != reflect.String {
|
||||
return reflect.Value{}, fmt.Errorf("encrypted value must be a string")
|
||||
}
|
||||
return reflect.ValueOf(NewEncrypted(src.String())), nil
|
||||
}
|
||||
68
modules/lagoon/fill_fuzz_test.go
Normal file
68
modules/lagoon/fill_fuzz_test.go
Normal file
@@ -0,0 +1,68 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fuzzFillFixture is the D-07 framework-level target: only name/notes are
|
||||
// allow-listed. collection_id, secret, and id must stay at their pre-Fill values.
|
||||
type fuzzFillFixture struct {
|
||||
ID uint `gorm:"column:id"`
|
||||
Name string `gorm:"column:name"`
|
||||
Notes *string `gorm:"column:notes"`
|
||||
CollectionID uint `gorm:"column:collection_id"`
|
||||
Secret string `gorm:"column:secret"`
|
||||
}
|
||||
|
||||
func FuzzFill(f *testing.F) {
|
||||
f.Add(`{"name":"ok","notes":"n"}`)
|
||||
f.Add(`{"name":"ok","collection_id":9,"secret":"leak","id":99}`)
|
||||
f.Add(`{"name":null,"notes":null}`)
|
||||
f.Add(`{"collection_id":1,"unknown":true}`)
|
||||
f.Add(`{"name":{"nested":true},"notes":[1,2]}`)
|
||||
f.Add(`{"name":1,"notes":false}`)
|
||||
f.Add(`[]`)
|
||||
f.Add(``)
|
||||
f.Add(`null`)
|
||||
f.Add(`{"name":"x","extra":{"deep":{"x":1}}}`)
|
||||
|
||||
f.Fuzz(func(t *testing.T, raw string) {
|
||||
var requested map[string]any
|
||||
_ = json.Unmarshal([]byte(raw), &requested)
|
||||
if requested == nil {
|
||||
requested = map[string]any{}
|
||||
}
|
||||
requested["collection_id"] = float64(99)
|
||||
requested["secret"] = "pwned"
|
||||
requested["id"] = float64(7)
|
||||
|
||||
row := fuzzFillFixture{ID: 1, CollectionID: 3, Secret: "keep"}
|
||||
if err := Fill(&row, []string{"name", "notes"}, requested, true); err != nil {
|
||||
// Type mismatches on allow-listed keys are errors, not panics.
|
||||
// Non-allow-listed fields must still be untouched.
|
||||
}
|
||||
if row.CollectionID != 3 {
|
||||
t.Fatalf("collection_id mutated to %d", row.CollectionID)
|
||||
}
|
||||
if row.Secret != "keep" {
|
||||
t.Fatalf("secret mutated to %q", row.Secret)
|
||||
}
|
||||
if row.ID != 1 {
|
||||
t.Fatalf("id mutated to %d", row.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFuzzFillNilAndEmptyNeverPanic(t *testing.T) {
|
||||
row := fuzzFillFixture{CollectionID: 3, Secret: "keep"}
|
||||
if err := Fill(&row, []string{"name"}, nil, true); err != nil {
|
||||
t.Fatalf("nil map: %v", err)
|
||||
}
|
||||
if err := Fill(&row, []string{"name"}, map[string]any{}, true); err != nil {
|
||||
t.Fatalf("empty map: %v", err)
|
||||
}
|
||||
if row.CollectionID != 3 || row.Secret != "keep" {
|
||||
t.Fatalf("zero-input fill mutated %+v", row)
|
||||
}
|
||||
}
|
||||
162
modules/lagoon/fill_test.go
Normal file
162
modules/lagoon/fill_test.go
Normal file
@@ -0,0 +1,162 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fillFixture struct {
|
||||
Name string `gorm:"column:name"`
|
||||
CollectionID uint `gorm:"column:collection_id"`
|
||||
Notes *string `gorm:"column:notes"`
|
||||
}
|
||||
|
||||
func (fillFixture) Fillable() []string { return []string{"name", "notes"} }
|
||||
func (fillFixture) Hidden() []string { return []string{"collection_id"} }
|
||||
|
||||
var (
|
||||
_ HasFillable = fillFixture{}
|
||||
_ HasHidden = fillFixture{}
|
||||
)
|
||||
|
||||
func TestFillAllowList(t *testing.T) {
|
||||
var row fillFixture
|
||||
row.CollectionID = 3
|
||||
err := Fill(&row, []string{"name"}, map[string]any{
|
||||
"name": "x",
|
||||
"collection_id": uint(9),
|
||||
}, true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if row.Name != "x" {
|
||||
t.Fatalf("name = %q", row.Name)
|
||||
}
|
||||
if row.CollectionID != 3 {
|
||||
t.Fatalf("collection_id mutated to %d", row.CollectionID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFillDroppedKeyLogsOnce(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
prev := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelWarn})))
|
||||
defer slog.SetDefault(prev)
|
||||
|
||||
var row fillFixture
|
||||
requested := map[string]any{"name": "once", "collection_id": uint(1)}
|
||||
if err := Fill(&row, []string{"name"}, requested, false); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Fill(&row, []string{"name"}, requested, false); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
log := buf.String()
|
||||
if strings.Count(log, "collection_id") != 1 {
|
||||
t.Fatalf("dropped key should log once, got %q", log)
|
||||
}
|
||||
if !strings.Contains(log, "lagoon: dropped non-fillable key") {
|
||||
t.Fatalf("missing warn message: %q", log)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFillProductionSilent(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
prev := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelWarn})))
|
||||
defer slog.SetDefault(prev)
|
||||
|
||||
var row fillFixture
|
||||
if err := Fill(&row, []string{"name"}, map[string]any{
|
||||
"name": "prod",
|
||||
"unknown_field": true,
|
||||
}, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if buf.Len() != 0 {
|
||||
t.Fatalf("production fill must be silent, got %q", buf.String())
|
||||
}
|
||||
if row.Name != "prod" {
|
||||
t.Fatalf("name = %q", row.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFillNilPointerClearsNullable(t *testing.T) {
|
||||
existing := "keep"
|
||||
row := fillFixture{Notes: &existing}
|
||||
if err := Fill(&row, []string{"notes"}, map[string]any{"notes": nil}, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if row.Notes != nil {
|
||||
t.Fatalf("notes = %v, want nil", row.Notes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFillDroppedKeyNeverErrors(t *testing.T) {
|
||||
var row fillFixture
|
||||
if err := Fill(&row, []string{"name"}, map[string]any{"collection_id": uint(9), "nope": 1}, true); err != nil {
|
||||
t.Fatalf("dropped keys must not error: %v", err)
|
||||
}
|
||||
if row.CollectionID != 0 {
|
||||
t.Fatalf("collection_id = %d", row.CollectionID)
|
||||
}
|
||||
}
|
||||
|
||||
type fillSecretFixture struct {
|
||||
Name string `gorm:"column:name"`
|
||||
APIKey Encrypted `gorm:"column:api_key" json:"-"`
|
||||
Token *Encrypted `gorm:"column:token" json:"-"`
|
||||
}
|
||||
|
||||
func TestFillEncryptedTakesPlaintext(t *testing.T) {
|
||||
if err := PublishEncryptionKeys(nil, bytes.Repeat([]byte("F"), 32), nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
allowed := []string{"name", "api_key", "token"}
|
||||
var row fillSecretFixture
|
||||
if err := Fill(&row, allowed, map[string]any{"api_key": "sk-plain", "token": "tok-plain"}, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if row.APIKey.Reveal() != "sk-plain" {
|
||||
t.Fatalf("api_key Reveal = %q", row.APIKey.Reveal())
|
||||
}
|
||||
if row.Token == nil || row.Token.Reveal() != "tok-plain" {
|
||||
t.Fatalf("token = %v", row.Token)
|
||||
}
|
||||
|
||||
// Another row's ciphertext is stored as literal text, never decrypted.
|
||||
stolen, err := NewEncrypted("victim-secret").Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ciphertext, ok := stolen.(string)
|
||||
if !ok {
|
||||
t.Fatalf("ciphertext driver value is %T", stolen)
|
||||
}
|
||||
var thief fillSecretFixture
|
||||
if err := Fill(&thief, allowed, map[string]any{"api_key": ciphertext}, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if thief.APIKey.Reveal() != ciphertext {
|
||||
t.Fatal("ciphertext in a request must not be decrypted into the victim's plaintext")
|
||||
}
|
||||
|
||||
for _, bad := range []any{42, true, map[string]any{"x": 1}, []byte("raw")} {
|
||||
var r fillSecretFixture
|
||||
if err := Fill(&r, allowed, map[string]any{"api_key": bad}, true); err == nil {
|
||||
t.Fatalf("api_key=%T must be rejected", bad)
|
||||
}
|
||||
if r.APIKey.Reveal() != "" {
|
||||
t.Fatalf("api_key=%T left a value behind", bad)
|
||||
}
|
||||
}
|
||||
|
||||
if err := Fill(&row, allowed, map[string]any{"api_key": nil, "token": nil}, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if row.APIKey.Reveal() != "" || row.Token != nil {
|
||||
t.Fatalf("nil must clear: api_key=%q token=%v", row.APIKey.Reveal(), row.Token)
|
||||
}
|
||||
}
|
||||
42
modules/lagoon/hidden_marshal_test.go
Normal file
42
modules/lagoon/hidden_marshal_test.go
Normal file
@@ -0,0 +1,42 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// hiddenMarshalFixture covers the HasHidden contract inside the framework
|
||||
// module. Enumerating fonoteka.go plugin models.All() from this package
|
||||
// would import the application into summercms.go, which CLAUDE.md forbids;
|
||||
// that registry walk lives in fonoteka.go (classes.TestHiddenNeverMarshals).
|
||||
type hiddenMarshalFixture struct {
|
||||
Name string `gorm:"column:name" json:"name"`
|
||||
CollectionID uint `gorm:"column:collection_id" json:"collection_id"`
|
||||
Secret string `gorm:"column:secret" json:"-"`
|
||||
}
|
||||
|
||||
func (hiddenMarshalFixture) Hidden() []string { return []string{"secret"} }
|
||||
|
||||
func TestHiddenNeverMarshals(t *testing.T) {
|
||||
row := hiddenMarshalFixture{Name: "ok", CollectionID: 3, Secret: hiddenSentinel}
|
||||
raw, err := json.Marshal(row)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(string(raw), hiddenSentinel) {
|
||||
t.Fatalf("hidden sentinel leaked: %s", raw)
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, ok := payload["secret"]; ok {
|
||||
t.Fatalf("json:\"-\" field present: %s", raw)
|
||||
}
|
||||
if payload["name"] != "ok" {
|
||||
t.Fatalf("visible field missing: %s", raw)
|
||||
}
|
||||
}
|
||||
|
||||
const hiddenSentinel = "HIDDEN-SENTINEL-05-06"
|
||||
97
modules/lagoon/jsonable.go
Normal file
97
modules/lagoon/jsonable.go
Normal file
@@ -0,0 +1,97 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// Jsonable is a JSON-text column (Winter $jsonable) stored as TEXT, not jsonb.
|
||||
// A struct (not `type Jsonable[T] T`, which Go rejects) is required so Scan
|
||||
// can set Valid independently of T's zero value: a nil []string (SQL NULL)
|
||||
// is distinct from a non-nil empty slice (SQL '[]'). Inner fields are
|
||||
// gorm:"-" so GORM uses Scanner/Valuer instead of walking them.
|
||||
//
|
||||
// Data is the payload (cannot be named Value: that name is driver.Valuer).
|
||||
// Valid=false ↔ SQL NULL. Valid=true encodes Data with encoding/json.
|
||||
// NullOnEmpty, when set on a slice/map value, stores SQL NULL instead of
|
||||
// '[]'/'{}' so callers pick empty-vs-null per column (Pitfall 4).
|
||||
type Jsonable[T any] struct {
|
||||
Data T `gorm:"-"`
|
||||
Valid bool `gorm:"-"`
|
||||
NullOnEmpty bool `gorm:"-"`
|
||||
}
|
||||
|
||||
// Get returns the underlying payload.
|
||||
func (j Jsonable[T]) Get() T { return j.Data }
|
||||
|
||||
// Scan decodes src as JSON text. SQL NULL (and JSON null) leave Valid=false.
|
||||
func (j *Jsonable[T]) Scan(src any) error {
|
||||
if j == nil {
|
||||
return fmt.Errorf("lagoon: jsonable scan on nil receiver")
|
||||
}
|
||||
if src == nil {
|
||||
var zero T
|
||||
j.Data = zero
|
||||
j.Valid = false
|
||||
return nil
|
||||
}
|
||||
var raw []byte
|
||||
switch v := src.(type) {
|
||||
case []byte:
|
||||
raw = v
|
||||
case string:
|
||||
raw = []byte(v)
|
||||
default:
|
||||
return fmt.Errorf("lagoon: jsonable scan unsupported type %T", src)
|
||||
}
|
||||
raw = bytes.TrimSpace(raw)
|
||||
if len(raw) == 0 || bytes.Equal(raw, []byte("null")) {
|
||||
var zero T
|
||||
j.Data = zero
|
||||
j.Valid = false
|
||||
return nil
|
||||
}
|
||||
var val T
|
||||
if err := json.Unmarshal(raw, &val); err != nil {
|
||||
return fmt.Errorf("lagoon: jsonable scan: %w", err)
|
||||
}
|
||||
j.Data = val
|
||||
j.Valid = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Value returns JSON text or nil. It never returns a numeric driver.Value.
|
||||
func (j Jsonable[T]) Value() (driver.Value, error) {
|
||||
if !j.Valid {
|
||||
return nil, nil
|
||||
}
|
||||
if j.NullOnEmpty && jsonableEmpty(j.Data) {
|
||||
return nil, nil
|
||||
}
|
||||
b, err := json.Marshal(j.Data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if bytes.Equal(b, []byte("null")) {
|
||||
return nil, nil
|
||||
}
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// GormDataType stores jsonable columns as TEXT.
|
||||
func (Jsonable[T]) GormDataType() string { return "text" }
|
||||
|
||||
func jsonableEmpty(v any) bool {
|
||||
rv := reflect.ValueOf(v)
|
||||
switch rv.Kind() {
|
||||
case reflect.Slice, reflect.Map, reflect.String:
|
||||
return rv.Len() == 0
|
||||
case reflect.Ptr, reflect.Interface:
|
||||
return rv.IsNil()
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
106
modules/lagoon/jsonable_test.go
Normal file
106
modules/lagoon/jsonable_test.go
Normal file
@@ -0,0 +1,106 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestJsonableStringSliceNullAndEmpty(t *testing.T) {
|
||||
var nullCol Jsonable[[]string]
|
||||
if err := nullCol.Scan(nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if nullCol.Valid || nullCol.Data != nil {
|
||||
t.Fatalf("NULL scan: valid=%v data=%v", nullCol.Valid, nullCol.Data)
|
||||
}
|
||||
got, err := nullCol.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != nil {
|
||||
t.Fatalf("NULL value = %#v, want nil", got)
|
||||
}
|
||||
|
||||
var emptyCol Jsonable[[]string]
|
||||
if err := emptyCol.Scan("[]"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !emptyCol.Valid {
|
||||
t.Fatal("[] must be valid")
|
||||
}
|
||||
if emptyCol.Data == nil || len(emptyCol.Data) != 0 {
|
||||
t.Fatalf("[] scan = %#v, want empty non-nil slice", emptyCol.Data)
|
||||
}
|
||||
got, err = emptyCol.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "[]" {
|
||||
t.Fatalf("[] value = %#v, want %q", got, "[]")
|
||||
}
|
||||
|
||||
var populated Jsonable[[]string]
|
||||
if err := populated.Scan(`["https://img.discogs.com/a.jpg"]`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !populated.Valid || len(populated.Data) != 1 || populated.Data[0] != "https://img.discogs.com/a.jpg" {
|
||||
t.Fatalf("populated scan = %#v", populated.Data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJsonableTrackEntryRoundTrip(t *testing.T) {
|
||||
type TrackEntry = map[string]any
|
||||
src := `[{"title":"Smells Like Teen Spirit","position":"A1"}]`
|
||||
var col Jsonable[[]TrackEntry]
|
||||
if err := col.Scan(src); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !col.Valid || len(col.Data) != 1 {
|
||||
t.Fatalf("scan = %#v", col.Data)
|
||||
}
|
||||
title, _ := col.Data[0]["title"].(string)
|
||||
if title != "Smells Like Teen Spirit" {
|
||||
t.Fatalf("title = %#v", col.Data[0]["title"])
|
||||
}
|
||||
if _, isStringSlice := any(col.Data).([]string); isStringSlice {
|
||||
t.Fatal("tracklist must not collapse to []string")
|
||||
}
|
||||
got, err := col.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s, ok := got.(string)
|
||||
if !ok {
|
||||
t.Fatalf("Value() type %T, want string", got)
|
||||
}
|
||||
var round Jsonable[[]TrackEntry]
|
||||
if err := round.Scan(s); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
title, _ = round.Data[0]["title"].(string)
|
||||
if title != "Smells Like Teen Spirit" {
|
||||
t.Fatalf("round-trip title = %#v", round.Data[0]["title"])
|
||||
}
|
||||
|
||||
var nilCol Jsonable[[]TrackEntry]
|
||||
if err := nilCol.Scan(nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err = nilCol.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != nil {
|
||||
t.Fatalf("NULL tracklist value = %#v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestJsonableValueNeverNumeric(t *testing.T) {
|
||||
col := Jsonable[[]string]{Data: []string{"a"}, Valid: true}
|
||||
got, err := col.Value()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, ok := got.(string); !ok && got != nil {
|
||||
t.Fatalf("Value() type %T, want string or nil", got)
|
||||
}
|
||||
}
|
||||
26
modules/lagoon/keygen.go
Normal file
26
modules/lagoon/keygen.go
Normal file
@@ -0,0 +1,26 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/bonfire"
|
||||
)
|
||||
|
||||
// KeyGenerateCommand prints a fresh 32-byte base64 app.key and performs no
|
||||
// other side effect (D-11).
|
||||
func KeyGenerateCommand() bonfire.Command {
|
||||
return bonfire.Command{
|
||||
Name: "key:generate",
|
||||
Description: "Print a fresh 32-byte base64 app.key",
|
||||
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
|
||||
key := make([]byte, encryptedKeySize)
|
||||
if _, err := rand.Read(key); err != nil {
|
||||
return err
|
||||
}
|
||||
out.Success(base64.StdEncoding.EncodeToString(key))
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
40
modules/lagoon/keygen_test.go
Normal file
40
modules/lagoon/keygen_test.go
Normal file
@@ -0,0 +1,40 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/bonfire"
|
||||
)
|
||||
|
||||
func TestEncryptedKeyGeneratePrints32ByteBase64(t *testing.T) {
|
||||
cmd := KeyGenerateCommand()
|
||||
if cmd.Name != "key:generate" {
|
||||
t.Fatalf("Name = %q", cmd.Name)
|
||||
}
|
||||
if cmd.Run == nil {
|
||||
t.Fatal("Run is nil")
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
out := bonfire.NewOutput(strings.NewReader(""), &buf, &buf)
|
||||
if err := cmd.Run(context.Background(), nil, out); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
line := strings.TrimSpace(buf.String())
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 0 {
|
||||
t.Fatal("key:generate produced no output")
|
||||
}
|
||||
encoded := fields[len(fields)-1]
|
||||
raw, err := base64.StdEncoding.DecodeString(encoded)
|
||||
if err != nil {
|
||||
t.Fatalf("output %q is not base64: %v", encoded, err)
|
||||
}
|
||||
if len(raw) != 32 {
|
||||
t.Fatalf("decoded key length = %d, want 32", len(raw))
|
||||
}
|
||||
}
|
||||
105
modules/lagoon/laravel_decrypt.go
Normal file
105
modules/lagoon/laravel_decrypt.go
Normal file
@@ -0,0 +1,105 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// laravelPayload is Laravel's Encrypter JSON body (iv/value/mac).
|
||||
type laravelPayload struct {
|
||||
IV string `json:"iv"`
|
||||
Value string `json:"value"`
|
||||
MAC string `json:"mac"`
|
||||
}
|
||||
|
||||
// DecryptLaravelPayload decrypts a Laravel AES-256-CBC+HMAC "encrypted" payload
|
||||
// (base64 JSON {iv,value,mac}) using the raw APP_KEY bytes.
|
||||
//
|
||||
// Cutover-import-only; never called from Encrypted's live Scan/Value path (D-10).
|
||||
func DecryptLaravelPayload(payloadJSON string, appKey []byte) ([]byte, error) {
|
||||
if len(appKey) != encryptedKeySize {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload key must be 32 bytes")
|
||||
}
|
||||
body, err := decodeLaravelJSON(payloadJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var payload laravelPayload
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload json: %w", err)
|
||||
}
|
||||
if payload.IV == "" || payload.Value == "" || payload.MAC == "" {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload missing iv, value, or mac")
|
||||
}
|
||||
wantMAC, err := hex.DecodeString(payload.MAC)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload mac: %w", err)
|
||||
}
|
||||
mac := hmac.New(sha256.New, appKey)
|
||||
_, _ = mac.Write([]byte(payload.IV + payload.Value))
|
||||
gotMAC := mac.Sum(nil)
|
||||
if !hmac.Equal(wantMAC, gotMAC) {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload mac mismatch")
|
||||
}
|
||||
iv, err := base64.StdEncoding.DecodeString(payload.IV)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload iv: %w", err)
|
||||
}
|
||||
ct, err := base64.StdEncoding.DecodeString(payload.Value)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload value: %w", err)
|
||||
}
|
||||
if len(iv) != aes.BlockSize {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload iv length %d", len(iv))
|
||||
}
|
||||
if len(ct) == 0 || len(ct)%aes.BlockSize != 0 {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload ciphertext length %d", len(ct))
|
||||
}
|
||||
block, err := aes.NewCipher(appKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
plain := make([]byte, len(ct))
|
||||
cipher.NewCBCDecrypter(block, iv).CryptBlocks(plain, ct)
|
||||
return pkcs7Unpad(plain, aes.BlockSize)
|
||||
}
|
||||
|
||||
func decodeLaravelJSON(payloadJSON string) ([]byte, error) {
|
||||
s := strings.TrimSpace(payloadJSON)
|
||||
if s == "" {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload is empty")
|
||||
}
|
||||
if decoded, err := base64.StdEncoding.DecodeString(s); err == nil {
|
||||
trimmed := strings.TrimSpace(string(decoded))
|
||||
if strings.HasPrefix(trimmed, "{") {
|
||||
return []byte(trimmed), nil
|
||||
}
|
||||
}
|
||||
if strings.HasPrefix(s, "{") {
|
||||
return []byte(s), nil
|
||||
}
|
||||
return nil, fmt.Errorf("lagoon: laravel payload is not base64 JSON")
|
||||
}
|
||||
|
||||
func pkcs7Unpad(b []byte, blockSize int) ([]byte, error) {
|
||||
if blockSize <= 0 || len(b) == 0 || len(b)%blockSize != 0 {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload padding")
|
||||
}
|
||||
pad := int(b[len(b)-1])
|
||||
if pad == 0 || pad > blockSize || pad > len(b) {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload padding")
|
||||
}
|
||||
for i := len(b) - pad; i < len(b); i++ {
|
||||
if int(b[i]) != pad {
|
||||
return nil, fmt.Errorf("lagoon: laravel payload padding")
|
||||
}
|
||||
}
|
||||
return b[:len(b)-pad], nil
|
||||
}
|
||||
21
modules/lagoon/laravel_decrypt_test.go
Normal file
21
modules/lagoon/laravel_decrypt_test.go
Normal file
@@ -0,0 +1,21 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Fixture produced by PHP openssl_encrypt AES-256-CBC + HMAC-SHA256 over iv+value,
|
||||
// matching Laravel's Encrypter payload (base64 JSON {iv,value,mac}) with serialize=false.
|
||||
const laravelTestPayload = "eyJpdiI6IlNVbEpTVWxKU1VsSlNVbEpTVWxKU1E9PSIsInZhbHVlIjoiYVJ6ODYrRitUM25oVFNoR0o5WGMxbUJtSW5DeWdaVWN6WktYR3RFV3liZz0iLCJtYWMiOiI4YmI4NTRmZDk0ZmQ0NDYzNjk3ZDhmYmQ3M2QyNzFhYzI1MGI1YjFhNzk2ZDdiOWVjN2YyYmQ3YzJjOTk4MzBiIn0="
|
||||
|
||||
func TestEncryptedDecryptLaravelPayload(t *testing.T) {
|
||||
key := bytes.Repeat([]byte("K"), 32)
|
||||
got, err := DecryptLaravelPayload(laravelTestPayload, key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(got) != "sk-test-super-secret-api-key" {
|
||||
t.Fatalf("decrypted = %q", got)
|
||||
}
|
||||
}
|
||||
48
modules/lagoon/lifecycle.go
Normal file
48
modules/lagoon/lifecycle.go
Normal file
@@ -0,0 +1,48 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// HasBeforeValidate is the Winter beforeValidate hook. GORM does not
|
||||
// dispatch this name; later validation wiring type-asserts it.
|
||||
type HasBeforeValidate interface {
|
||||
BeforeValidate(tx *gorm.DB) error
|
||||
}
|
||||
|
||||
// HasBeforeCreate matches GORM's native BeforeCreate hook so a model
|
||||
// needs no adapter. GORM dispatches it by method name.
|
||||
type HasBeforeCreate interface {
|
||||
BeforeCreate(tx *gorm.DB) error
|
||||
}
|
||||
|
||||
// HasBeforeSave matches GORM's native BeforeSave hook so a model
|
||||
// needs no adapter. GORM dispatches it by method name.
|
||||
type HasBeforeSave interface {
|
||||
BeforeSave(tx *gorm.DB) error
|
||||
}
|
||||
|
||||
// HasBeforeDelete matches GORM's native BeforeDelete hook so a model
|
||||
// needs no adapter. GORM dispatches it by method name.
|
||||
type HasBeforeDelete interface {
|
||||
BeforeDelete(tx *gorm.DB) error
|
||||
}
|
||||
|
||||
// HasAfterDelete matches GORM's native AfterDelete hook so a model
|
||||
// needs no adapter. GORM dispatches it by method name.
|
||||
type HasAfterDelete interface {
|
||||
AfterDelete(tx *gorm.DB) error
|
||||
}
|
||||
|
||||
// WithSoftDeleteCascade names the "Collection cascades to Album" pattern
|
||||
// (DATA-03). Callers invoke it from their own BeforeDelete method: GORM
|
||||
// already runs BeforeDelete inside the parent Delete transaction, so this
|
||||
// does not open a new transaction. A cascade error aborts the parent delete.
|
||||
func WithSoftDeleteCascade(tx *gorm.DB, cascade func(tx *gorm.DB) error) error {
|
||||
if tx == nil {
|
||||
return fmt.Errorf("lagoon: soft-delete cascade tx is nil")
|
||||
}
|
||||
return cascade(tx)
|
||||
}
|
||||
183
modules/lagoon/lifecycle_test.go
Normal file
183
modules/lagoon/lifecycle_test.go
Normal file
@@ -0,0 +1,183 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type hookRow struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
Flag bool `gorm:"column:flag"`
|
||||
}
|
||||
|
||||
func (hookRow) TableName() string { return "lagoon_hook_rows" }
|
||||
|
||||
func (h *hookRow) BeforeCreate(tx *gorm.DB) error {
|
||||
h.Flag = true
|
||||
return nil
|
||||
}
|
||||
|
||||
type cascadeParent struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"`
|
||||
}
|
||||
|
||||
func (cascadeParent) TableName() string { return "lagoon_cascade_parents" }
|
||||
|
||||
func (p *cascadeParent) BeforeDelete(tx *gorm.DB) error {
|
||||
return WithSoftDeleteCascade(tx, func(tx *gorm.DB) error {
|
||||
return tx.Where("parent_id = ?", p.ID).Delete(&cascadeChild{}).Error
|
||||
})
|
||||
}
|
||||
|
||||
type cascadeChild struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
ParentID uint `gorm:"column:parent_id"`
|
||||
Name string `gorm:"column:name"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"`
|
||||
}
|
||||
|
||||
func (cascadeChild) TableName() string { return "lagoon_cascade_children" }
|
||||
|
||||
type failingParent struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
DeletedAt gorm.DeletedAt `gorm:"column:deleted_at"`
|
||||
}
|
||||
|
||||
func (failingParent) TableName() string { return "lagoon_cascade_parents" }
|
||||
|
||||
func (p *failingParent) BeforeDelete(tx *gorm.DB) error {
|
||||
return WithSoftDeleteCascade(tx, func(tx *gorm.DB) error {
|
||||
if err := tx.Where("parent_id = ?", p.ID).Delete(&cascadeChild{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return gorm.ErrInvalidData
|
||||
})
|
||||
}
|
||||
|
||||
func TestLifecycleBeforeCreate(t *testing.T) {
|
||||
sqlDB, _ := dedicatedDB(t, "lagoon_lifecycle_create")
|
||||
gdb, err := Use(t.Context(), sqlDB)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`
|
||||
CREATE TABLE lagoon_hook_rows (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
flag BOOLEAN NOT NULL DEFAULT FALSE
|
||||
)`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
row := hookRow{Name: "created"}
|
||||
if err := gdb.Create(&row).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !row.Flag {
|
||||
t.Fatal("BeforeCreate did not run")
|
||||
}
|
||||
var stored hookRow
|
||||
if err := gdb.Take(&stored, row.ID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !stored.Flag {
|
||||
t.Fatal("BeforeCreate flag was not persisted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithSoftDeleteCascadeSameTransaction(t *testing.T) {
|
||||
gdb := cascadeSchema(t, "lagoon_lifecycle_cascade")
|
||||
parent := cascadeParent{Name: "p"}
|
||||
if err := gdb.Create(&parent).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
child := cascadeChild{ParentID: parent.ID, Name: "c"}
|
||||
if err := gdb.Create(&child).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Delete(&parent).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var parents, children int64
|
||||
if err := gdb.Unscoped().Model(&cascadeParent{}).Where("id = ? AND deleted_at IS NOT NULL", parent.ID).Count(&parents).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Unscoped().Model(&cascadeChild{}).Where("id = ? AND deleted_at IS NOT NULL", child.ID).Count(&children).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parents != 1 || children != 1 {
|
||||
t.Fatalf("want both soft-deleted, parents=%d children=%d", parents, children)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithSoftDeleteCascadeRollback(t *testing.T) {
|
||||
gdb := cascadeSchema(t, "lagoon_lifecycle_rollback")
|
||||
parent := failingParent{Name: "p"}
|
||||
if err := gdb.Create(&parent).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
child := cascadeChild{ParentID: parent.ID, Name: "c"}
|
||||
if err := gdb.Create(&child).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err := gdb.Delete(&parent).Error
|
||||
if err == nil {
|
||||
t.Fatal("want cascade error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), gorm.ErrInvalidData.Error()) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
var parents, children int64
|
||||
if err := gdb.Model(&cascadeParent{}).Where("id = ?", parent.ID).Count(&parents).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Model(&cascadeChild{}).Where("id = ?", child.ID).Count(&children).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parents != 1 || children != 1 {
|
||||
t.Fatalf("cascade error must roll back both, parents=%d children=%d", parents, children)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithSoftDeleteCascadeNilTx(t *testing.T) {
|
||||
err := WithSoftDeleteCascade(nil, func(tx *gorm.DB) error { return nil })
|
||||
if err == nil {
|
||||
t.Fatal("want nil tx error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "nil") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func cascadeSchema(t *testing.T, dbName string) *gorm.DB {
|
||||
t.Helper()
|
||||
sqlDB, _ := dedicatedDB(t, dbName)
|
||||
gdb, err := Use(t.Context(), sqlDB)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stmts := []string{
|
||||
`CREATE TABLE lagoon_cascade_parents (
|
||||
id SERIAL PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
deleted_at TIMESTAMPTZ
|
||||
)`,
|
||||
`CREATE TABLE lagoon_cascade_children (
|
||||
id SERIAL PRIMARY KEY,
|
||||
parent_id INTEGER NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
deleted_at TIMESTAMPTZ
|
||||
)`,
|
||||
}
|
||||
for _, stmt := range stmts {
|
||||
if err := gdb.Exec(stmt).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return gdb
|
||||
}
|
||||
183
modules/lagoon/migrations.go
Normal file
183
modules/lagoon/migrations.go
Normal file
@@ -0,0 +1,183 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/lagoon/attach"
|
||||
"git.golem15.com/golem15/summercms/modules/pact"
|
||||
"git.golem15.com/golem15/summercms/modules/party"
|
||||
"github.com/go-gormigrate/gormigrate/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const historyTablePrefix = "summer_migrations_"
|
||||
|
||||
var (
|
||||
// ErrUnknownPlugin is returned when migrate:rollback names a plugin that is not activated.
|
||||
ErrUnknownPlugin = errors.New("lagoon: plugin is not activated")
|
||||
// ErrNoMigrations is returned when there is no plugin migration set to roll back.
|
||||
ErrNoMigrations = errors.New("lagoon: no plugin migrations to roll back")
|
||||
)
|
||||
|
||||
// HistoryTableName returns the isolated gormigrate table for pluginID.
|
||||
func HistoryTableName(pluginID string) (string, error) {
|
||||
if err := validatePluginID(pluginID); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return historyTablePrefix + strings.ReplaceAll(pluginID, ".", "_"), nil
|
||||
}
|
||||
|
||||
func validatePluginID(id string) error {
|
||||
if id == "" {
|
||||
return fmt.Errorf("lagoon: plugin id is empty")
|
||||
}
|
||||
for _, r := range id {
|
||||
if r == '.' || unicode.IsLower(r) || unicode.IsDigit(r) {
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("lagoon: plugin id %q is not a valid history table name", id)
|
||||
}
|
||||
if strings.Contains(id, "..") || strings.HasPrefix(id, ".") || strings.HasSuffix(id, ".") {
|
||||
return fmt.Errorf("lagoon: plugin id %q is not a valid history table name", id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func migrator(gdb *gorm.DB, pluginID string, migrations []*gormigrate.Migration) (*gormigrate.Gormigrate, error) {
|
||||
table, err := HistoryTableName(pluginID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return gormigrate.New(gdb, &gormigrate.Options{
|
||||
TableName: table,
|
||||
UseTransaction: true,
|
||||
}, migrations), nil
|
||||
}
|
||||
|
||||
// Migrate runs the framework-owned system_files and backend-admin sets
|
||||
// first, then each plugin's HasMigrations set in party.Activate order.
|
||||
func Migrate(gdb *gorm.DB, plugins []party.Plugin) error {
|
||||
if gdb == nil {
|
||||
return fmt.Errorf("lagoon: gorm db is nil")
|
||||
}
|
||||
m, err := migrator(gdb, "summercms.attach", attach.Migrations)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.Migrate(); err != nil {
|
||||
return fmt.Errorf("lagoon: migrate system_files: %w", err)
|
||||
}
|
||||
admin, err := migrator(gdb, "summercms.cabana", BackendAdminMigrations)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := admin.Migrate(); err != nil {
|
||||
return fmt.Errorf("lagoon: migrate backend admin: %w", err)
|
||||
}
|
||||
for _, p := range plugins {
|
||||
hm, ok := p.(pact.HasMigrations)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
migrations := hm.Migrations()
|
||||
if len(migrations) == 0 {
|
||||
continue
|
||||
}
|
||||
m, err := migrator(gdb, p.ID(), migrations)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.Migrate(); err != nil {
|
||||
return fmt.Errorf("lagoon: migrate %s: %w", p.ID(), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RollbackLast rolls back the last migration of pluginID.
|
||||
func RollbackLast(gdb *gorm.DB, plugins []party.Plugin, pluginID string) error {
|
||||
if gdb == nil {
|
||||
return fmt.Errorf("lagoon: gorm db is nil")
|
||||
}
|
||||
pluginID = strings.TrimSpace(pluginID)
|
||||
if pluginID == "" {
|
||||
pluginID = lastMigrationPlugin(plugins)
|
||||
}
|
||||
if pluginID == "" {
|
||||
return ErrNoMigrations
|
||||
}
|
||||
var target party.Plugin
|
||||
for _, p := range plugins {
|
||||
if p.ID() == pluginID {
|
||||
target = p
|
||||
break
|
||||
}
|
||||
}
|
||||
if target == nil {
|
||||
return fmt.Errorf("%w: %q", ErrUnknownPlugin, pluginID)
|
||||
}
|
||||
hm, ok := target.(pact.HasMigrations)
|
||||
if !ok || len(hm.Migrations()) == 0 {
|
||||
return fmt.Errorf("%w: %q", ErrNoMigrations, pluginID)
|
||||
}
|
||||
m, err := migrator(gdb, target.ID(), hm.Migrations())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.RollbackLast(); err != nil {
|
||||
return fmt.Errorf("lagoon: rollback %s: %w", pluginID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func lastMigrationPlugin(plugins []party.Plugin) string {
|
||||
for i := len(plugins) - 1; i >= 0; i-- {
|
||||
p := plugins[i]
|
||||
hm, ok := p.(pact.HasMigrations)
|
||||
if !ok || len(hm.Migrations()) == 0 {
|
||||
continue
|
||||
}
|
||||
return p.ID()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// StatusRow is one plugin's recorded migration ids.
|
||||
type StatusRow struct {
|
||||
Plugin string
|
||||
Table string
|
||||
IDs []string
|
||||
}
|
||||
|
||||
// Status reads each plugin history table.
|
||||
func Status(gdb *gorm.DB, plugins []party.Plugin) ([]StatusRow, error) {
|
||||
if gdb == nil {
|
||||
return nil, fmt.Errorf("lagoon: gorm db is nil")
|
||||
}
|
||||
var rows []StatusRow
|
||||
for _, p := range plugins {
|
||||
if _, ok := p.(pact.HasMigrations); !ok {
|
||||
continue
|
||||
}
|
||||
table, err := HistoryTableName(p.ID())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row := StatusRow{Plugin: p.ID(), Table: table}
|
||||
if !gdb.Migrator().HasTable(table) {
|
||||
rows = append(rows, row)
|
||||
continue
|
||||
}
|
||||
if err := gdb.Table(table).Order("id").Pluck("id", &row.IDs).Error; err != nil {
|
||||
return nil, fmt.Errorf("lagoon: status %s: %w", p.ID(), err)
|
||||
}
|
||||
if row.IDs == nil {
|
||||
row.IDs = []string{}
|
||||
}
|
||||
rows = append(rows, row)
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
277
modules/lagoon/migrations_test.go
Normal file
277
modules/lagoon/migrations_test.go
Normal file
@@ -0,0 +1,277 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/backpack"
|
||||
"git.golem15.com/golem15/summercms/modules/bonfire"
|
||||
"git.golem15.com/golem15/summercms/modules/pact"
|
||||
"git.golem15.com/golem15/summercms/modules/party"
|
||||
"github.com/go-gormigrate/gormigrate/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func TestHistoryTableName(t *testing.T) {
|
||||
got, err := HistoryTableName("golem15.user")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "summer_migrations_golem15_user" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
got, err = HistoryTableName("golem15.fonoteka")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "summer_migrations_golem15_fonoteka" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if _, err := HistoryTableName("golem15.user;drop"); err == nil {
|
||||
t.Fatal("want invalid plugin id error")
|
||||
}
|
||||
if _, err := HistoryTableName(""); err == nil {
|
||||
t.Fatal("want empty plugin id error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckLocaleMessage(t *testing.T) {
|
||||
if err := checkLocale("i", "pl-PL"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err := checkLocale("c", "en_US")
|
||||
if err == nil {
|
||||
t.Fatal("want locale error")
|
||||
}
|
||||
msg := err.Error()
|
||||
for _, want := range []string{"ICU", "pl-PL", "CREATE DATABASE", "LOCALE_PROVIDER icu", "en_US"} {
|
||||
if !strings.Contains(msg, want) {
|
||||
t.Fatalf("missing %q in %s", want, msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeCommandsRegisterBareAndColonNames(t *testing.T) {
|
||||
cmds := RuntimeCommands(nil, nil)
|
||||
names := map[string]bool{}
|
||||
for _, c := range cmds {
|
||||
names[c.Name] = true
|
||||
}
|
||||
for _, want := range []string{"migrate", "migrate:rollback", "migrate:status", "key:generate"} {
|
||||
if !names[want] {
|
||||
t.Fatalf("missing %s", want)
|
||||
}
|
||||
}
|
||||
root, err := bonfire.NewRoot("app", cmds, io.Discard)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root.SetArgs([]string{"migrate", "--help"})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoAutoMigrate(t *testing.T) {
|
||||
entries, err := os.ReadDir(".")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".go") || strings.HasSuffix(entry.Name(), "_test.go") {
|
||||
continue
|
||||
}
|
||||
body, err := os.ReadFile(entry.Name())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if bytes.Contains(body, []byte("AutoMigrate")) {
|
||||
t.Fatalf("%s must not call AutoMigrate", entry.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenRequiresDSN(t *testing.T) {
|
||||
_, _, err := Open(t.Context(), " ")
|
||||
if err == nil || !strings.Contains(err.Error(), "database.dsn") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
type stubPlugin struct{ id string }
|
||||
|
||||
func (s stubPlugin) ID() string { return s.id }
|
||||
func (s stubPlugin) Requires() []string { return nil }
|
||||
func (s stubPlugin) Register(*backpack.App) error { return nil }
|
||||
func (s stubPlugin) Boot(*backpack.App) error { return nil }
|
||||
|
||||
func TestRollbackUnknownPluginIsNamedError(t *testing.T) {
|
||||
err := RollbackLast(&gorm.DB{}, []party.Plugin{stubPlugin{id: "golem15.user"}}, "golem15.missing")
|
||||
if !errors.Is(err, ErrUnknownPlugin) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "golem15.missing") {
|
||||
t.Fatalf("error should name plugin: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackMissingPluginIsNamedError(t *testing.T) {
|
||||
err := RollbackLast(&gorm.DB{}, nil, "")
|
||||
if !errors.Is(err, ErrNoMigrations) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
type migPlugin struct {
|
||||
id string
|
||||
migrations []*gormigrate.Migration
|
||||
}
|
||||
|
||||
func (p migPlugin) ID() string { return p.id }
|
||||
func (p migPlugin) Requires() []string { return nil }
|
||||
func (p migPlugin) Register(*backpack.App) error { return nil }
|
||||
func (p migPlugin) Boot(*backpack.App) error { return nil }
|
||||
func (p migPlugin) Migrations() []*gormigrate.Migration { return p.migrations }
|
||||
|
||||
func TestMigrateRunsSystemFilesFirst(t *testing.T) {
|
||||
db, _ := dedicatedDB(t, "lagoon_attach_first")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("system_files") {
|
||||
t.Fatal("system_files must exist after Migrate with an empty plugin list")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTwoPluginMigrationSetsIsolated(t *testing.T) {
|
||||
db, _ := dedicatedDB(t, "lagoon_mig_iso")
|
||||
gdb, err := Use(t.Context(), db)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
alpha := migPlugin{
|
||||
id: "demo.alpha",
|
||||
migrations: []*gormigrate.Migration{{
|
||||
ID: "202609170001_create_alpha",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`CREATE TABLE lagoon_alpha (id BIGSERIAL PRIMARY KEY, name TEXT NOT NULL)`).Error
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`DROP TABLE IF EXISTS lagoon_alpha`).Error
|
||||
},
|
||||
}},
|
||||
}
|
||||
beta := migPlugin{
|
||||
id: "demo.beta",
|
||||
migrations: []*gormigrate.Migration{
|
||||
{
|
||||
ID: "202609170001_create_beta",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`CREATE TABLE lagoon_beta (id BIGSERIAL PRIMARY KEY, name TEXT NOT NULL)`).Error
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`DROP TABLE IF EXISTS lagoon_beta`).Error
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "202609170002_seed_beta",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`INSERT INTO lagoon_beta (name) VALUES ('seed')`).Error
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
return tx.Exec(`DELETE FROM lagoon_beta WHERE name = 'seed'`).Error
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
plugins := []party.Plugin{alpha, beta}
|
||||
|
||||
if err := Migrate(gdb, plugins); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := Migrate(gdb, plugins); err != nil {
|
||||
t.Fatalf("idempotent migrate: %v", err)
|
||||
}
|
||||
|
||||
status, err := Status(gdb, plugins)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(status) != 2 {
|
||||
t.Fatalf("status rows = %+v", status)
|
||||
}
|
||||
if status[0].Plugin != "demo.alpha" || status[0].Table != "summer_migrations_demo_alpha" || strings.Join(status[0].IDs, ",") != "202609170001_create_alpha" {
|
||||
t.Fatalf("alpha status = %+v", status[0])
|
||||
}
|
||||
if status[1].Plugin != "demo.beta" || status[1].Table != "summer_migrations_demo_beta" || strings.Join(status[1].IDs, ",") != "202609170001_create_beta,202609170002_seed_beta" {
|
||||
t.Fatalf("beta status = %+v", status[1])
|
||||
}
|
||||
|
||||
if !gdb.Migrator().HasTable("lagoon_alpha") || !gdb.Migrator().HasTable("lagoon_beta") {
|
||||
t.Fatal("both plugin tables must exist after migrate")
|
||||
}
|
||||
var seedName string
|
||||
if err := gdb.Raw(`SELECT name FROM lagoon_beta`).Scan(&seedName).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if seedName != "seed" {
|
||||
t.Fatalf("seed = %q", seedName)
|
||||
}
|
||||
|
||||
if err := RollbackLast(gdb, plugins, "demo.beta"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int64
|
||||
if err := gdb.Table("lagoon_beta").Count(&n).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("beta seed rollback left %d rows", n)
|
||||
}
|
||||
if !gdb.Migrator().HasTable("lagoon_alpha") || !gdb.Migrator().HasTable("lagoon_beta") {
|
||||
t.Fatal("schema rollback must not run on the first beta RollbackLast")
|
||||
}
|
||||
|
||||
status, err = Status(gdb, plugins)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Join(status[0].IDs, ",") != "202609170001_create_alpha" {
|
||||
t.Fatalf("alpha history changed: %+v", status[0])
|
||||
}
|
||||
if strings.Join(status[1].IDs, ",") != "202609170001_create_beta" {
|
||||
t.Fatalf("beta history = %+v", status[1])
|
||||
}
|
||||
|
||||
if err := RollbackLast(gdb, plugins, "demo.beta"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if gdb.Migrator().HasTable("lagoon_beta") {
|
||||
t.Fatal("second beta rollback must drop lagoon_beta")
|
||||
}
|
||||
if !gdb.Migrator().HasTable("lagoon_alpha") {
|
||||
t.Fatal("alpha table must survive beta rollback")
|
||||
}
|
||||
|
||||
status, err = Status(gdb, plugins)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(status[0].IDs) != 1 {
|
||||
t.Fatalf("alpha history after beta down = %+v", status[0])
|
||||
}
|
||||
if len(status[1].IDs) != 0 {
|
||||
t.Fatalf("beta history should be empty, got %+v", status[1])
|
||||
}
|
||||
}
|
||||
|
||||
var _ pact.HasMigrations = migPlugin{}
|
||||
46
modules/lagoon/order.go
Normal file
46
modules/lagoon/order.go
Normal file
@@ -0,0 +1,46 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// OrderBy appends a database-default ORDER BY for an allow-listed qualified
|
||||
// column. Identifiers and directions are never taken from untrusted input:
|
||||
// column must match allowed exactly, and dir must be asc or desc. No COLLATE
|
||||
// is emitted; Postgres ICU pl-PL is the database default (CheckLocale).
|
||||
func OrderBy(db *gorm.DB, column, dir string, allowed []string) (*gorm.DB, error) {
|
||||
if db == nil {
|
||||
return nil, fmt.Errorf("lagoon: gorm db is nil")
|
||||
}
|
||||
clause, err := orderClause(column, dir, allowed)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return db.Order(clause), nil
|
||||
}
|
||||
|
||||
func orderClause(column, dir string, allowed []string) (string, error) {
|
||||
if !allowListed(column, allowed) {
|
||||
return "", fmt.Errorf("lagoon: order column %q is not allow-listed", column)
|
||||
}
|
||||
switch strings.ToLower(strings.TrimSpace(dir)) {
|
||||
case "asc":
|
||||
return column + " ASC", nil
|
||||
case "desc":
|
||||
return column + " DESC", nil
|
||||
default:
|
||||
return "", fmt.Errorf("lagoon: order direction %q is not allow-listed", dir)
|
||||
}
|
||||
}
|
||||
|
||||
func allowListed(column string, allowed []string) bool {
|
||||
for _, a := range allowed {
|
||||
if a == column {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
42
modules/lagoon/order_test.go
Normal file
42
modules/lagoon/order_test.go
Normal file
@@ -0,0 +1,42 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestOrderClauseAllowList(t *testing.T) {
|
||||
allowed := []string{"golem15_fonoteka_genres.name", "items.title"}
|
||||
|
||||
got, err := orderClause("golem15_fonoteka_genres.name", "asc", allowed)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "golem15_fonoteka_genres.name ASC" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if strings.Contains(strings.ToLower(got), "collate") {
|
||||
t.Fatalf("must not emit COLLATE: %q", got)
|
||||
}
|
||||
|
||||
got, err = orderClause("items.title", "DESC", allowed)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "items.title DESC" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
|
||||
if _, err := orderClause("golem15_fonoteka_genres.name;drop table x", "asc", allowed); err == nil {
|
||||
t.Fatal("want reject unknown column")
|
||||
}
|
||||
if _, err := orderClause("golem15_fonoteka_genres.name", "ascending", allowed); err == nil {
|
||||
t.Fatal("want reject unknown direction")
|
||||
}
|
||||
if _, err := OrderBy(nil, "items.title", "asc", allowed); err == nil {
|
||||
t.Fatal("want nil db error")
|
||||
}
|
||||
if _, err := orderClause("items.title", "asc", nil); err == nil {
|
||||
t.Fatal("empty allow-list must reject")
|
||||
}
|
||||
}
|
||||
40
modules/lagoon/paginate.go
Normal file
40
modules/lagoon/paginate.go
Normal file
@@ -0,0 +1,40 @@
|
||||
package lagoon
|
||||
|
||||
// PageMeta is the pagination block of the house REST envelope (DATA-10).
|
||||
type PageMeta struct {
|
||||
CurrentPage int `json:"current_page"`
|
||||
LastPage int `json:"last_page"`
|
||||
PerPage int `json:"per_page"`
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
|
||||
// Page wraps already-built row DTOs. There is no links field.
|
||||
type Page[T any] struct {
|
||||
Data []T `json:"data"`
|
||||
Meta PageMeta `json:"meta"`
|
||||
}
|
||||
|
||||
// Paginate builds {data, meta{current_page,last_page,per_page,total}} from
|
||||
// a caller-sliced row set. LastPage is ceil(total/perPage); perPage<=0
|
||||
// yields LastPage 1 instead of dividing by zero. Nil rows become [].
|
||||
func Paginate[T any](rows []T, page, perPage int, total int64) Page[T] {
|
||||
last := 1
|
||||
if perPage > 0 {
|
||||
last = int((total + int64(perPage) - 1) / int64(perPage))
|
||||
if last < 1 {
|
||||
last = 1
|
||||
}
|
||||
}
|
||||
if rows == nil {
|
||||
rows = []T{}
|
||||
}
|
||||
return Page[T]{
|
||||
Data: rows,
|
||||
Meta: PageMeta{
|
||||
CurrentPage: page,
|
||||
LastPage: last,
|
||||
PerPage: perPage,
|
||||
Total: total,
|
||||
},
|
||||
}
|
||||
}
|
||||
42
modules/lagoon/paginate_test.go
Normal file
42
modules/lagoon/paginate_test.go
Normal file
@@ -0,0 +1,42 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPaginateEnvelope(t *testing.T) {
|
||||
got := Paginate([]string{"a", "b", "c"}, 1, 2, 3)
|
||||
if got.Meta.CurrentPage != 1 || got.Meta.LastPage != 2 || got.Meta.PerPage != 2 || got.Meta.Total != 3 {
|
||||
t.Fatalf("meta = %+v", got.Meta)
|
||||
}
|
||||
raw, err := json.Marshal(got)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(string(raw), "links") {
|
||||
t.Fatalf("envelope must not contain links: %s", raw)
|
||||
}
|
||||
if !strings.Contains(string(raw), `"data"`) || !strings.Contains(string(raw), `"meta"`) {
|
||||
t.Fatalf("missing data/meta: %s", raw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPaginateZeroPerPage(t *testing.T) {
|
||||
got := Paginate([]int{1}, 1, 0, 10)
|
||||
if got.Meta.LastPage != 1 {
|
||||
t.Fatalf("last_page = %d, want 1", got.Meta.LastPage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPaginateNilDataIsEmptyArray(t *testing.T) {
|
||||
got := Paginate[string](nil, 1, 10, 0)
|
||||
raw, err := json.Marshal(got)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(raw), `"data":[]`) {
|
||||
t.Fatalf("nil rows must marshal as []: %s", raw)
|
||||
}
|
||||
}
|
||||
139
modules/lagoon/postgres_test.go
Normal file
139
modules/lagoon/postgres_test.go
Normal file
@@ -0,0 +1,139 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"github.com/testcontainers/testcontainers-go"
|
||||
"github.com/testcontainers/testcontainers-go/modules/postgres"
|
||||
)
|
||||
|
||||
var (
|
||||
lagoonPG *postgres.PostgresContainer
|
||||
lagoonSQL *sql.DB
|
||||
lagoonDSN string
|
||||
lagoonPGErr error
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
code := 1
|
||||
if !testShort() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
lagoonPGErr = startLagoonPostgres(ctx)
|
||||
cancel()
|
||||
if lagoonPGErr != nil {
|
||||
fmt.Fprintf(os.Stderr, "lagoon: testcontainers postgres: %v\n", lagoonPGErr)
|
||||
stopLagoonPostgres()
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
code = m.Run()
|
||||
stopLagoonPostgres()
|
||||
os.Exit(code)
|
||||
}
|
||||
|
||||
func testShort() bool {
|
||||
for _, a := range os.Args {
|
||||
if a == "-test.short" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func startLagoonPostgres(ctx context.Context) error {
|
||||
ctr, err := postgres.Run(ctx,
|
||||
"postgres:16-alpine",
|
||||
postgres.WithDatabase("lagoon"),
|
||||
postgres.WithUsername("lagoon"),
|
||||
postgres.WithPassword("lagoon"),
|
||||
postgres.BasicWaitStrategies(),
|
||||
testcontainers.WithEnv(map[string]string{
|
||||
"POSTGRES_INITDB_ARGS": "--locale-provider=icu --icu-locale=pl-PL --encoding=UTF8",
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
lagoonPG = ctr
|
||||
dsn, err := ctr.ConnectionString(ctx, "sslmode=disable")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
_ = db.Close()
|
||||
return err
|
||||
}
|
||||
lagoonSQL = db
|
||||
lagoonDSN = dsn
|
||||
return nil
|
||||
}
|
||||
|
||||
func stopLagoonPostgres() {
|
||||
if lagoonSQL != nil {
|
||||
_ = lagoonSQL.Close()
|
||||
}
|
||||
if lagoonPG != nil {
|
||||
_ = testcontainers.TerminateContainer(lagoonPG)
|
||||
}
|
||||
}
|
||||
|
||||
func lagoonDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
if testing.Short() {
|
||||
t.Skip("requires testcontainers postgres")
|
||||
}
|
||||
if lagoonPGErr != nil {
|
||||
t.Fatalf("postgres unavailable: %v", lagoonPGErr)
|
||||
}
|
||||
if lagoonSQL == nil {
|
||||
t.Fatal("postgres unavailable: container was not started")
|
||||
}
|
||||
return lagoonSQL
|
||||
}
|
||||
|
||||
func dedicatedDB(t *testing.T, name string) (*sql.DB, string) {
|
||||
t.Helper()
|
||||
admin := lagoonDB(t)
|
||||
ctx := t.Context()
|
||||
quoted := `"` + strings.ReplaceAll(name, `"`, `""`) + `"`
|
||||
if _, err := admin.ExecContext(ctx, `CREATE DATABASE `+quoted+` TEMPLATE template0 ENCODING 'UTF8' LOCALE_PROVIDER icu ICU_LOCALE 'pl-PL'`); err != nil && !strings.Contains(err.Error(), "already exists") {
|
||||
t.Fatalf("create %s: %v", name, err)
|
||||
}
|
||||
dsn, err := dsnWithDB(lagoonDSN, name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = db.Close()
|
||||
_, _ = admin.ExecContext(context.Background(), `DROP DATABASE IF EXISTS `+quoted+` WITH (FORCE)`)
|
||||
})
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db, dsn
|
||||
}
|
||||
|
||||
func dsnWithDB(dsn, name string) (string, error) {
|
||||
u, err := url.Parse(dsn)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
u.Path = "/" + name
|
||||
return u.String(), nil
|
||||
}
|
||||
28
modules/lagoon/relations.go
Normal file
28
modules/lagoon/relations.go
Normal file
@@ -0,0 +1,28 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// RegisterJoinTable wraps db.SetupJoinTable so a nil db returns an error
|
||||
// instead of panicking.
|
||||
//
|
||||
// Pivot-write contract: pivot tables with business columns (sort_order,
|
||||
// role, granted_at, granted_by) are never written through
|
||||
// db.Model(&owner).Association(field).Append/Replace(...). GORM's
|
||||
// Association Mode has no documented path to set those columns. Writes go
|
||||
// through an explicit delete-then-bulk-insert (or ON CONFLICT DO UPDATE)
|
||||
// against the join table directly, in the same transaction as the parent
|
||||
// save. Reads use Preload(field) plus .Order(...) on the pivot's own
|
||||
// columns once RegisterJoinTable has run.
|
||||
func RegisterJoinTable(db *gorm.DB, owner any, field string, joinModel any) error {
|
||||
if db == nil {
|
||||
return fmt.Errorf("lagoon: register join table: db is nil")
|
||||
}
|
||||
if err := db.SetupJoinTable(owner, field, joinModel); err != nil {
|
||||
return fmt.Errorf("lagoon: register join table: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
16
modules/lagoon/relations_test.go
Normal file
16
modules/lagoon/relations_test.go
Normal file
@@ -0,0 +1,16 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRegisterJoinTableNilDB(t *testing.T) {
|
||||
err := RegisterJoinTable(nil, struct{}{}, "Artists", struct{}{})
|
||||
if err == nil {
|
||||
t.Fatal("want non-nil error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "nil") {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
477
modules/lagoon/validate.go
Normal file
477
modules/lagoon/validate.go
Normal file
@@ -0,0 +1,477 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math"
|
||||
"math/big"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/phrasebook"
|
||||
"github.com/go-playground/validator/v10"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var identName = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`)
|
||||
|
||||
var validateOnce = validator.New(validator.WithRequiredStructEnabled())
|
||||
|
||||
// Validate translates Laravel-style rule strings onto validator.Var() plus a
|
||||
// unique:table DB check. Unrecognized tokens fail loudly. The returned map is
|
||||
// the Laravel-shaped errors object (field -> messages); the HTTP envelope is
|
||||
// a later-phase concern.
|
||||
func Validate(ctx context.Context, tx *gorm.DB, model any, rules map[string]string, values map[string]any, tr *phrasebook.Translator) (map[string][]string, error) {
|
||||
out := map[string][]string{}
|
||||
for field, rule := range rules {
|
||||
msgs, err := validateField(ctx, tx, model, field, rule, values, tr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(msgs) > 0 {
|
||||
out[field] = msgs
|
||||
}
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func validateField(ctx context.Context, tx *gorm.DB, model any, field, rule string, values map[string]any, tr *phrasebook.Translator) ([]string, error) {
|
||||
val := values[field]
|
||||
tokens := splitRule(rule)
|
||||
nullable := false
|
||||
required := false
|
||||
var tags []string
|
||||
var uniqueTable string
|
||||
var betweenMin, betweenMax, minArg, maxArg string
|
||||
hasNumeric := false
|
||||
hasInteger := false
|
||||
for _, tok := range tokens {
|
||||
name, arg, _ := strings.Cut(tok, ":")
|
||||
switch name {
|
||||
case "nullable":
|
||||
nullable = true
|
||||
case "required":
|
||||
required = true
|
||||
tags = append(tags, "required")
|
||||
case "integer":
|
||||
hasInteger = true
|
||||
if !isIntegerValue(val) && !isEmptyValue(val) {
|
||||
return []string{validateMessage(ctx, tr, "integer", field, nil)}, nil
|
||||
}
|
||||
case "numeric":
|
||||
hasNumeric = true
|
||||
tags = append(tags, "numeric")
|
||||
case "between":
|
||||
x, y, ok := strings.Cut(arg, ",")
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("lagoon: unrecognized validation rule %q", tok)
|
||||
}
|
||||
betweenMin, betweenMax = x, y
|
||||
case "min":
|
||||
minArg = arg
|
||||
case "max":
|
||||
maxArg = arg
|
||||
case "in":
|
||||
tags = append(tags, oneofTag(strings.Split(arg, ",")))
|
||||
case "unique":
|
||||
uniqueTable = arg
|
||||
case "boolean":
|
||||
if !isLaravelBoolean(val) {
|
||||
return []string{validateMessage(ctx, tr, "boolean", field, nil)}, nil
|
||||
}
|
||||
case "email":
|
||||
tags = append(tags, "email")
|
||||
case "confirmed":
|
||||
if fmt.Sprint(val) != fmt.Sprint(values[field+"_confirmation"]) {
|
||||
return []string{validateMessage(ctx, tr, "confirmed", field, nil)}, nil
|
||||
}
|
||||
case "different":
|
||||
if fmt.Sprint(val) == fmt.Sprint(values[arg]) {
|
||||
return []string{validateMessage(ctx, tr, "different", field, map[string]string{"other": arg})}, nil
|
||||
}
|
||||
case "mimes":
|
||||
got := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(fmt.Sprint(val))), ".")
|
||||
match := false
|
||||
for _, ext := range strings.Split(arg, ",") {
|
||||
want := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(ext)), ".")
|
||||
if got != "" && got != "<nil>" && got == want {
|
||||
match = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !match {
|
||||
return []string{validateMessage(ctx, tr, "mimes", field, nil)}, nil
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("lagoon: unrecognized validation rule %q", tok)
|
||||
}
|
||||
}
|
||||
if nullable && isEmptyValue(val) && !required {
|
||||
return nil, nil
|
||||
}
|
||||
numericRange := hasNumeric || hasInteger
|
||||
var rangeMin, rangeMax string
|
||||
if numericRange {
|
||||
if betweenMin != "" {
|
||||
rangeMin, rangeMax = betweenMin, betweenMax
|
||||
}
|
||||
if minArg != "" {
|
||||
rangeMin = minArg
|
||||
}
|
||||
if maxArg != "" {
|
||||
rangeMax = maxArg
|
||||
}
|
||||
} else {
|
||||
if betweenMin != "" {
|
||||
tags = append(tags, "min="+betweenMin, "max="+betweenMax)
|
||||
}
|
||||
if minArg != "" {
|
||||
tags = append(tags, "min="+minArg)
|
||||
}
|
||||
if maxArg != "" {
|
||||
tags = append(tags, "max="+maxArg)
|
||||
}
|
||||
}
|
||||
if numericRange && (rangeMin != "" || rangeMax != "") {
|
||||
s, ok := numericString(val)
|
||||
if !ok {
|
||||
ruleName := "numeric"
|
||||
if hasInteger {
|
||||
ruleName = "integer"
|
||||
}
|
||||
return []string{validateMessage(ctx, tr, ruleName, field, nil)}, nil
|
||||
}
|
||||
if !moneyInRange(s, rangeMin, rangeMax) {
|
||||
return []string{validateMessage(ctx, tr, "max", field, map[string]string{"max": rangeMax, "min": rangeMin})}, nil
|
||||
}
|
||||
tags = withoutTag(tags, "numeric")
|
||||
}
|
||||
if required && isEmptyValue(val) {
|
||||
return []string{validateMessage(ctx, tr, "required", field, nil)}, nil
|
||||
}
|
||||
usedBetween := betweenMin != "" && betweenMax != ""
|
||||
if len(tags) > 0 {
|
||||
msgs := make([]string, 0, len(tags))
|
||||
for _, t := range tags {
|
||||
name, _, _ := strings.Cut(t, "=")
|
||||
if name == "required" || name == "min" || name == "max" {
|
||||
continue
|
||||
}
|
||||
if err := validateOnce.Var(val, t); err != nil {
|
||||
ruleName := name
|
||||
if name == "oneof" {
|
||||
ruleName = "oneof"
|
||||
}
|
||||
msgs = append(msgs, validateMessage(ctx, tr, ruleName, field, map[string]string{
|
||||
"min": betweenMin, "max": betweenMax,
|
||||
}))
|
||||
}
|
||||
}
|
||||
if usedBetween && !numericRange {
|
||||
n := 0
|
||||
if s, ok := val.(string); ok {
|
||||
n = len(s)
|
||||
} else if !isEmptyValue(val) {
|
||||
n = len(fmt.Sprint(val))
|
||||
}
|
||||
minN, maxN := atoiOr(betweenMin, 0), atoiOr(betweenMax, 0)
|
||||
if n < minN || (maxN > 0 && n > maxN) {
|
||||
msgs = append(msgs, validateMessage(ctx, tr, "between", field, map[string]string{
|
||||
"min": betweenMin, "max": betweenMax,
|
||||
}))
|
||||
}
|
||||
}
|
||||
if len(msgs) > 0 {
|
||||
return msgs, nil
|
||||
}
|
||||
}
|
||||
if uniqueTable != "" {
|
||||
ok, err := uniqueOK(tx, model, uniqueTable, field, val)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !ok {
|
||||
return []string{validateMessage(ctx, tr, "unique", field, nil)}, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func splitRule(rule string) []string {
|
||||
var out []string
|
||||
for _, p := range strings.Split(rule, "|") {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
out = append(out, p)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func oneofTag(vals []string) string {
|
||||
parts := make([]string, 0, len(vals))
|
||||
for _, v := range vals {
|
||||
v = strings.TrimSpace(v)
|
||||
if v == "" {
|
||||
continue
|
||||
}
|
||||
if strings.ContainsAny(v, " \t\"'") {
|
||||
parts = append(parts, "'"+v+"'")
|
||||
} else {
|
||||
parts = append(parts, v)
|
||||
}
|
||||
}
|
||||
return "oneof=" + strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
func isEmptyValue(val any) bool {
|
||||
if val == nil {
|
||||
return true
|
||||
}
|
||||
rv := reflect.ValueOf(val)
|
||||
switch rv.Kind() {
|
||||
case reflect.Ptr, reflect.Interface:
|
||||
return rv.IsNil()
|
||||
case reflect.String, reflect.Slice, reflect.Map:
|
||||
return rv.Len() == 0
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// isIntegerValue reports whether val is an integer in Laravel's sense. It
|
||||
// dereferences pointers (model fields such as *int) and accepts whole
|
||||
// floating-point numbers, the shape encoding/json gives a map[string]any.
|
||||
func isIntegerValue(val any) bool {
|
||||
if val == nil {
|
||||
return true
|
||||
}
|
||||
rv := reflect.ValueOf(val)
|
||||
for rv.Kind() == reflect.Ptr || rv.Kind() == reflect.Interface {
|
||||
if rv.IsNil() {
|
||||
return true
|
||||
}
|
||||
rv = rv.Elem()
|
||||
}
|
||||
switch rv.Kind() {
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return true
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||
return true
|
||||
case reflect.Float32, reflect.Float64:
|
||||
f := rv.Float()
|
||||
return f == math.Trunc(f) && !math.IsInf(f, 0)
|
||||
case reflect.String:
|
||||
s := rv.String()
|
||||
if s == "" {
|
||||
return true
|
||||
}
|
||||
_, ok := new(big.Int).SetString(s, 10)
|
||||
return ok
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func numericString(val any) (string, bool) {
|
||||
if val == nil {
|
||||
return "", false
|
||||
}
|
||||
rv := reflect.ValueOf(val)
|
||||
for rv.Kind() == reflect.Ptr || rv.Kind() == reflect.Interface {
|
||||
if rv.IsNil() {
|
||||
return "", false
|
||||
}
|
||||
rv = rv.Elem()
|
||||
}
|
||||
switch rv.Kind() {
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return strconv.FormatInt(rv.Int(), 10), true
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||
return strconv.FormatUint(rv.Uint(), 10), true
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return strconv.FormatFloat(rv.Float(), 'f', -1, 64), true
|
||||
case reflect.String:
|
||||
s := strings.TrimSpace(rv.String())
|
||||
if s == "" {
|
||||
return "", false
|
||||
}
|
||||
return s, true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
func moneyInRange(val any, min, max string) bool {
|
||||
s, ok := numericString(val)
|
||||
if !ok {
|
||||
s = strings.TrimSpace(fmt.Sprint(val))
|
||||
if s == "" || s == "<nil>" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
r := new(big.Rat)
|
||||
if _, ok := r.SetString(s); !ok {
|
||||
return false
|
||||
}
|
||||
if min != "" {
|
||||
m := new(big.Rat)
|
||||
if _, ok := m.SetString(min); ok && r.Cmp(m) < 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if max != "" {
|
||||
m := new(big.Rat)
|
||||
if _, ok := m.SetString(max); ok && r.Cmp(m) > 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func uniqueOK(tx *gorm.DB, model any, table, column string, val any) (bool, error) {
|
||||
if tx == nil {
|
||||
return false, fmt.Errorf("lagoon: unique:%s requires a database handle", table)
|
||||
}
|
||||
if !identName.MatchString(table) || !identName.MatchString(column) {
|
||||
return false, fmt.Errorf("lagoon: unique identifier %q.%q is not safe", table, column)
|
||||
}
|
||||
if isEmptyValue(val) {
|
||||
return true, nil
|
||||
}
|
||||
q := tx.Table(table).Where(column+" = ?", val)
|
||||
if tx.Migrator().HasColumn(table, "deleted_at") {
|
||||
q = q.Where("deleted_at IS NULL")
|
||||
}
|
||||
if id := modelUintID(model); id != 0 && tx.Migrator().HasColumn(table, "id") {
|
||||
q = q.Where("id <> ?", id)
|
||||
}
|
||||
var n int64
|
||||
if err := q.Count(&n).Error; err != nil {
|
||||
return false, err
|
||||
}
|
||||
return n == 0, nil
|
||||
}
|
||||
|
||||
func modelUintID(model any) uint {
|
||||
if model == nil {
|
||||
return 0
|
||||
}
|
||||
rv := reflect.ValueOf(model)
|
||||
if rv.Kind() == reflect.Ptr {
|
||||
if rv.IsNil() {
|
||||
return 0
|
||||
}
|
||||
rv = rv.Elem()
|
||||
}
|
||||
if rv.Kind() != reflect.Struct {
|
||||
return 0
|
||||
}
|
||||
f := rv.FieldByName("ID")
|
||||
if !f.IsValid() {
|
||||
return 0
|
||||
}
|
||||
switch f.Kind() {
|
||||
case reflect.Uint, reflect.Uint32, reflect.Uint64:
|
||||
return uint(f.Uint())
|
||||
case reflect.Int, reflect.Int32, reflect.Int64:
|
||||
if f.Int() < 0 {
|
||||
return 0
|
||||
}
|
||||
return uint(f.Int())
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func withoutTag(tags []string, name string) []string {
|
||||
out := tags[:0]
|
||||
for _, t := range tags {
|
||||
n, _, _ := strings.Cut(t, "=")
|
||||
if n != name {
|
||||
out = append(out, t)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func validateMessage(ctx context.Context, tr *phrasebook.Translator, rule, field string, params map[string]string) string {
|
||||
if params == nil {
|
||||
params = map[string]string{}
|
||||
}
|
||||
params["attribute"] = laravelAttribute(field)
|
||||
key := "lagoon::validate." + rule
|
||||
if tr != nil {
|
||||
s := tr.Get(ctx, key, params)
|
||||
if s != "" && s != key {
|
||||
return s
|
||||
}
|
||||
}
|
||||
attr := params["attribute"]
|
||||
switch rule {
|
||||
case "required":
|
||||
return "The " + attr + " field is required."
|
||||
case "integer":
|
||||
return "The " + attr + " must be an integer."
|
||||
case "numeric":
|
||||
return "The " + attr + " must be a number."
|
||||
case "unique":
|
||||
return "The " + attr + " has already been taken."
|
||||
case "max":
|
||||
return "The " + attr + " may not be greater than " + params["max"] + "."
|
||||
case "min":
|
||||
return "The " + attr + " must be at least " + params["min"] + "."
|
||||
case "oneof":
|
||||
return "The selected " + attr + " is invalid."
|
||||
case "email":
|
||||
return "The " + attr + " must be a valid email address."
|
||||
case "confirmed":
|
||||
return "The " + attr + " confirmation does not match."
|
||||
case "different":
|
||||
return "The " + attr + " and " + params["other"] + " must be different."
|
||||
case "mimes":
|
||||
return "The " + attr + " must be a file of the allowed types."
|
||||
case "between":
|
||||
return "The " + attr + " must be between " + params["min"] + " and " + params["max"] + " characters."
|
||||
case "boolean":
|
||||
return "The " + attr + " field must be true or false."
|
||||
default:
|
||||
return "The " + attr + " is invalid."
|
||||
}
|
||||
}
|
||||
|
||||
func laravelAttribute(field string) string {
|
||||
return strings.ReplaceAll(field, "_", " ")
|
||||
}
|
||||
|
||||
func isLaravelBoolean(val any) bool {
|
||||
switch v := val.(type) {
|
||||
case bool:
|
||||
return true
|
||||
case string:
|
||||
switch strings.ToLower(strings.TrimSpace(v)) {
|
||||
case "0", "1", "true", "false":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
case float64:
|
||||
return v == 0 || v == 1
|
||||
case int:
|
||||
return v == 0 || v == 1
|
||||
default:
|
||||
s := strings.TrimSpace(fmt.Sprint(val))
|
||||
return s == "0" || s == "1" || s == "true" || s == "false"
|
||||
}
|
||||
}
|
||||
|
||||
func atoiOr(s string, fallback int) int {
|
||||
n, err := strconv.Atoi(strings.TrimSpace(s))
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return n
|
||||
}
|
||||
231
modules/lagoon/validate_test.go
Normal file
231
modules/lagoon/validate_test.go
Normal file
@@ -0,0 +1,231 @@
|
||||
package lagoon
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateBetweenYear(t *testing.T) {
|
||||
rules := map[string]string{"year": "nullable|integer|between:1889,2100"}
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"year": 1700}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs["year"]) == 0 {
|
||||
t.Fatal("year=1700 must fail")
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"year": nil}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("nil year = %v", errs)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"year": 1991}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("1991 = %v", errs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateIntegerPointerAndJSONNumber(t *testing.T) {
|
||||
rules := map[string]string{"year": "nullable|integer|between:1889,2100"}
|
||||
year := 1991
|
||||
var nilYear *int
|
||||
pass := map[string]any{
|
||||
"*int": &year,
|
||||
"nil *int": nilYear,
|
||||
"float64": float64(1991),
|
||||
"digit string": "1991",
|
||||
}
|
||||
for name, val := range pass {
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"year": val}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", name, err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("%s must pass, got %v", name, errs)
|
||||
}
|
||||
}
|
||||
zero := 0
|
||||
early := 1700
|
||||
fail := map[string]any{
|
||||
"fractional float": 1991.5,
|
||||
"out of range *int": &early,
|
||||
"bool": true,
|
||||
"non-digit string": "19x1",
|
||||
"int 0": 0,
|
||||
"float64 0": float64(0),
|
||||
"*int 0": &zero,
|
||||
"string 0": "0",
|
||||
}
|
||||
for name, val := range fail {
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"year": val}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", name, err)
|
||||
}
|
||||
if len(errs["year"]) == 0 {
|
||||
t.Fatalf("%s must fail", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateOneofQuotedSpace(t *testing.T) {
|
||||
rules := map[string]string{"format": `nullable|in:LP,2LP,CD,2CD,MC,Box,EP 7"`}
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"format": `EP 7"`}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("EP 7\" rejected: %v", errs)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"format": "tape"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs["format"]) == 0 {
|
||||
t.Fatal("invalid format must fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateMoneyRange(t *testing.T) {
|
||||
rules := map[string]string{"market_price_stored": "nullable|numeric|min:0|max:999999.9999"}
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"market_price_stored": "1000000.0000"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs["market_price_stored"]) == 0 {
|
||||
t.Fatal("1000000.0000 must fail max")
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"market_price_stored": "25.0000"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("25.0000 = %v", errs)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"market_price_stored": 0}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("numeric 0 must pass min:0, got %v", errs)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"market_price_stored": nil}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("nil money = %v", errs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateBooleanNoop(t *testing.T) {
|
||||
rules := map[string]string{"search_use_typesense": "boolean"}
|
||||
for _, v := range []any{true, false} {
|
||||
errs, err := Validate(t.Context(), nil, nil, rules, map[string]any{"search_use_typesense": v}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("boolean %v = %v", v, errs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateEmailConfirmedDifferentMimes(t *testing.T) {
|
||||
errs, err := Validate(t.Context(), nil, nil, map[string]string{"email": "email"}, map[string]any{"email": "not-an-email"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs["email"]) == 0 {
|
||||
t.Fatal("not-an-email must fail")
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, map[string]string{"email": "email"}, map[string]any{"email": "a@b.com"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("a@b.com = %v", errs)
|
||||
}
|
||||
|
||||
rules := map[string]string{"password": "confirmed"}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{
|
||||
"password": "secret", "password_confirmation": "secret",
|
||||
}, nil)
|
||||
if err != nil || len(errs) != 0 {
|
||||
t.Fatalf("confirmed match: %v %v", errs, err)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{
|
||||
"password": "secret", "password_confirmation": "other",
|
||||
}, nil)
|
||||
if err != nil || len(errs["password"]) == 0 {
|
||||
t.Fatalf("confirmed mismatch: %v %v", errs, err)
|
||||
}
|
||||
|
||||
rules = map[string]string{"password": "different:current_password"}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{
|
||||
"password": "same", "current_password": "same",
|
||||
}, nil)
|
||||
if err != nil || len(errs["password"]) == 0 {
|
||||
t.Fatalf("different equal: %v %v", errs, err)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{
|
||||
"password": "new", "current_password": "old",
|
||||
}, nil)
|
||||
if err != nil || len(errs) != 0 {
|
||||
t.Fatalf("different distinct: %v %v", errs, err)
|
||||
}
|
||||
|
||||
rules = map[string]string{"avatar": "mimes:jpeg,jpg,png,webp,gif"}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"avatar": "PNG"}, nil)
|
||||
if err != nil || len(errs) != 0 {
|
||||
t.Fatalf("png: %v %v", errs, err)
|
||||
}
|
||||
errs, err = Validate(t.Context(), nil, nil, rules, map[string]any{"avatar": "svg"}, nil)
|
||||
if err != nil || len(errs["avatar"]) == 0 {
|
||||
t.Fatalf("svg: %v %v", errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateUnrecognizedRule(t *testing.T) {
|
||||
_, err := Validate(t.Context(), nil, nil, map[string]string{"name": "required|nope"}, map[string]any{"name": "x"}, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "nope") {
|
||||
t.Fatalf("want unrecognized nope, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateUniqueRespectsDeletedAt(t *testing.T) {
|
||||
sqlDB, _ := dedicatedDB(t, "lagoon_validate_unique")
|
||||
gdb, err := Use(t.Context(), sqlDB)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`CREATE TABLE lagoon_unique_rows (
|
||||
id SERIAL PRIMARY KEY,
|
||||
slug TEXT NOT NULL,
|
||||
deleted_at TIMESTAMPTZ
|
||||
)`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := gdb.Exec(`INSERT INTO lagoon_unique_rows (slug, deleted_at) VALUES ('live', NULL), ('gone', NOW())`).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rules := map[string]string{"slug": "unique:lagoon_unique_rows"}
|
||||
errs, err := Validate(t.Context(), gdb, nil, rules, map[string]any{"slug": "live"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs["slug"]) == 0 {
|
||||
t.Fatal("live slug must fail unique")
|
||||
}
|
||||
errs, err = Validate(t.Context(), gdb, nil, rules, map[string]any{"slug": "gone"}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(errs) != 0 {
|
||||
t.Fatalf("soft-deleted slug should pass unique: %v", errs)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user