Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 16 additions & 6 deletions generator/templates/client.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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" }}
Expand Down
12 changes: 12 additions & 0 deletions generator/templates/model_create.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}

Expand Down
12 changes: 11 additions & 1 deletion generator/templates/model_structs.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
5 changes: 2 additions & 3 deletions generator/templates/tx.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions integration/schema.prisma
Original file line number Diff line number Diff line change
Expand Up @@ -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[]
Expand Down
76 changes: 76 additions & 0 deletions integration/validation_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ package main

import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"strings"
"sync"
Expand Down Expand Up @@ -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)
}
}
24 changes: 23 additions & 1 deletion integration/valkyrie/category.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}

Expand Down
24 changes: 23 additions & 1 deletion integration/valkyrie/categoryToPost.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}

Expand Down
43 changes: 27 additions & 16 deletions integration/valkyrie/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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()
Expand All @@ -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,
Expand Down
24 changes: 23 additions & 1 deletion integration/valkyrie/comment.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}

Expand Down
Loading
Loading