diff --git a/generator/templates/client.gotpl b/generator/templates/client.gotpl index 25eb8a0..207ea2b 100644 --- a/generator/templates/client.gotpl +++ b/generator/templates/client.gotpl @@ -64,15 +64,26 @@ func Open(provider, dataSourceName string) (*DB, error) { {{ $enum.Name }}: {{ $enum.Name }}, {{- end }} } - {{- range $model := .Schema.Models }} - q.{{ $model.Name }} = &{{ $model.Name }}Delegate{client: q} - {{- end }} + q.initDelegates() return &DB{ Queries: q, sqlDB: sqlDB, }, nil } +func (q *Queries) initDelegates() { + {{- range $model := .Schema.Models }} + q.{{ $model.Name }} = &{{ $model.Name }}Delegate{client: q} + {{- end }} +} + +func (q *Queries) copyHooksFrom(other *Queries) { + {{- range $model := .Schema.Models }} + q.{{ $model.Name }}.beforeCreate = other.{{ $model.Name }}.beforeCreate + q.{{ $model.Name }}.afterCreate = other.{{ $model.Name }}.afterCreate + {{- end }} +} + // Close closes the database connection. func (db *DB) Close() error { return db.sqlDB.Close() @@ -216,9 +227,8 @@ func (q *Queries) transaction(ctx context.Context, fn func(txQ *Queries) error) {{ $enum.Name }}: q.{{ $enum.Name }}, {{- end }} } - {{- range $model := .Schema.Models }} - txQueries.{{ $model.Name }} = &{{ $model.Name }}Delegate{client: txQueries} - {{- end }} + txQueries.initDelegates() + txQueries.copyHooksFrom(q) if err := fn(txQueries); err != nil { {{- if hasLog "query" }} diff --git a/generator/templates/model_create.gotpl b/generator/templates/model_create.gotpl index 1acd15f..beb2ca4 100644 --- a/generator/templates/model_create.gotpl +++ b/generator/templates/model_create.gotpl @@ -24,6 +24,12 @@ func (d *{{ .Model.Name }}Delegate) Create(input {{ .Model.Name }}CreateInput) * } func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, input {{ .Model.Name }}CreateInput, selects *{{ .Model.Name }}Select, omits *{{ .Model.Name }}Omit) (*{{ .Model.Name }}, error) { + if q.{{ .Model.Name }}.beforeCreate != nil { + if err := q.{{ .Model.Name }}.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -58,6 +64,12 @@ func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, input {{ . return nil, err } + if q.{{ .Model.Name }}.afterCreate != nil { + if err := q.{{ .Model.Name }}.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/generator/templates/model_structs.gotpl b/generator/templates/model_structs.gotpl index 34ff73c..61457b9 100644 --- a/generator/templates/model_structs.gotpl +++ b/generator/templates/model_structs.gotpl @@ -36,7 +36,17 @@ type {{ .Model.Name }}Omit struct { } type {{ .Model.Name }}Delegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *{{ .Model.Name }}CreateInput) error + afterCreate func(context.Context, *{{ .Model.Name }}) error +} + +func (d *{{ .Model.Name }}Delegate) BeforeCreate(hook func(context.Context, *{{ .Model.Name }}CreateInput) error) { + d.beforeCreate = hook +} + +func (d *{{ .Model.Name }}Delegate) AfterCreate(hook func(context.Context, *{{ .Model.Name }}) error) { + d.afterCreate = hook } func (m *{{ .Model.Name }}) ScanFields(cols []string) []any { diff --git a/generator/templates/tx.gotpl b/generator/templates/tx.gotpl index 7e5e9f1..5501a4a 100644 --- a/generator/templates/tx.gotpl +++ b/generator/templates/tx.gotpl @@ -20,9 +20,8 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { {{ $enum.Name }}: db.{{ $enum.Name }}, {{- end }} } - {{- range $model := .Schema.Models }} - q.{{ $model.Name }} = &{{ $model.Name }}Delegate{client: q} - {{- end }} + q.initDelegates() + q.copyHooksFrom(db.Queries) return &Tx{ Queries: q, tx: sqlTx, diff --git a/integration/schema.prisma b/integration/schema.prisma index 0eccc03..b596f4f 100644 --- a/integration/schema.prisma +++ b/integration/schema.prisma @@ -14,6 +14,7 @@ model User { id String @id @default(cuid()) email String @unique phoneNum String @unique + password String? role UserRole @default(STUDENT) profile Profile? posts Post[] diff --git a/integration/validation_test.go b/integration/validation_test.go index 3b35d26..a6bd9a6 100644 --- a/integration/validation_test.go +++ b/integration/validation_test.go @@ -2,6 +2,8 @@ package main import ( "context" + "crypto/sha256" + "encoding/hex" "fmt" "strings" "sync" @@ -447,3 +449,77 @@ func countAllUsers(t *testing.T, ctx context.Context, db *valkyrie.DB) int { } return count } + +func TestCreate_Hooks(t *testing.T) { + db, cleanup := setupTestDB(t) + defer cleanup() + ctx := context.Background() + + db.User.BeforeCreate(func(ctx context.Context, input *valkyrie.UserCreateInput) error { + if input.Email == "hook@example.com" { + input.PhoneNum = "+188888888" + } + return nil + }) + + var afterCalled bool + db.User.AfterCreate(func(ctx context.Context, u *valkyrie.User) error { + if u.Email == "hook@example.com" { + afterCalled = true + } + return nil + }) + + u, err := db.User.Create(valkyrie.UserCreateInput{ + Email: "hook@example.com", + PhoneNum: "+100000000", // Will be modified by hook + }).Exec(ctx) + + if err != nil { + t.Fatalf("failed to create user: %v", err) + } + + if u.PhoneNum != "+188888888" { + t.Errorf("expected PhoneNum mutated to '+188888888', got %q", u.PhoneNum) + } + + if !afterCalled { + t.Error("expected AfterCreate hook to be called") + } +} + +func TestCreate_Hooks_PasswordHashing(t *testing.T) { + db, cleanup := setupTestDB(t) + defer cleanup() + ctx := context.Background() + + db.User.BeforeCreate(func(ctx context.Context, input *valkyrie.UserCreateInput) error { + if input.Email == "hash@example.com" && input.Password != nil { + + h := sha256.Sum256([]byte(*input.Password)) + hashed := hex.EncodeToString(h[:]) + input.Password = &hashed + } + return nil + }) + + rawPassword := "12345678" + + u, err := db.User.Create(valkyrie.UserCreateInput{ + Email: "hash@example.com", + PhoneNum: "+199999999", + Password: &rawPassword, + }).Exec(ctx) + + if err != nil { + t.Fatalf("failed to create user: %v", err) + } + + if u.Password == nil { + t.Fatal("expected Password field to be populated, got nil") + } + + if *u.Password == rawPassword { + t.Errorf("expected Password to be hashed, got %q", *u.Password) + } +} diff --git a/integration/valkyrie/category.go b/integration/valkyrie/category.go index 6046608..7b176f3 100644 --- a/integration/valkyrie/category.go +++ b/integration/valkyrie/category.go @@ -46,7 +46,17 @@ type CategoryOmit struct { } type CategoryDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *CategoryCreateInput) error + afterCreate func(context.Context, *Category) error +} + +func (d *CategoryDelegate) BeforeCreate(hook func(context.Context, *CategoryCreateInput) error) { + d.beforeCreate = hook +} + +func (d *CategoryDelegate) AfterCreate(hook func(context.Context, *Category) error) { + d.afterCreate = hook } func (m *Category) ScanFields(cols []string) []any { @@ -129,6 +139,12 @@ func (d *CategoryDelegate) Create(input CategoryCreateInput) *CreateBuilder[Cate } func (q *Queries) executeCategoryCreate(ctx context.Context, input CategoryCreateInput, selects *CategorySelect, omits *CategoryOmit) (*Category, error) { + if q.Category.beforeCreate != nil { + if err := q.Category.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -163,6 +179,12 @@ func (q *Queries) executeCategoryCreate(ctx context.Context, input CategoryCreat return nil, err } + if q.Category.afterCreate != nil { + if err := q.Category.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/integration/valkyrie/categoryToPost.go b/integration/valkyrie/categoryToPost.go index 43d00db..e5a850a 100644 --- a/integration/valkyrie/categoryToPost.go +++ b/integration/valkyrie/categoryToPost.go @@ -49,7 +49,17 @@ type CategoryToPostOmit struct { } type CategoryToPostDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *CategoryToPostCreateInput) error + afterCreate func(context.Context, *CategoryToPost) error +} + +func (d *CategoryToPostDelegate) BeforeCreate(hook func(context.Context, *CategoryToPostCreateInput) error) { + d.beforeCreate = hook +} + +func (d *CategoryToPostDelegate) AfterCreate(hook func(context.Context, *CategoryToPost) error) { + d.afterCreate = hook } func (m *CategoryToPost) ScanFields(cols []string) []any { @@ -132,6 +142,12 @@ func (d *CategoryToPostDelegate) Create(input CategoryToPostCreateInput) *Create } func (q *Queries) executeCategoryToPostCreate(ctx context.Context, input CategoryToPostCreateInput, selects *CategoryToPostSelect, omits *CategoryToPostOmit) (*CategoryToPost, error) { + if q.CategoryToPost.beforeCreate != nil { + if err := q.CategoryToPost.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -166,6 +182,12 @@ func (q *Queries) executeCategoryToPostCreate(ctx context.Context, input Categor return nil, err } + if q.CategoryToPost.afterCreate != nil { + if err := q.CategoryToPost.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/integration/valkyrie/client.go b/integration/valkyrie/client.go index f050ef4..f38499e 100644 --- a/integration/valkyrie/client.go +++ b/integration/valkyrie/client.go @@ -161,16 +161,35 @@ func Open(provider, dataSourceName string) (*DB, error) { dialect: sqliteDialect{}, UserRole: UserRole, } + q.initDelegates() + return &DB{ + Queries: q, + sqlDB: sqlDB, + }, nil +} + +func (q *Queries) initDelegates() { q.User = &UserDelegate{client: q} q.Profile = &ProfileDelegate{client: q} q.Post = &PostDelegate{client: q} q.Comment = &CommentDelegate{client: q} q.Category = &CategoryDelegate{client: q} q.CategoryToPost = &CategoryToPostDelegate{client: q} - return &DB{ - Queries: q, - sqlDB: sqlDB, - }, nil +} + +func (q *Queries) copyHooksFrom(other *Queries) { + q.User.beforeCreate = other.User.beforeCreate + q.User.afterCreate = other.User.afterCreate + q.Profile.beforeCreate = other.Profile.beforeCreate + q.Profile.afterCreate = other.Profile.afterCreate + q.Post.beforeCreate = other.Post.beforeCreate + q.Post.afterCreate = other.Post.afterCreate + q.Comment.beforeCreate = other.Comment.beforeCreate + q.Comment.afterCreate = other.Comment.afterCreate + q.Category.beforeCreate = other.Category.beforeCreate + q.Category.afterCreate = other.Category.afterCreate + q.CategoryToPost.beforeCreate = other.CategoryToPost.beforeCreate + q.CategoryToPost.afterCreate = other.CategoryToPost.afterCreate } // Close closes the database connection. @@ -255,12 +274,8 @@ func (q *Queries) transaction(ctx context.Context, fn func(txQ *Queries) error) dialect: q.dialect, UserRole: q.UserRole, } - txQueries.User = &UserDelegate{client: txQueries} - txQueries.Profile = &ProfileDelegate{client: txQueries} - txQueries.Post = &PostDelegate{client: txQueries} - txQueries.Comment = &CommentDelegate{client: txQueries} - txQueries.Category = &CategoryDelegate{client: txQueries} - txQueries.CategoryToPost = &CategoryToPostDelegate{client: txQueries} + txQueries.initDelegates() + txQueries.copyHooksFrom(q) if err := fn(txQueries); err != nil { _ = tx.Rollback() @@ -286,12 +301,8 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { dialect: db.dialect, UserRole: db.UserRole, } - q.User = &UserDelegate{client: q} - q.Profile = &ProfileDelegate{client: q} - q.Post = &PostDelegate{client: q} - q.Comment = &CommentDelegate{client: q} - q.Category = &CategoryDelegate{client: q} - q.CategoryToPost = &CategoryToPostDelegate{client: q} + q.initDelegates() + q.copyHooksFrom(db.Queries) return &Tx{ Queries: q, tx: sqlTx, diff --git a/integration/valkyrie/comment.go b/integration/valkyrie/comment.go index 00e3133..eff4b44 100644 --- a/integration/valkyrie/comment.go +++ b/integration/valkyrie/comment.go @@ -69,7 +69,17 @@ type CommentOmit struct { } type CommentDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *CommentCreateInput) error + afterCreate func(context.Context, *Comment) error +} + +func (d *CommentDelegate) BeforeCreate(hook func(context.Context, *CommentCreateInput) error) { + d.beforeCreate = hook +} + +func (d *CommentDelegate) AfterCreate(hook func(context.Context, *Comment) error) { + d.afterCreate = hook } func (m *Comment) ScanFields(cols []string) []any { @@ -213,6 +223,12 @@ func (d *CommentDelegate) Create(input CommentCreateInput) *CreateBuilder[Commen } func (q *Queries) executeCommentCreate(ctx context.Context, input CommentCreateInput, selects *CommentSelect, omits *CommentOmit) (*Comment, error) { + if q.Comment.beforeCreate != nil { + if err := q.Comment.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -247,6 +263,12 @@ func (q *Queries) executeCommentCreate(ctx context.Context, input CommentCreateI return nil, err } + if q.Comment.afterCreate != nil { + if err := q.Comment.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/integration/valkyrie/migrations/00002_add_password.sql b/integration/valkyrie/migrations/00002_add_password.sql new file mode 100644 index 0000000..8a1fc3e --- /dev/null +++ b/integration/valkyrie/migrations/00002_add_password.sql @@ -0,0 +1,22 @@ +-- +goose Up +ALTER TABLE `User` ADD COLUMN `password` text NULL; + +-- +goose Down +PRAGMA foreign_keys = off; +CREATE TABLE `new_User` ( + `id` text NOT NULL, + `email` text NOT NULL, + `phoneNum` text NOT NULL, + `role` text NOT NULL DEFAULT 'student', + `referredById` text NULL, + PRIMARY KEY (`id`), + CONSTRAINT `User_referredById_fkey` FOREIGN KEY (`referredById`) REFERENCES `User` (`id`) ON UPDATE NO ACTION ON DELETE NO ACTION, + CONSTRAINT `User_role_check` CHECK ("role" IN ('ADMIN', 'student', 'TEACHER')) +); +INSERT INTO `new_User` (`id`, `email`, `phoneNum`, `role`, `referredById`) SELECT `id`, `email`, `phoneNum`, `role`, `referredById` FROM `User`; +DROP TABLE `User`; +ALTER TABLE `new_User` RENAME TO `User`; +CREATE UNIQUE INDEX `User_email_key` ON `User` (`email`); +CREATE UNIQUE INDEX `User_phoneNum_key` ON `User` (`phoneNum`); +CREATE UNIQUE INDEX `User_email_phoneNum_key` ON `User` (`email`, `phoneNum`); +PRAGMA foreign_keys = on; diff --git a/integration/valkyrie/post.go b/integration/valkyrie/post.go index e6a47c5..62d1e7b 100644 --- a/integration/valkyrie/post.go +++ b/integration/valkyrie/post.go @@ -64,7 +64,17 @@ type PostOmit struct { } type PostDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *PostCreateInput) error + afterCreate func(context.Context, *Post) error +} + +func (d *PostDelegate) BeforeCreate(hook func(context.Context, *PostCreateInput) error) { + d.beforeCreate = hook +} + +func (d *PostDelegate) AfterCreate(hook func(context.Context, *Post) error) { + d.afterCreate = hook } func (m *Post) ScanFields(cols []string) []any { @@ -180,6 +190,12 @@ func (d *PostDelegate) Create(input PostCreateInput) *CreateBuilder[Post, PostCr } func (q *Queries) executePostCreate(ctx context.Context, input PostCreateInput, selects *PostSelect, omits *PostOmit) (*Post, error) { + if q.Post.beforeCreate != nil { + if err := q.Post.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -214,6 +230,12 @@ func (q *Queries) executePostCreate(ctx context.Context, input PostCreateInput, return nil, err } + if q.Post.afterCreate != nil { + if err := q.Post.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/integration/valkyrie/profile.go b/integration/valkyrie/profile.go index b6386c2..4d1d3e5 100644 --- a/integration/valkyrie/profile.go +++ b/integration/valkyrie/profile.go @@ -50,7 +50,17 @@ type ProfileOmit struct { } type ProfileDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *ProfileCreateInput) error + afterCreate func(context.Context, *Profile) error +} + +func (d *ProfileDelegate) BeforeCreate(hook func(context.Context, *ProfileCreateInput) error) { + d.beforeCreate = hook +} + +func (d *ProfileDelegate) AfterCreate(hook func(context.Context, *Profile) error) { + d.afterCreate = hook } func (m *Profile) ScanFields(cols []string) []any { @@ -147,6 +157,12 @@ func (d *ProfileDelegate) Create(input ProfileCreateInput) *CreateBuilder[Profil } func (q *Queries) executeProfileCreate(ctx context.Context, input ProfileCreateInput, selects *ProfileSelect, omits *ProfileOmit) (*Profile, error) { + if q.Profile.beforeCreate != nil { + if err := q.Profile.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -181,6 +197,12 @@ func (q *Queries) executeProfileCreate(ctx context.Context, input ProfileCreateI return nil, err } + if q.Profile.afterCreate != nil { + if err := q.Profile.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } diff --git a/integration/valkyrie/user.go b/integration/valkyrie/user.go index a7cf4ed..6a89171 100644 --- a/integration/valkyrie/user.go +++ b/integration/valkyrie/user.go @@ -23,6 +23,7 @@ type User struct { Id string `db:"id" json:"id"` Email string `db:"email" json:"email"` PhoneNum string `db:"phoneNum" json:"phoneNum"` + Password *string `db:"password" json:"password"` Role UserRoleType `db:"role" json:"role"` ReferredById *string `db:"referredById" json:"referredById"` Profile *Profile `json:"profile,omitempty"` @@ -37,6 +38,7 @@ type UserCreateInput struct { Id *string `json:"id"` Email string `json:"email"` PhoneNum string `json:"phoneNum"` + Password *string `json:"password"` Role *UserRoleType `json:"role"` ReferredById *string `json:"referredById"` } @@ -46,6 +48,7 @@ type UserSelect struct { Id bool `json:"id"` Email bool `json:"email"` PhoneNum bool `json:"phoneNum"` + Password bool `json:"password"` Role bool `json:"role"` ReferredById bool `json:"referredById"` Profile *ProfileSelect `json:"profile,omitempty"` @@ -60,6 +63,7 @@ type UserOmit struct { Id bool `json:"id"` Email bool `json:"email"` PhoneNum bool `json:"phoneNum"` + Password bool `json:"password"` Role bool `json:"role"` ReferredById bool `json:"referredById"` Profile *ProfileOmit `json:"profile,omitempty"` @@ -70,7 +74,17 @@ type UserOmit struct { } type UserDelegate struct { - client *Queries + client *Queries + beforeCreate func(context.Context, *UserCreateInput) error + afterCreate func(context.Context, *User) error +} + +func (d *UserDelegate) BeforeCreate(hook func(context.Context, *UserCreateInput) error) { + d.beforeCreate = hook +} + +func (d *UserDelegate) AfterCreate(hook func(context.Context, *User) error) { + d.afterCreate = hook } func (m *User) ScanFields(cols []string) []any { @@ -83,6 +97,8 @@ func (m *User) ScanFields(cols []string) []any { targets[i] = &m.Email case "phoneNum": targets[i] = &m.PhoneNum + case "password": + targets[i] = &m.Password case "role": targets[i] = &m.Role case "referredById": @@ -96,6 +112,7 @@ var userDefaultCols = []string{ "id", "email", "phoneNum", + "password", "role", "referredById", } @@ -105,12 +122,13 @@ func (q *Queries) selectUserCols(selects *UserSelect, omits *UserOmit, forceCols return userDefaultCols } - anySelected := selects != nil && (selects.Id || selects.Email || selects.PhoneNum || selects.Role || selects.ReferredById || selects.Profile != nil || selects.Posts != nil || selects.Comments != nil || selects.ReferredBy != nil || selects.Referrals != nil) + anySelected := selects != nil && (selects.Id || selects.Email || selects.PhoneNum || selects.Password || selects.Role || selects.ReferredById || selects.Profile != nil || selects.Posts != nil || selects.Comments != nil || selects.ReferredBy != nil || selects.Referrals != nil) specs := []colSpec{ {"id", selects != nil && selects.Id, omits != nil && omits.Id, false}, {"email", selects != nil && selects.Email, omits != nil && omits.Email, false}, {"phoneNum", selects != nil && selects.PhoneNum, omits != nil && omits.PhoneNum, false}, + {"password", selects != nil && selects.Password, omits != nil && omits.Password, false}, {"role", selects != nil && selects.Role, omits != nil && omits.Role, false}, {"referredById", selects != nil && selects.ReferredById, omits != nil && omits.ReferredById, selects != nil && selects.ReferredBy != nil}, } @@ -171,6 +189,7 @@ var UserColOrder = []string{ "id", "email", "phoneNum", + "password", "role", "referredById", } @@ -191,6 +210,12 @@ func (d *UserDelegate) Create(input UserCreateInput) *CreateBuilder[User, UserCr } func (q *Queries) executeUserCreate(ctx context.Context, input UserCreateInput, selects *UserSelect, omits *UserOmit) (*User, error) { + if q.User.beforeCreate != nil { + if err := q.User.beforeCreate(ctx, &input); err != nil { + return nil, err + } + } + if err := input.Validate(); err != nil { return nil, err } @@ -225,6 +250,12 @@ func (q *Queries) executeUserCreate(ctx context.Context, input UserCreateInput, return nil, err } + if q.User.afterCreate != nil { + if err := q.User.afterCreate(ctx, res); err != nil { + return nil, err + } + } + return res, nil } @@ -237,6 +268,9 @@ func (q *Queries) UserInputToMap(input UserCreateInput) map[string]any { } m["email"] = input.Email m["phoneNum"] = input.PhoneNum + if input.Password != nil { + m["password"] = *input.Password + } if input.Role != nil { m["role"] = *input.Role }