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
15 changes: 13 additions & 2 deletions cli/handleGenerate.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,14 +47,25 @@ func handleGenerate() {
pkgName = "valk"
}

outputs, err := generator.GenerateClient(*schemaDef, pkgName, embedRelDir, config.Output.Migrations, config.Log)
parentImportPath, err := generator.ResolveImportPath(config.Output.Client)
if err != nil {
fmt.Printf("failed to resolve parent import path: %v\n", err)
return
}

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

for filename, content := range outputs {
if err := os.WriteFile(filepath.Join(config.Output.Client, filename), []byte(content), 0644); err != nil {
outPath := filepath.Join(config.Output.Client, filename)
if err := os.MkdirAll(filepath.Dir(outPath), 0755); err != nil {
fmt.Println(err)
return
}
if err := os.WriteFile(outPath, []byte(content), 0644); err != nil {
fmt.Println(err)
return
}
Expand Down
121 changes: 116 additions & 5 deletions generator/generator.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,13 @@ package generator
import (
"bytes"
"embed"
"fmt"
"go/format"
"os"
"path/filepath"
"strings"
"text/template"

"github.com/voidclancy/valk/schema"
)

Expand All @@ -22,15 +26,77 @@ type templateData struct {
}

type modelTemplateData struct {
PackageName string
Model *schema.Model
PackageName string
Model *schema.Model
ParentImportPath string
ParentPackageName string
}

func ResolveImportPath(clientDir string) (string, error) {
absClientDir, err := filepath.Abs(clientDir)
if err != nil {
return "", err
}

current := absClientDir
for {
modFile := filepath.Join(current, "go.mod")
if _, err := os.Stat(modFile); err == nil {
content, err := os.ReadFile(modFile)
if err != nil {
return "", err
}
var modName string
lines := strings.Split(string(content), "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, "module ") {
modName = strings.TrimSpace(strings.TrimPrefix(line, "module"))
break
}
}
if modName == "" {
return "", fmt.Errorf("go.mod found but no module declaration found")
}

rel, err := filepath.Rel(current, absClientDir)
if err != nil {
return "", err
}
if rel == "." {
return modName, nil
}
return filepath.ToSlash(filepath.Join(modName, rel)), nil
}

parent := filepath.Dir(current)
if parent == current {
break
}
current = parent
}

return filepath.Base(clientDir), nil
}

func GenerateClient(sch schema.Schema, pkgName string, embedPath string, defaultDiskPath string, defaultLogs []string) (map[string]string, error) {
func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, embedPath string, defaultDiskPath string, defaultLogs []string) (map[string]string, error) {
tmpl := template.New("").Funcs(template.FuncMap{
"capitalize": capitalize,
"lowercase": lowercase,
"fkForRelation": fkForRelation,
"fieldPredType": func(f *schema.ScalarField, parentPkg string) string {
if f.EnumRef != nil {
if f.IsArray {
return "[]" + parentPkg + "." + f.EnumRef.Name + "Type"
}
return parentPkg + "." + f.EnumRef.Name + "Type"
}
t := f.GoType
if f.Optional {
t = strings.TrimPrefix(t, "*")
}
return t
},
"hasLog": func(level string) bool {
for _, l := range defaultLogs {
if l == "all" || l == level {
Expand All @@ -47,6 +113,30 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
}
return false
},
"hasJsonField": func(m *schema.Model) bool {
for _, sf := range m.ScalarFields {
if sf.Type == "Json" || strings.Contains(sf.GoType, "json.RawMessage") {
return true
}
}
return false
},
"hasTimeField": func(m *schema.Model) bool {
for _, sf := range m.ScalarFields {
if sf.Type == "DateTime" || strings.Contains(sf.GoType, "time.Time") {
return true
}
}
return false
},
"hasStringField": func(m *schema.Model) bool {
for _, sf := range m.ScalarFields {
if sf.GoType == "string" || strings.Contains(sf.GoType, "string") {
return true
}
}
return false
},
})
tmpl, err := tmpl.ParseFS(templatesFS, "templates/*.gotpl")
if err != nil {
Expand Down Expand Up @@ -76,6 +166,7 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
"client.gotpl",
"tx.gotpl",
"builders_create.gotpl",
"builders_read.gotpl",
"relations_runtime.gotpl",
}
for _, file := range files {
Expand All @@ -93,8 +184,10 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
for _, m := range sch.Models {
var mBuf bytes.Buffer
mData := modelTemplateData{
PackageName: pkgName,
Model: m,
PackageName: pkgName,
Model: m,
ParentImportPath: parentImportPath,
ParentPackageName: pkgName,
}

if err := tmpl.ExecuteTemplate(&mBuf, "model_header.gotpl", mData); err != nil {
Expand All @@ -104,6 +197,7 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
mFiles := []string{
"model_structs.gotpl",
"model_create.gotpl",
"model_read.gotpl",
"model_relations.gotpl",
}
for _, file := range mFiles {
Expand All @@ -117,6 +211,23 @@ func GenerateClient(sch schema.Schema, pkgName string, embedPath string, default
return nil, err
}
outputs[lowercase(m.Name)+".go"] = string(mFormatted)

// Generate the sub-package predicate file (e.g. user/user.go)
var pBuf bytes.Buffer
pData := modelTemplateData{
PackageName: lowercase(m.Name),
Model: m,
ParentImportPath: parentImportPath,
ParentPackageName: pkgName,
}
if err := tmpl.ExecuteTemplate(&pBuf, "model_predicate.gotpl", pData); err != nil {
return nil, err
}
pFormatted, err := format.Source(pBuf.Bytes())
if err != nil {
return nil, err
}
outputs[lowercase(m.Name)+"/"+lowercase(m.Name)+".go"] = string(pFormatted)
}

return outputs, nil
Expand Down
2 changes: 1 addition & 1 deletion generator/generator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ func TestGenerateClient_NativeDBConstraints(t *testing.T) {
},
}

outputs, err := GenerateClient(sch, "valk", "", "", nil)
outputs, err := GenerateClient(sch, "valk", "github.com/voidclancy/valk", "", "", nil)
if err != nil {
t.Fatalf("failed to generate client: %v", err)
}
Expand Down
Loading
Loading