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
10 changes: 6 additions & 4 deletions cli/handleGenerate.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,15 +47,17 @@ func handleGenerate() {
pkgName = "valkyrie"
}

content, err := generator.GenerateClient(*schemaDef, pkgName, embedRelDir, config.Output.Migrations)
outputs, err := generator.GenerateClient(*schemaDef, pkgName, embedRelDir, config.Output.Migrations)
if err != nil {
fmt.Printf("failed to generate client: %v\n", err)
return
}

if err := os.WriteFile(filepath.Join(config.Output.Client, "client.go"), []byte(content), 0644); err != nil {
fmt.Println(err)
return
for filename, content := range outputs {
if err := os.WriteFile(filepath.Join(config.Output.Client, filename), []byte(content), 0644); err != nil {
fmt.Println(err)
return
}
}

fmt.Println("Generating Client...")
Expand Down
54 changes: 45 additions & 9 deletions generator/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,14 +20,20 @@ type templateData struct {
Schema schema.Schema
}

func GenerateClient(sch schema.Schema, pkgName string, embedPath string, defaultDiskPath string) (string, error) {
type modelTemplateData struct {
PackageName string
Model *schema.Model
}

func GenerateClient(sch schema.Schema, pkgName string, embedPath string, defaultDiskPath string) (map[string]string, error) {
tmpl := template.New("").Funcs(template.FuncMap{
"capitalize": capitalize,
"lowercase": lowercase,
"capitalize": capitalize,
"lowercase": lowercase,
"fkForRelation": fkForRelation,
})
tmpl, err := tmpl.ParseFS(templatesFS, "templates/*.gotpl")
if err != nil {
return "", err
return nil, err
}

var embedDir string
Expand All @@ -43,26 +49,56 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
Schema: sch,
}

outputs := make(map[string]string)

var buf bytes.Buffer
// sequentially !!!!
files := []string{
"header.gotpl",
"enums.gotpl",
"client.gotpl",
"tx.gotpl",
"builders_create.gotpl",
"delegates.gotpl",
}
for _, file := range files {
if err := tmpl.ExecuteTemplate(&buf, file, data); err != nil {
return "", err
return nil, err
}
}

formatted, err := format.Source(buf.Bytes())
if err != nil {
return buf.String(), err
return nil, err
}
outputs["client.go"] = string(formatted)

for _, m := range sch.Models {
var mBuf bytes.Buffer
mData := modelTemplateData{
PackageName: pkgName,
Model: m,
}

if err := tmpl.ExecuteTemplate(&mBuf, "model_header.gotpl", mData); err != nil {
return nil, err
}

mFiles := []string{
"model_structs.gotpl",
"model_create.gotpl",
"model_relations.gotpl",
}
for _, file := range mFiles {
if err := tmpl.ExecuteTemplate(&mBuf, file, mData); err != nil {
return nil, err
}
}

mFormatted, err := format.Source(mBuf.Bytes())
if err != nil {
return nil, err
}
outputs[lowercase(m.Name)+".go"] = string(mFormatted)
}

return string(formatted), nil
return outputs, nil
}
9 changes: 7 additions & 2 deletions generator/generator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,14 @@ func TestGenerateClient(t *testing.T) {
t.Fatalf("parser errors: %v", errs)
}

code, err := GenerateClient(*sch, "client", "migrations/*.sql", "migrations")
outputs, err := GenerateClient(*sch, "client", "migrations/*.sql", "migrations")
if err != nil {
t.Fatalf("failed to generate client: %v\nCode output:\n%s", err, code)
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 {") {
Expand Down
13 changes: 13 additions & 0 deletions generator/helpers.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package generator

import (
"strings"
"valkyrie/schema"
)

func capitalize(s string) string {
Expand All @@ -23,3 +24,15 @@ func lowercase(s string) string {
}
return strings.ToLower(s[:1]) + s[1:]
}

// returns the relation name if this scalar field is a FK for a relation on the model, empty string if not
func fkForRelation(model *schema.Model, field *schema.ScalarField) string {
for _, rel := range model.RelationFields {
for _, fk := range rel.FKFields {
if fk.Name == field.Name {
return rel.Name
}
}
}
return ""
}
145 changes: 125 additions & 20 deletions generator/templates/builders_create.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -43,25 +43,40 @@ func executeInsert[M any](
idCol string,
scanFunc func(record *M, cols []string) []any,
) (*M, error) {
placeholders := make([]string, len(cols))
var sb strings.Builder
sb.Grow(128 + len(table) + len(cols)*15 + len(returningCols)*15)

sb.WriteString("INSERT INTO ")
sb.WriteString(q.dialect.Quote(table))
sb.WriteString(" (")
for i, col := range cols {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(col)
}
sb.WriteString(") VALUES (")
for i := range cols {
placeholders[i] = q.dialect.BindVar(i + 1)
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(q.dialect.BindVar(i + 1))
}
sb.WriteString(")")

var res M
quotedReturningCols := make([]string, len(returningCols))
for i, col := range returningCols {
quotedReturningCols[i] = q.dialect.Quote(col)
if q.dialect.SupportsReturning() && len(returningCols) > 0 {
sb.WriteString(" RETURNING ")
for i, col := range returningCols {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(q.dialect.Quote(col))
}
}
query := sb.String()

query := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)",
q.dialect.Quote(table),
strings.Join(cols, ", "),
strings.Join(placeholders, ", "),
)

var res M
if q.dialect.SupportsReturning() {
query += " RETURNING " + strings.Join(quotedReturningCols, ", ")
row := q.db.QueryRowContext(ctx, query, vals...)

scanTargets := scanFunc(&res, returningCols)
Expand All @@ -78,8 +93,9 @@ func executeInsert[M any](
}

var idVal any
quotedIdCol := q.dialect.Quote(idCol)
for i, c := range cols {
if c == q.dialect.Quote(idCol) {
if c == quotedIdCol {
idVal = vals[i]
break
}
Expand All @@ -92,15 +108,104 @@ func executeInsert[M any](
idVal = lastID
}

selectQuery := fmt.Sprintf("SELECT %s FROM %s WHERE %s = ?",
strings.Join(quotedReturningCols, ", "),
q.dialect.Quote(table),
q.dialect.Quote(idCol),
)
row := q.db.QueryRowContext(ctx, selectQuery, idVal)
var selectSb strings.Builder
selectSb.Grow(64 + len(returningCols)*15 + len(table) + len(idCol))
selectSb.WriteString("SELECT ")
for i, col := range returningCols {
if i > 0 {
selectSb.WriteString(", ")
}
selectSb.WriteString(q.dialect.Quote(col))
}
selectSb.WriteString(" FROM ")
selectSb.WriteString(q.dialect.Quote(table))
selectSb.WriteString(" WHERE ")
selectSb.WriteString(q.dialect.Quote(idCol))
selectSb.WriteString(" = ?")

row := q.db.QueryRowContext(ctx, selectSb.String(), idVal)
scanTargets := scanFunc(&res, returningCols)
if err := row.Scan(scanTargets...); err != nil {
return nil, err
}
return &res, nil
}


func loadRelation[P any, C any](
ctx context.Context,
q *Queries,
parents []*P,
parentKey func(*P) (string, bool),
table string,
fkCol string,
returningCols []string,
scan func(*sql.Rows, *C) error,
childKey func(*C) (string, bool),
assign func(*P, []*C),
) ([]*C, error) {
var parentKeys []any
for _, p := range parents {
if p == nil {
continue
}
if key, ok := parentKey(p); ok {
parentKeys = append(parentKeys, key)
}
}
if len(parentKeys) == 0 {
return nil, nil
}

var sb strings.Builder
sb.Grow(128 + len(returningCols)*15 + len(table) + len(fkCol) + len(parentKeys)*3)
sb.WriteString("SELECT ")
for i, col := range returningCols {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(q.dialect.Quote(col))
}
sb.WriteString(" FROM ")
sb.WriteString(q.dialect.Quote(table))
sb.WriteString(" WHERE ")
sb.WriteString(q.dialect.Quote(fkCol))
sb.WriteString(" IN (")
sb.WriteString(q.bindVars(len(parentKeys)))
sb.WriteString(")")
query := sb.String()

rows, err := q.db.QueryContext(ctx, query, parentKeys...)
if err != nil {
return nil, err
}
defer rows.Close()

childMap := make(map[string][]*C)
var allChildren []*C

for rows.Next() {
var child C
if err := scan(rows, &child); err != nil {
return nil, err
}
if key, ok := childKey(&child); ok {
childMap[key] = append(childMap[key], &child)
}
allChildren = append(allChildren, &child)
}
if err := rows.Err(); err != nil {
return nil, err
}

for _, p := range parents {
if p == nil {
continue
}
if key, ok := parentKey(p); ok {
assign(p, childMap[key])
}
}

return allChildren, nil
}
37 changes: 24 additions & 13 deletions generator/templates/client.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -4,15 +4,19 @@ type Dialect interface {
SupportsReturning() bool
}

{{- if or (eq .Schema.Datasource.Provider "postgres") (eq .Schema.Datasource.Provider "postgresql") }}
type postgresDialect struct{}
func (postgresDialect) Quote(ident string) string { return `"` + ident + `"` }
func (postgresDialect) BindVar(idx int) string { return fmt.Sprintf("$%d", idx) }
func (postgresDialect) SupportsReturning() bool { return true }
{{- end }}

{{- if or (eq .Schema.Datasource.Provider "sqlite") (eq .Schema.Datasource.Provider "sqlite3") }}
type sqliteDialect struct{}
func (sqliteDialect) Quote(ident string) string { return `"` + ident + `"` }
func (sqliteDialect) BindVar(idx int) string { return "?" }
func (sqliteDialect) SupportsReturning() bool { return true }
{{- end }}

type DBTX interface {
ExecContext(context.Context, string, ...any) (sql.Result, error)
Expand Down Expand Up @@ -45,22 +49,14 @@ func Open(provider, dataSourceName string) (*DB, error) {
return nil, err
}

var d Dialect
switch provider {
case "postgres", "postgresql":
d = postgresDialect{}
case "sqlite", "sqlite3":
d = sqliteDialect{}

default:
sqlDB.Close()
return nil, fmt.Errorf("unsupported database provider: %s", provider)
}

q := &Queries{
db: sqlDB,
provider: provider,
dialect: d,
{{- if or (eq .Schema.Datasource.Provider "postgres") (eq .Schema.Datasource.Provider "postgresql") }}
dialect: postgresDialect{},
{{- else if or (eq .Schema.Datasource.Provider "sqlite") (eq .Schema.Datasource.Provider "sqlite3") }}
dialect: sqliteDialect{},
{{- end }}
{{- range $enum := .Schema.Enums }}
{{ $enum.Name }}: {{ $enum.Name }},
{{- end }}
Expand Down Expand Up @@ -104,3 +100,18 @@ func (db *DB) RunMigrations(ctx context.Context) error {
return goose.UpContext(ctx, db.sqlDB, "{{ .DefaultDiskPath }}")
}
{{- end }}

func (q *Queries) bindVars(count int) string {
if count <= 0 {
return ""
}
var sb strings.Builder
sb.Grow(count * 3)
for i := 0; i < count; i++ {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(q.dialect.BindVar(i + 1))
}
return sb.String()
}
Loading
Loading