From de836c05c6655ccb4d78d275efa317e1a5bf0f2d Mon Sep 17 00:00:00 2001 From: Clancy Date: Sat, 4 Jul 2026 16:39:00 +0300 Subject: [PATCH 1/2] Feat: Implement query and migration logging in client (configurable) 1- add log levels and runtime validation for logs in getconfig() 2- pass a hasLog() func to templates to embed what's needed instead of bloating the client 3- refactor raw query calls on the query struct, and enable logging there instead of writing dialect specific logic --- cli/getConfig.go | 25 ++++++ cli/handleGenerate.go | 2 +- generator/generator.go | 12 ++- generator/generator_test.go | 97 ----------------------- generator/templates/builders_create.gotpl | 8 +- generator/templates/client.gotpl | 63 ++++++++++++++- generator/templates/header.gotpl | 1 + integration/valkyrie.json | 4 +- integration/valkyrie/client.go | 41 ++++++++-- 9 files changed, 142 insertions(+), 111 deletions(-) delete mode 100644 generator/generator_test.go diff --git a/cli/getConfig.go b/cli/getConfig.go index 9a82e59..0a7de80 100644 --- a/cli/getConfig.go +++ b/cli/getConfig.go @@ -4,12 +4,23 @@ import ( "encoding/json" "log" "os" + "slices" ) +var LogLevels = []string{ + "query", + "info", + "warn", + "error", + "all", + "none", +} + type Config struct { Database DatabaseConfig `json:"database"` Schema string `json:"schema"` Output OutputConfig `json:"output"` + Log []string `json:"log"` } type DatabaseConfig struct { @@ -35,6 +46,20 @@ func GetConfig() *Config { log.Fatal(err) return nil } + // hasAll := false + for _, l := range config.Log { + if l == "all" { + // hasAll = true + } + if !slices.Contains(LogLevels, l) && l != "all" { + log.Fatalf("invalid log level in valkyrie.json: %q (must be one of: query, info, warn, error, all)", l) + return nil + } + } + // if hasAll && len(config.Log) > 1 { + // log.Fatal("invalid log configuration: 'all' must be the only log level specified") + // return nil + // } return &config } diff --git a/cli/handleGenerate.go b/cli/handleGenerate.go index 4c03946..872a3b7 100644 --- a/cli/handleGenerate.go +++ b/cli/handleGenerate.go @@ -47,7 +47,7 @@ func handleGenerate() { pkgName = "valkyrie" } - outputs, err := generator.GenerateClient(*schemaDef, pkgName, embedRelDir, config.Output.Migrations) + outputs, err := generator.GenerateClient(*schemaDef, pkgName, embedRelDir, config.Output.Migrations, config.Log) if err != nil { fmt.Printf("failed to generate client: %v\n", err) return diff --git a/generator/generator.go b/generator/generator.go index dc53927..50f7857 100644 --- a/generator/generator.go +++ b/generator/generator.go @@ -18,6 +18,7 @@ type templateData struct { EmbedDir string DefaultDiskPath string Schema schema.Schema + DefaultLogs []string } type modelTemplateData struct { @@ -25,11 +26,19 @@ type modelTemplateData struct { Model *schema.Model } -func GenerateClient(sch schema.Schema, pkgName string, embedPath string, defaultDiskPath string) (map[string]string, error) { +func GenerateClient(sch schema.Schema, pkgName string, embedPath string, defaultDiskPath string, defaultLogs []string) (map[string]string, error) { tmpl := template.New("").Funcs(template.FuncMap{ "capitalize": capitalize, "lowercase": lowercase, "fkForRelation": fkForRelation, + "hasLog": func(level string) bool { + for _, l := range defaultLogs { + if l == "all" || l == level { + return true + } + } + return false + }, }) tmpl, err := tmpl.ParseFS(templatesFS, "templates/*.gotpl") if err != nil { @@ -47,6 +56,7 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default EmbedDir: embedDir, DefaultDiskPath: defaultDiskPath, Schema: sch, + DefaultLogs: defaultLogs, } outputs := make(map[string]string) diff --git a/generator/generator_test.go b/generator/generator_test.go deleted file mode 100644 index 372d51e..0000000 --- a/generator/generator_test.go +++ /dev/null @@ -1,97 +0,0 @@ -package generator - -import ( - "strings" - "testing" - "valkyrie/schema" -) - -func TestGenerateClient(t *testing.T) { - input := ` - datasource db { - provider = "postgresql" - } - - enum Role { - USER - ADMIN - } - - model User { - id Int @id @default(autoincrement()) - email String @unique - role Role @default(USER) - - @@map("users") - } - - model Post { - id String @id @default(uuid()) - title String - published Boolean @default(false) - authorId Int - author User @relation(fields: [authorId], references: [id]) - } - ` - - sch, errs := schema.ParseSchema(input) - if len(errs) > 0 { - t.Fatalf("parser errors: %v", errs) - } - - outputs, err := GenerateClient(*sch, "client", "migrations/*.sql", "migrations") - if err != nil { - t.Fatalf("failed to generate client: %v\n", err) - } - - code := "" - for _, content := range outputs { - code += content + "\n" - } - - if !strings.Contains(code, "type CreateBuilder[M any, I any, S any, O any] struct {") { - t.Errorf("expected CreateBuilder struct in code, got:\n%s", code) - } - - if !strings.Contains(code, "package client") { - t.Errorf("expected package client, got:\n%s", code) - } - - if !strings.Contains(code, "type DB struct {") { - t.Errorf("expected DB struct, got:\n%s", code) - } - if !strings.Contains(code, "User") || !strings.Contains(code, "*UserDelegate") { - t.Errorf("expected User delegate field on DB, got:\n%s", code) - } - if !strings.Contains(code, "Post") || !strings.Contains(code, "*PostDelegate") { - t.Errorf("expected Post delegate field on DB, got:\n%s", code) - } - if !strings.Contains(code, "type UserDelegate struct {") { - t.Errorf("expected UserDelegate struct, got:\n%s", code) - } - if !strings.Contains(code, "type PostDelegate struct {") { - t.Errorf("expected PostDelegate struct, got:\n%s", code) - } - - if !strings.Contains(code, "type RoleType string") { - t.Errorf("expected RoleType type definition, got:\n%s", code) - } - if !strings.Contains(code, "RoleTypeUser") || !strings.Contains(code, "\"USER\"") { - t.Errorf("expected RoleTypeUser constant, got:\n%s", code) - } - if !strings.Contains(code, "RoleTypeAdmin") || !strings.Contains(code, "\"ADMIN\"") { - t.Errorf("expected RoleTypeAdmin constant, got:\n%s", code) - } - if !strings.Contains(code, "type roleNamespace struct {") { - t.Errorf("expected roleNamespace struct definition, got:\n%s", code) - } - if !strings.Contains(code, "var Role = roleNamespace{") { - t.Errorf("expected Role namespace variable declaration, got:\n%s", code) - } - if !strings.Contains(code, "User:") || !strings.Contains(code, "RoleTypeUser") { - t.Errorf("expected User: RoleTypeUser mapping in Role namespace, got:\n%s", code) - } - if !strings.Contains(code, "Admin:") || !strings.Contains(code, "RoleTypeAdmin") { - t.Errorf("expected Admin: RoleTypeAdmin mapping in Role namespace, got:\n%s", code) - } -} diff --git a/generator/templates/builders_create.gotpl b/generator/templates/builders_create.gotpl index efee1b1..412c2d1 100644 --- a/generator/templates/builders_create.gotpl +++ b/generator/templates/builders_create.gotpl @@ -77,7 +77,7 @@ func executeInsert[M any]( var res M if q.dialect.SupportsReturning() { - row := q.db.QueryRowContext(ctx, query, vals...) + row := q.queryRow(ctx, query, vals...) scanTargets := scanFunc(&res, returningCols) if err := row.Scan(scanTargets...); err != nil { @@ -87,7 +87,7 @@ func executeInsert[M any]( } // Fallback for dialects without RETURNING (MySQL) - result, err := q.db.ExecContext(ctx, query, vals...) + result, err := q.exec(ctx, query, vals...) if err != nil { return nil, err } @@ -123,7 +123,7 @@ func executeInsert[M any]( selectSb.WriteString(q.dialect.Quote(idCol)) selectSb.WriteString(" = ?") - row := q.db.QueryRowContext(ctx, selectSb.String(), idVal) + row := q.queryRow(ctx, selectSb.String(), idVal) scanTargets := scanFunc(&res, returningCols) if err := row.Scan(scanTargets...); err != nil { return nil, err @@ -175,7 +175,7 @@ func loadRelation[P any, C any]( sb.WriteString(")") query := sb.String() - rows, err := q.db.QueryContext(ctx, query, parentKeys...) + rows, err := q.query(ctx, query, parentKeys...) if err != nil { return nil, err } diff --git a/generator/templates/client.gotpl b/generator/templates/client.gotpl index bfd3173..677e438 100644 --- a/generator/templates/client.gotpl +++ b/generator/templates/client.gotpl @@ -83,21 +83,47 @@ func (db *DB) Raw() *sql.DB { {{- if .EmbedPath }} // RunMigrations runs all pending migrations from the embedded folder. func (db *DB) RunMigrations(ctx context.Context) error { + {{- if hasLog "info" }} + log.Println("Running migrations...") + {{- end }} if err := goose.SetDialect(db.provider); err != nil { return err } goose.SetLogger(goose.NopLogger()) goose.SetBaseFS(migrationsFS) - return goose.UpContext(ctx, db.sqlDB, "{{ .EmbedDir }}") + err := goose.UpContext(ctx, db.sqlDB, "{{ .EmbedDir }}") + if err != nil { + {{- if hasLog "error" }} + log.Printf("Migrations failed: %v", err) + {{- end }} + return err + } + {{- if hasLog "info" }} + log.Println("Migrations completed successfully.") + {{- end }} + return nil } {{- else }} // RunMigrations runs all pending migrations from the disk folder. func (db *DB) RunMigrations(ctx context.Context) error { + {{- if hasLog "info" }} + log.Println("Running migrations...") + {{- end }} if err := goose.SetDialect(db.provider); err != nil { return err } goose.SetLogger(goose.NopLogger()) - return goose.UpContext(ctx, db.sqlDB, "{{ .DefaultDiskPath }}") + err := goose.UpContext(ctx, db.sqlDB, "{{ .DefaultDiskPath }}") + if err != nil { + {{- if hasLog "error" }} + log.Printf("Migrations failed: %v", err) + {{- end }} + return err + } + {{- if hasLog "info" }} + log.Println("Migrations completed successfully.") + {{- end }} + return nil } {{- end }} @@ -115,3 +141,36 @@ func (q *Queries) bindVars(count int) string { } return sb.String() } + +func (q *Queries) query(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Query: %s | Args: %v", strings.ToUpper(q.provider), query, args) + {{- end }} + res, err := q.db.QueryContext(ctx, query, args...) + {{- if hasLog "error" }} + if err != nil { + log.Printf("[%s] SQL Error: %v | Query: %s | Args: %v", strings.ToUpper(q.provider), err, query, args) + } + {{- end }} + return res, err +} + +func (q *Queries) queryRow(ctx context.Context, query string, args ...any) *sql.Row { + {{- if hasLog "query" }} + log.Printf("[%s] SQL QueryRow: %s | Args: %v", strings.ToUpper(q.provider), query, args) + {{- end }} + return q.db.QueryRowContext(ctx, query, args...) +} + +func (q *Queries) exec(ctx context.Context, query string, args ...any) (sql.Result, error) { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Exec: %s | Args: %v", strings.ToUpper(q.provider), query, args) + {{- end }} + res, err := q.db.ExecContext(ctx, query, args...) + {{- if hasLog "error" }} + if err != nil { + log.Printf("[%s] SQL Error: %v | Query: %s | Args: %v", strings.ToUpper(q.provider), err, query, args) + } + {{- end }} + return res, err +} diff --git a/generator/templates/header.gotpl b/generator/templates/header.gotpl index b167d3e..372b531 100644 --- a/generator/templates/header.gotpl +++ b/generator/templates/header.gotpl @@ -9,6 +9,7 @@ import ( "embed" {{- end }} "fmt" + "log" "strconv" "strings" "time" diff --git a/integration/valkyrie.json b/integration/valkyrie.json index 4fa9a00..72edeb4 100644 --- a/integration/valkyrie.json +++ b/integration/valkyrie.json @@ -9,5 +9,7 @@ "output": { "client": "./valkyrie", "migrations": "./valkyrie/migrations" - } + }, + + "log": ["query", "info", "error"] } diff --git a/integration/valkyrie/client.go b/integration/valkyrie/client.go index 7867c8e..807787c 100644 --- a/integration/valkyrie/client.go +++ b/integration/valkyrie/client.go @@ -7,6 +7,7 @@ import ( "embed" "encoding/json" "fmt" + "log" "strconv" "strings" "time" @@ -137,12 +138,19 @@ func (db *DB) Raw() *sql.DB { // RunMigrations runs all pending migrations from the embedded folder. func (db *DB) RunMigrations(ctx context.Context) error { + log.Println("Running migrations...") if err := goose.SetDialect(db.provider); err != nil { return err } goose.SetLogger(goose.NopLogger()) goose.SetBaseFS(migrationsFS) - return goose.UpContext(ctx, db.sqlDB, "migrations") + err := goose.UpContext(ctx, db.sqlDB, "migrations") + if err != nil { + log.Printf("Migrations failed: %v", err) + return err + } + log.Println("Migrations completed successfully.") + return nil } func (q *Queries) bindVars(count int) string { @@ -160,6 +168,29 @@ func (q *Queries) bindVars(count int) string { return sb.String() } +func (q *Queries) query(ctx context.Context, query string, args ...any) (*sql.Rows, error) { + log.Printf("[%s] SQL Query: %s | Args: %v", strings.ToUpper(q.provider), query, args) + res, err := q.db.QueryContext(ctx, query, args...) + if err != nil { + log.Printf("[%s] SQL Error: %v | Query: %s | Args: %v", strings.ToUpper(q.provider), err, query, args) + } + return res, err +} + +func (q *Queries) queryRow(ctx context.Context, query string, args ...any) *sql.Row { + log.Printf("[%s] SQL QueryRow: %s | Args: %v", strings.ToUpper(q.provider), query, args) + return q.db.QueryRowContext(ctx, query, args...) +} + +func (q *Queries) exec(ctx context.Context, query string, args ...any) (sql.Result, error) { + log.Printf("[%s] SQL Exec: %s | Args: %v", strings.ToUpper(q.provider), query, args) + res, err := q.db.ExecContext(ctx, query, args...) + if err != nil { + log.Printf("[%s] SQL Error: %v | Query: %s | Args: %v", strings.ToUpper(q.provider), err, query, args) + } + return res, err +} + type Tx struct { *Queries tx *sql.Tx @@ -309,7 +340,7 @@ func executeInsert[M any]( var res M if q.dialect.SupportsReturning() { - row := q.db.QueryRowContext(ctx, query, vals...) + row := q.queryRow(ctx, query, vals...) scanTargets := scanFunc(&res, returningCols) if err := row.Scan(scanTargets...); err != nil { @@ -319,7 +350,7 @@ func executeInsert[M any]( } // Fallback for dialects without RETURNING (MySQL) - result, err := q.db.ExecContext(ctx, query, vals...) + result, err := q.exec(ctx, query, vals...) if err != nil { return nil, err } @@ -355,7 +386,7 @@ func executeInsert[M any]( selectSb.WriteString(q.dialect.Quote(idCol)) selectSb.WriteString(" = ?") - row := q.db.QueryRowContext(ctx, selectSb.String(), idVal) + row := q.queryRow(ctx, selectSb.String(), idVal) scanTargets := scanFunc(&res, returningCols) if err := row.Scan(scanTargets...); err != nil { return nil, err @@ -406,7 +437,7 @@ func loadRelation[P any, C any]( sb.WriteString(")") query := sb.String() - rows, err := q.db.QueryContext(ctx, query, parentKeys...) + rows, err := q.query(ctx, query, parentKeys...) if err != nil { return nil, err } From cb886e57381ea59a09d0fce3e34e03f100340566 Mon Sep 17 00:00:00 2001 From: Clancy Date: Sat, 4 Jul 2026 21:08:32 +0300 Subject: [PATCH 2/2] implement transactional creation for nested relations selects only, model scalar fields or no select at all is non-transactional --- generator/templates/client.gotpl | 56 ++++++++++++++++++++++++++ generator/templates/model_create.gotpl | 25 +++++++++--- generator/templates/tx.gotpl | 9 +++++ integration/valkyrie/category.go | 22 +++++++--- integration/valkyrie/categoryToPost.go | 22 +++++++--- integration/valkyrie/client.go | 50 +++++++++++++++++++++++ integration/valkyrie/comment.go | 22 +++++++--- integration/valkyrie/post.go | 22 +++++++--- integration/valkyrie/profile.go | 22 +++++++--- integration/valkyrie/user.go | 20 ++++++--- 10 files changed, 230 insertions(+), 40 deletions(-) diff --git a/generator/templates/client.gotpl b/generator/templates/client.gotpl index 677e438..75503b4 100644 --- a/generator/templates/client.gotpl +++ b/generator/templates/client.gotpl @@ -174,3 +174,59 @@ func (q *Queries) exec(ctx context.Context, query string, args ...any) (sql.Resu {{- end }} return res, err } + +func (q *Queries) transaction(ctx context.Context, fn func(txQ *Queries) error) error { + if _, ok := q.db.(*sql.Tx); ok { + return fn(q) + } + + starter, ok := q.db.(interface { + BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) + }) + if !ok { + return fn(q) + } + + {{- if hasLog "query" }} + log.Printf("[%s] SQL Begin Transaction", strings.ToUpper(q.provider)) + {{- end }} + tx, err := starter.BeginTx(ctx, nil) + if err != nil { + return err + } + + defer func() { + if p := recover(); p != nil { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(q.provider)) + {{- end }} + _ = tx.Rollback() + panic(p) + } + }() + + txQueries := &Queries{ + db: tx, + provider: q.provider, + dialect: q.dialect, + {{- range $enum := .Schema.Enums }} + {{ $enum.Name }}: q.{{ $enum.Name }}, + {{- end }} + } + {{- range $model := .Schema.Models }} + txQueries.{{ $model.Name }} = &{{ $model.Name }}Delegate{client: txQueries} + {{- end }} + + if err := fn(txQueries); err != nil { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(q.provider)) + {{- end }} + _ = tx.Rollback() + return err + } + + {{- if hasLog "query" }} + log.Printf("[%s] SQL Commit Transaction", strings.ToUpper(q.provider)) + {{- end }} + return tx.Commit() +} diff --git a/generator/templates/model_create.gotpl b/generator/templates/model_create.gotpl index ae67f76..82ada59 100644 --- a/generator/templates/model_create.gotpl +++ b/generator/templates/model_create.gotpl @@ -74,12 +74,27 @@ func (q *Queries) execute{{ .Model.Name }}Create(ctx context.Context, input {{ . idCol := "{{ range $field := .Model.ScalarFields }}{{ if $field.IsID }}{{ $field.EffectiveColName }}{{ end }}{{ end }}" - res, err := executeInsert(ctx, q, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err - } + {{- if .Model.RelationFields }} + hasRelations := selects != nil && ({{ range $i, $rel := .Model.RelationFields }}{{ if $i }} || {{ end }}selects.{{ capitalize $rel.Name }} != nil{{ end }}) + {{- else }} + hasRelations := false + {{- end }} - if err := q.load{{ .Model.Name }}Relations(ctx, []*{{ .Model.Name }}{res}, selects); err != nil { + var res *{{ .Model.Name }} + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.load{{ .Model.Name }}Relations(ctx, []*{{ .Model.Name }}{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "{{ .Model.EffectiveTableName }}", cols, vals, returningCols, idCol, scanFunc) + } + if err != nil { return nil, err } diff --git a/generator/templates/tx.gotpl b/generator/templates/tx.gotpl index 8d828e3..7e5e9f1 100644 --- a/generator/templates/tx.gotpl +++ b/generator/templates/tx.gotpl @@ -9,6 +9,9 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { if err != nil { return nil, err } + {{- if hasLog "query" }} + log.Printf("[%s] SQL Begin Transaction", strings.ToUpper(db.provider)) + {{- end }} q := &Queries{ db: sqlTx, provider: db.provider, @@ -28,11 +31,17 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { // Commit commits the transaction. func (tx *Tx) Commit() error { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Commit Transaction", strings.ToUpper(tx.provider)) + {{- end }} return tx.tx.Commit() } // Rollback aborts the transaction. func (tx *Tx) Rollback() error { + {{- if hasLog "query" }} + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(tx.provider)) + {{- end }} return tx.tx.Rollback() } diff --git a/integration/valkyrie/category.go b/integration/valkyrie/category.go index 3e57a67..12ea2b2 100644 --- a/integration/valkyrie/category.go +++ b/integration/valkyrie/category.go @@ -114,13 +114,23 @@ func (q *Queries) executeCategoryCreate(ctx context.Context, input CategoryCreat } idCol := "id" - - res, err := executeInsert(ctx, q, "Category", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + hasRelations := selects != nil && (selects.Posts != nil) + + var res *Category + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "Category", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadCategoryRelations(ctx, []*Category{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "Category", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadCategoryRelations(ctx, []*Category{res}, selects); err != nil { + if err != nil { return nil, err } diff --git a/integration/valkyrie/categoryToPost.go b/integration/valkyrie/categoryToPost.go index 106dd43..77790dd 100644 --- a/integration/valkyrie/categoryToPost.go +++ b/integration/valkyrie/categoryToPost.go @@ -113,13 +113,23 @@ func (q *Queries) executeCategoryToPostCreate(ctx context.Context, input Categor } idCol := "" - - res, err := executeInsert(ctx, q, "CategoryToPost", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + hasRelations := selects != nil && (selects.Post != nil || selects.Category != nil) + + var res *CategoryToPost + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "CategoryToPost", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadCategoryToPostRelations(ctx, []*CategoryToPost{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "CategoryToPost", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadCategoryToPostRelations(ctx, []*CategoryToPost{res}, selects); err != nil { + if err != nil { return nil, err } diff --git a/integration/valkyrie/client.go b/integration/valkyrie/client.go index 666ddcc..20682a3 100644 --- a/integration/valkyrie/client.go +++ b/integration/valkyrie/client.go @@ -191,6 +191,53 @@ func (q *Queries) exec(ctx context.Context, query string, args ...any) (sql.Resu return res, err } +func (q *Queries) transaction(ctx context.Context, fn func(txQ *Queries) error) error { + if _, ok := q.db.(*sql.Tx); ok { + return fn(q) + } + + starter, ok := q.db.(interface { + BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) + }) + if !ok { + return fn(q) + } + log.Printf("[%s] SQL Begin Transaction", strings.ToUpper(q.provider)) + tx, err := starter.BeginTx(ctx, nil) + if err != nil { + return err + } + + defer func() { + if p := recover(); p != nil { + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(q.provider)) + _ = tx.Rollback() + panic(p) + } + }() + + txQueries := &Queries{ + db: tx, + provider: q.provider, + 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} + + if err := fn(txQueries); err != nil { + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(q.provider)) + _ = tx.Rollback() + return err + } + log.Printf("[%s] SQL Commit Transaction", strings.ToUpper(q.provider)) + return tx.Commit() +} + type Tx struct { *Queries tx *sql.Tx @@ -202,6 +249,7 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { if err != nil { return nil, err } + log.Printf("[%s] SQL Begin Transaction", strings.ToUpper(db.provider)) q := &Queries{ db: sqlTx, provider: db.provider, @@ -222,11 +270,13 @@ func (db *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { // Commit commits the transaction. func (tx *Tx) Commit() error { + log.Printf("[%s] SQL Commit Transaction", strings.ToUpper(tx.provider)) return tx.tx.Commit() } // Rollback aborts the transaction. func (tx *Tx) Rollback() error { + log.Printf("[%s] SQL Rollback Transaction", strings.ToUpper(tx.provider)) return tx.tx.Rollback() } diff --git a/integration/valkyrie/comment.go b/integration/valkyrie/comment.go index 79c7a9e..829eafd 100644 --- a/integration/valkyrie/comment.go +++ b/integration/valkyrie/comment.go @@ -168,13 +168,23 @@ func (q *Queries) executeCommentCreate(ctx context.Context, input CommentCreateI } idCol := "id" - - res, err := executeInsert(ctx, q, "Comment", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + hasRelations := selects != nil && (selects.Post != nil || selects.Author != nil) + + var res *Comment + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "Comment", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadCommentRelations(ctx, []*Comment{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "Comment", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadCommentRelations(ctx, []*Comment{res}, selects); err != nil { + if err != nil { return nil, err } diff --git a/integration/valkyrie/post.go b/integration/valkyrie/post.go index dbccd91..5b49f8a 100644 --- a/integration/valkyrie/post.go +++ b/integration/valkyrie/post.go @@ -155,13 +155,23 @@ func (q *Queries) executePostCreate(ctx context.Context, input PostCreateInput, } idCol := "id" - - res, err := executeInsert(ctx, q, "Post", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + hasRelations := selects != nil && (selects.Author != nil || selects.Comments != nil || selects.Categories != nil) + + var res *Post + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "Post", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadPostRelations(ctx, []*Post{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "Post", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadPostRelations(ctx, []*Post{res}, selects); err != nil { + if err != nil { return nil, err } diff --git a/integration/valkyrie/profile.go b/integration/valkyrie/profile.go index 200fd03..201d0d9 100644 --- a/integration/valkyrie/profile.go +++ b/integration/valkyrie/profile.go @@ -127,13 +127,23 @@ func (q *Queries) executeProfileCreate(ctx context.Context, input ProfileCreateI } idCol := "id" - - res, err := executeInsert(ctx, q, "Profile", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + hasRelations := selects != nil && (selects.User != nil) + + var res *Profile + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "Profile", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadProfileRelations(ctx, []*Profile{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "Profile", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadProfileRelations(ctx, []*Profile{res}, selects); err != nil { + if err != nil { return nil, err } diff --git a/integration/valkyrie/user.go b/integration/valkyrie/user.go index f4d4aed..0671723 100644 --- a/integration/valkyrie/user.go +++ b/integration/valkyrie/user.go @@ -161,13 +161,23 @@ func (q *Queries) executeUserCreate(ctx context.Context, input UserCreateInput, } idCol := "id" + hasRelations := selects != nil && (selects.Profile != nil || selects.Posts != nil || selects.Comments != nil || selects.ReferredBy != nil || selects.Referrals != nil) - res, err := executeInsert(ctx, q, "User", cols, vals, returningCols, idCol, scanFunc) - if err != nil { - return nil, err + var res *User + var err error + if hasRelations { + err = q.transaction(ctx, func(txQ *Queries) error { + var err error + res, err = executeInsert(ctx, txQ, "User", cols, vals, returningCols, idCol, scanFunc) + if err != nil { + return err + } + return txQ.loadUserRelations(ctx, []*User{res}, selects) + }) + } else { + res, err = executeInsert(ctx, q, "User", cols, vals, returningCols, idCol, scanFunc) } - - if err := q.loadUserRelations(ctx, []*User{res}, selects); err != nil { + if err != nil { return nil, err }