Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bed80b046b | ||
|
|
495a21b67b | ||
|
|
a32647ee16 | ||
|
|
08e1417393 | ||
|
|
70282fff73 | ||
|
|
43265dac0f | ||
|
|
66b90ca54b | ||
|
|
47108809aa | ||
|
|
720476fd6e | ||
|
|
572d03fe42 | ||
|
|
f1b9079b2d | ||
|
|
bb671c3680 | ||
|
|
bc8284db25 | ||
|
|
29e747393d | ||
|
|
df980a3434 | ||
|
|
778379538b | ||
|
|
7a9219b6e3 | ||
|
|
fd9c37cd25 | ||
|
|
0235a28add | ||
|
|
d961536186 | ||
|
|
d36806047b | ||
|
|
948419ffd3 | ||
|
|
b38f53c603 | ||
|
|
ccba53c494 | ||
|
|
53327b9a5a | ||
|
|
734b14d48d | ||
|
|
938f0ed51f | ||
|
|
6e2e7eb19e |
@@ -1,4 +1,4 @@
|
|||||||
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc vet fmt fmt-check staticcheck govulncheck check
|
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version rerelease godoc vet fmt fmt-check staticcheck govulncheck check
|
||||||
|
|
||||||
# Binary name
|
# Binary name
|
||||||
BINARY_NAME=relspec
|
BINARY_NAME=relspec
|
||||||
@@ -260,5 +260,13 @@ release-version: lint fmt-check ## Run lint and format check, then auto-incremen
|
|||||||
git push origin HEAD "$$NEXT"; \
|
git push origin HEAD "$$NEXT"; \
|
||||||
echo "Pushed $$NEXT — release workflow triggered"
|
echo "Pushed $$NEXT — release workflow triggered"
|
||||||
|
|
||||||
|
rerelease: lint fmt-check ## Move the latest tag to HEAD and force push it
|
||||||
|
@TAG=$$(git describe --tags --abbrev=0 2>/dev/null); \
|
||||||
|
if [ -z "$$TAG" ]; then echo "No existing tags found"; exit 1; fi; \
|
||||||
|
echo "Moving $$TAG to $$(git rev-parse --short HEAD)"; \
|
||||||
|
git tag -f -a "$$TAG" -m "Release $$TAG" HEAD; \
|
||||||
|
git push --force origin "$$TAG"; \
|
||||||
|
echo "Pushed $$TAG — release workflow triggered"
|
||||||
|
|
||||||
help: ## Display this help screen
|
help: ## Display this help screen
|
||||||
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
|
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
|
||||||
|
|||||||
@@ -23,6 +23,8 @@ go install -v git.warky.dev/wdevs/relspecgo/cmd/relspec@latest
|
|||||||
| **Readers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqldir` `sqlite` `typeorm` `yaml` |
|
| **Readers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqldir` `sqlite` `typeorm` `yaml` |
|
||||||
| **Writers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqlexec` `sqlite` `template` `typeorm` `yaml` |
|
| **Writers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqlexec` `sqlite` `template` `typeorm` `yaml` |
|
||||||
|
|
||||||
|
See [docs/FORMAT_EXAMPLES.md](docs/FORMAT_EXAMPLES.md) for usage examples covering every format.
|
||||||
|
|
||||||
## Commands
|
## Commands
|
||||||
|
|
||||||
### `convert` — Schema conversion
|
### `convert` — Schema conversion
|
||||||
@@ -40,6 +42,26 @@ relspec convert --from pgsql --from-conn "postgres://..." --to sqlite --to-path
|
|||||||
|
|
||||||
# Multiple input files merged
|
# Multiple input files merged
|
||||||
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
||||||
|
|
||||||
|
# Watch mode: regenerate whenever the source file(s) change (Ctrl-C to stop)
|
||||||
|
relspec convert --from dbml --from-path schema.dbml --to gorm --to-path models/ --package models --watch
|
||||||
|
```
|
||||||
|
|
||||||
|
`--watch` works with `--from-path` and `--from-list` (not live database
|
||||||
|
connections or `--dry-run`). Source files are polled every `--watch-interval`
|
||||||
|
(default 500ms), a directory source is watched recursively, and the output path
|
||||||
|
is ignored so generating into the source tree does not loop. Conversion errors
|
||||||
|
are printed and watching continues.
|
||||||
|
|
||||||
|
### `batch` — Convert many inputs in one run
|
||||||
|
|
||||||
|
Converts each input independently (one output per input, unlike `--from-list`
|
||||||
|
which merges). `--input` takes paths or globs; outputs go to `--to-dir`.
|
||||||
|
Use `--keep-going` to continue past failures (exit is still non-zero) and
|
||||||
|
`--dry-run` to validate without writing. For named workflows see `relspec job run`.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec batch --from dbml --input "schemas/*.dbml" --to json --to-dir out/
|
||||||
```
|
```
|
||||||
|
|
||||||
PostgreSQL connections opened by relspec set `application_name` by default to
|
PostgreSQL connections opened by relspec set `application_name` by default to
|
||||||
@@ -216,6 +238,23 @@ see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
|
|||||||
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
||||||
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
||||||
|
|
||||||
|
#### Custom type mapping
|
||||||
|
|
||||||
|
Override the built-in SQL → Go mapping of the `bun` and `gorm` writers with the
|
||||||
|
repeatable `--type-map sqltype=gotype` flag:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec convert --from pgsql --from-conn "$DSN" --to gorm --to-path models.go \
|
||||||
|
--type-map uuid=string --type-map jsonb=json.RawMessage
|
||||||
|
```
|
||||||
|
|
||||||
|
SQL type names are matched case-insensitively on the base type (modifiers such
|
||||||
|
as `(10,2)` are ignored; aliases like `int4` resolve to `integer`). NOT NULL
|
||||||
|
columns use the Go type verbatim, nullable columns get a `*` prefix (unless the
|
||||||
|
type is already a pointer, slice, map or `any`), and arrays become `[]gotype`.
|
||||||
|
Unmapped types keep their defaults. The flag does not add imports: use types
|
||||||
|
that need none, or add the import afterwards (e.g. with `goimports`).
|
||||||
|
|
||||||
## Contributing
|
## Contributing
|
||||||
|
|
||||||
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
batchSourceType string
|
||||||
|
batchInputs []string
|
||||||
|
batchTargetType string
|
||||||
|
batchTargetDir string
|
||||||
|
batchPackageName string
|
||||||
|
batchSchemaFilter string
|
||||||
|
batchFlattenSchema bool
|
||||||
|
batchNullableTypes string
|
||||||
|
batchNullableArrays string
|
||||||
|
batchContinueOnError bool
|
||||||
|
batchKeepGoing bool
|
||||||
|
batchDryRun bool
|
||||||
|
)
|
||||||
|
|
||||||
|
var batchCmd = &cobra.Command{
|
||||||
|
Use: "batch",
|
||||||
|
Short: "Convert many input files to a target format in one run",
|
||||||
|
Long: `Convert each input file independently to the target format.
|
||||||
|
|
||||||
|
Unlike 'convert --from-list', which merges all inputs into one output, batch
|
||||||
|
mode writes one output per input into --to-dir. The output is named after the
|
||||||
|
input file (without its extension). Directory-style targets (gorm, bun,
|
||||||
|
drizzle) get a sub-directory per input.
|
||||||
|
|
||||||
|
Inputs are given with --input, which accepts file paths and glob patterns and
|
||||||
|
may be repeated or comma-separated. Inputs are processed in sorted order and
|
||||||
|
duplicates are removed. The command exits non-zero if any input fails.
|
||||||
|
|
||||||
|
For named, multi-step workflows use 'relspec job run' instead.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
# Convert every DBML file in a directory to JSON
|
||||||
|
relspec batch --from dbml --input "schemas/*.dbml" --to json --to-dir out/
|
||||||
|
|
||||||
|
# Convert specific files to GORM models, one package directory per input
|
||||||
|
relspec batch --from json --input a.json,b.json \
|
||||||
|
--to gorm --to-dir models/ --package models
|
||||||
|
|
||||||
|
# Validate everything first, writing nothing
|
||||||
|
relspec batch --from yaml --input "specs/*.yaml" --to pgsql --to-dir sql/ --dry-run
|
||||||
|
|
||||||
|
# Report all failures instead of stopping at the first
|
||||||
|
relspec batch --from json --input "*.json" --to yaml --to-dir out/ --keep-going`,
|
||||||
|
RunE: runBatch,
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
batchCmd.Flags().StringVar(&batchSourceType, "from", "", "Source format for every input (dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, sqlite)")
|
||||||
|
batchCmd.Flags().StringSliceVar(&batchInputs, "input", nil, "Input file path or glob pattern (repeatable, comma-separated)")
|
||||||
|
batchCmd.Flags().StringVar(&batchTargetType, "to", "", "Target format")
|
||||||
|
batchCmd.Flags().StringVar(&batchTargetDir, "to-dir", "", "Output directory; one output per input is written here")
|
||||||
|
batchCmd.Flags().StringVar(&batchPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
||||||
|
batchCmd.Flags().StringVar(&batchSchemaFilter, "schema", "", "Filter to a specific schema by name")
|
||||||
|
batchCmd.Flags().BoolVar(&batchFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table")
|
||||||
|
batchCmd.Flags().StringVar(&batchNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm)")
|
||||||
|
batchCmd.Flags().StringVar(&batchNullableArrays, "array-nullable", "", "Nullable array representation for the Bun writer")
|
||||||
|
batchCmd.Flags().BoolVar(&batchContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL (pgsql output only)")
|
||||||
|
batchCmd.Flags().BoolVar(&batchKeepGoing, "keep-going", false, "Process remaining inputs after a failure; still exits non-zero")
|
||||||
|
batchCmd.Flags().BoolVar(&batchDryRun, "dry-run", false, "Read and validate every input and print the plan without writing any output")
|
||||||
|
|
||||||
|
for _, f := range []string{"from", "input", "to", "to-dir"} {
|
||||||
|
if err := batchCmd.MarkFlagRequired(f); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Error marking %s flag as required: %v\n", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// batchDirTargets are writers that emit a directory of files rather than one file.
|
||||||
|
var batchDirTargets = map[string]bool{"gorm": true, "bun": true, "drizzle": true}
|
||||||
|
|
||||||
|
// batchExtensions maps single-file target formats to their output extension.
|
||||||
|
var batchExtensions = map[string]string{
|
||||||
|
"dbml": ".dbml", "dctx": ".dctx", "drawdb": ".ddb", "json": ".json",
|
||||||
|
"yaml": ".yaml", "yml": ".yaml", "pgsql": ".sql", "postgres": ".sql",
|
||||||
|
"postgresql": ".sql", "sql": ".sql", "mssql": ".sql", "sqlserver": ".sql",
|
||||||
|
"mssql2016": ".sql", "mssql2017": ".sql", "mssql2019": ".sql", "mssql2022": ".sql",
|
||||||
|
"sqlite": ".sql", "sqlite3": ".sql", "prisma": ".prisma", "typeorm": ".ts",
|
||||||
|
"graphql": ".graphql", "gql": ".graphql",
|
||||||
|
}
|
||||||
|
|
||||||
|
// expandBatchInputs resolves paths and glob patterns into a sorted,
|
||||||
|
// de-duplicated file list. A pattern that matches nothing is an error.
|
||||||
|
func expandBatchInputs(patterns []string) ([]string, error) {
|
||||||
|
seen := map[string]bool{}
|
||||||
|
var files []string
|
||||||
|
for _, p := range patterns {
|
||||||
|
p = strings.TrimSpace(p)
|
||||||
|
if p == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
matches, err := filepath.Glob(p)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid pattern %q: %w", p, err)
|
||||||
|
}
|
||||||
|
if len(matches) == 0 {
|
||||||
|
return nil, fmt.Errorf("no files match %q", p)
|
||||||
|
}
|
||||||
|
for _, m := range matches {
|
||||||
|
if !seen[m] {
|
||||||
|
seen[m] = true
|
||||||
|
files = append(files, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(files) == 0 {
|
||||||
|
return nil, fmt.Errorf("no input files given")
|
||||||
|
}
|
||||||
|
sort.Strings(files)
|
||||||
|
return files, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// batchOutputPaths returns the output path for each input. It errors when two
|
||||||
|
// inputs would collide on the same output name.
|
||||||
|
func batchOutputPaths(files []string, targetType, dir string) ([]string, error) {
|
||||||
|
key := strings.ToLower(targetType)
|
||||||
|
ext := ""
|
||||||
|
if !batchDirTargets[key] {
|
||||||
|
var ok bool
|
||||||
|
if ext, ok = batchExtensions[key]; !ok {
|
||||||
|
return nil, fmt.Errorf("unsupported target format: %s", targetType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
outs := make([]string, len(files))
|
||||||
|
owner := map[string]string{}
|
||||||
|
for i, f := range files {
|
||||||
|
stem := strings.TrimSuffix(filepath.Base(f), filepath.Ext(f))
|
||||||
|
out := filepath.Join(dir, stem+ext)
|
||||||
|
if prev, dup := owner[out]; dup {
|
||||||
|
return nil, fmt.Errorf("inputs %s and %s would both write %s", prev, f, out)
|
||||||
|
}
|
||||||
|
owner[out] = f
|
||||||
|
outs[i] = out
|
||||||
|
}
|
||||||
|
return outs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func runBatch(cmd *cobra.Command, args []string) error {
|
||||||
|
files, err := expandBatchInputs(batchInputs)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
outs, err := batchOutputPaths(files, batchTargetType, batchTargetDir)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== RelSpec Batch Converter ===\n")
|
||||||
|
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
|
||||||
|
fmt.Fprintf(os.Stderr, "Inputs: %d file(s), %s -> %s\n\n", len(files), batchSourceType, batchTargetType)
|
||||||
|
|
||||||
|
out := outWriter(cmd)
|
||||||
|
if batchDryRun {
|
||||||
|
fmt.Fprintf(out, "RelSpec batch plan (dry run - nothing written):\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
var failed []string
|
||||||
|
for i, f := range files {
|
||||||
|
fmt.Fprintf(os.Stderr, "[%d/%d] %s -> %s\n", i+1, len(files), f, outs[i])
|
||||||
|
if err := processBatchItem(cmd, f, outs[i]); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, " ✗ %v\n", err)
|
||||||
|
failed = append(failed, fmt.Sprintf("%s: %v", f, err))
|
||||||
|
if !batchKeepGoing {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
fmt.Fprintf(os.Stderr, " ✓ done\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Batch Complete: %d ok, %d failed ===\n", len(files)-len(failed), len(failed))
|
||||||
|
if len(failed) > 0 {
|
||||||
|
return fmt.Errorf("batch finished with %d failure(s):\n %s", len(failed), strings.Join(failed, "\n "))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func processBatchItem(cmd *cobra.Command, in, outPath string) error {
|
||||||
|
db, err := readDatabaseForConvert(batchSourceType, in, "")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to read source: %w", err)
|
||||||
|
}
|
||||||
|
finalizeCommentedRefs(db, stderrWarn)
|
||||||
|
|
||||||
|
if batchDryRun {
|
||||||
|
if err := validateWriteTarget(db, batchTargetType, batchPackageName, batchSchemaFilter, ""); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
w := outWriter(cmd)
|
||||||
|
fmt.Fprintf(w, " %s -> %s (database '%s')\n", in, outPath, db.Name)
|
||||||
|
printDryRunPlan(w, db)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(batchTargetDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("failed to create output directory: %w", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, batchTargetType, outPath, batchPackageName, batchSchemaFilter, batchFlattenSchema, batchNullableTypes, batchNullableArrays, batchContinueOnError, ""); err != nil {
|
||||||
|
return fmt.Errorf("failed to write target: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func saveBatchState(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
a, b, c, d, e, f, g := batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun
|
||||||
|
t.Cleanup(func() {
|
||||||
|
batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun = a, b, c, d, e, f, g
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_ConvertsEachInput(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "b.json"), []string{"posts"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchKeepGoing, batchDryRun = "", false, false
|
||||||
|
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||||
|
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
if err := runBatch(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("batch: %v", err)
|
||||||
|
}
|
||||||
|
for _, name := range []string{"a.yaml", "b.yaml"} {
|
||||||
|
if _, err := os.Stat(filepath.Join(outDir, name)); err != nil {
|
||||||
|
t.Errorf("expected %s: %v", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_DryRunWritesNothing(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchKeepGoing, batchDryRun = "", false, true
|
||||||
|
batchInputs = []string{filepath.Join(dir, "a.json")}
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runBatch(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(outDir); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output directory")
|
||||||
|
}
|
||||||
|
if !strings.Contains(buf.String(), "users") || !strings.Contains(buf.String(), "a.yaml") {
|
||||||
|
t.Errorf("plan incomplete:\n%s", buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_FailureHandling(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "b.json"), []byte("{not json"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "c.json"), []string{"posts"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchDryRun = "", false
|
||||||
|
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
|
||||||
|
// Default: stop at first failure.
|
||||||
|
batchKeepGoing = false
|
||||||
|
err := runBatch(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "b.json") {
|
||||||
|
t.Fatalf("expected failure naming b.json, got %v", err)
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); !os.IsNotExist(statErr) {
|
||||||
|
t.Error("c.json should not be processed without --keep-going")
|
||||||
|
}
|
||||||
|
|
||||||
|
// --keep-going: remaining inputs are processed, exit still fails.
|
||||||
|
batchKeepGoing = true
|
||||||
|
if err := runBatch(cmd, nil); err == nil {
|
||||||
|
t.Fatal("expected non-zero result with --keep-going")
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); statErr != nil {
|
||||||
|
t.Errorf("c.yaml should be written with --keep-going: %v", statErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExpandBatchInputs(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
for _, n := range []string{"b.json", "a.json"} {
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, n), []byte("{}"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
got, err := expandBatchInputs([]string{filepath.Join(dir, "*.json"), filepath.Join(dir, "a.json")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(got) != 2 || filepath.Base(got[0]) != "a.json" || filepath.Base(got[1]) != "b.json" {
|
||||||
|
t.Errorf("want sorted deduped [a b], got %v", got)
|
||||||
|
}
|
||||||
|
if _, err := expandBatchInputs([]string{filepath.Join(dir, "*.nope")}); err == nil {
|
||||||
|
t.Error("unmatched pattern should error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchOutputPaths(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
files []string
|
||||||
|
target string
|
||||||
|
want []string
|
||||||
|
wantErr string
|
||||||
|
}{
|
||||||
|
{"file target", []string{"x/a.dbml"}, "json", []string{"out/a.json"}, ""},
|
||||||
|
{"dir target", []string{"x/a.json"}, "gorm", []string{"out/a"}, ""},
|
||||||
|
{"collision", []string{"x/a.json", "y/a.json"}, "yaml", nil, "both write"},
|
||||||
|
{"unsupported", []string{"a.json"}, "nope", nil, "unsupported target"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := batchOutputPaths(tt.files, tt.target, "out")
|
||||||
|
if tt.wantErr != "" {
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||||
|
t.Fatalf("want error %q, got %v", tt.wantErr, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil || len(got) != len(tt.want) || got[0] != filepath.FromSlash(tt.want[0]) {
|
||||||
|
t.Fatalf("got %v, %v; want %v", got, err, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+140
-15
@@ -3,6 +3,7 @@ package main
|
|||||||
import (
|
import (
|
||||||
stdjson "encoding/json"
|
stdjson "encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -21,6 +22,7 @@ import (
|
|||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/graphql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/graphql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/mssql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/mssql"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/mysql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
||||||
@@ -36,6 +38,7 @@ import (
|
|||||||
wgraphql "git.warky.dev/wdevs/relspecgo/pkg/writers/graphql"
|
wgraphql "git.warky.dev/wdevs/relspecgo/pkg/writers/graphql"
|
||||||
wjson "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
|
wjson "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
|
||||||
wmssql "git.warky.dev/wdevs/relspecgo/pkg/writers/mssql"
|
wmssql "git.warky.dev/wdevs/relspecgo/pkg/writers/mssql"
|
||||||
|
wmysql "git.warky.dev/wdevs/relspecgo/pkg/writers/mysql"
|
||||||
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
|
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
|
||||||
wprisma "git.warky.dev/wdevs/relspecgo/pkg/writers/prisma"
|
wprisma "git.warky.dev/wdevs/relspecgo/pkg/writers/prisma"
|
||||||
wsqlite "git.warky.dev/wdevs/relspecgo/pkg/writers/sqlite"
|
wsqlite "git.warky.dev/wdevs/relspecgo/pkg/writers/sqlite"
|
||||||
@@ -57,6 +60,9 @@ var (
|
|||||||
convertNullableArrays string
|
convertNullableArrays string
|
||||||
convertContinueOnError bool
|
convertContinueOnError bool
|
||||||
convertExtraFields string
|
convertExtraFields string
|
||||||
|
convertDryRun bool
|
||||||
|
convertWatch bool
|
||||||
|
convertWatchInterval time.Duration
|
||||||
)
|
)
|
||||||
|
|
||||||
var convertCmd = &cobra.Command{
|
var convertCmd = &cobra.Command{
|
||||||
@@ -165,7 +171,11 @@ Examples:
|
|||||||
|
|
||||||
# Convert SQLite to PostgreSQL SQL
|
# Convert SQLite to PostgreSQL SQL
|
||||||
relspec convert --from sqlite --from-path database.db \
|
relspec convert --from sqlite --from-path database.db \
|
||||||
--to pgsql --to-path schema.sql`,
|
--to pgsql --to-path schema.sql
|
||||||
|
|
||||||
|
# Regenerate GORM models every time the DBML file changes
|
||||||
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
|
--to gorm --to-path models/ --package models --watch`,
|
||||||
RunE: runConvert,
|
RunE: runConvert,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -185,6 +195,11 @@ func init() {
|
|||||||
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
||||||
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
||||||
|
|
||||||
|
convertCmd.Flags().BoolVar(&convertDryRun, "dry-run", false, "Read and validate the input and print the plan without writing any output")
|
||||||
|
|
||||||
|
convertCmd.Flags().BoolVar(&convertWatch, "watch", false, "Watch the source files (--from-path or --from-list) and regenerate the output whenever they change")
|
||||||
|
convertCmd.Flags().DurationVar(&convertWatchInterval, "watch-interval", 500*time.Millisecond, "Polling interval used by --watch")
|
||||||
|
|
||||||
err := convertCmd.MarkFlagRequired("from")
|
err := convertCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
||||||
@@ -200,6 +215,13 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func runConvert(cmd *cobra.Command, args []string) error {
|
func runConvert(cmd *cobra.Command, args []string) error {
|
||||||
|
if convertWatch {
|
||||||
|
return runConvertWatch(cmd.Context(), os.Stderr, func() error { return runConvertOnce(cmd) })
|
||||||
|
}
|
||||||
|
return runConvertOnce(cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runConvertOnce(cmd *cobra.Command) error {
|
||||||
fmt.Fprintf(os.Stderr, "\n=== RelSpec Schema Converter ===\n")
|
fmt.Fprintf(os.Stderr, "\n=== RelSpec Schema Converter ===\n")
|
||||||
fmt.Fprintf(os.Stderr, "Started at: %s\n\n", getCurrentTimestamp())
|
fmt.Fprintf(os.Stderr, "Started at: %s\n\n", getCurrentTimestamp())
|
||||||
|
|
||||||
@@ -240,6 +262,22 @@ func runConvert(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
fmt.Fprintf(os.Stderr, " Found: %d table(s)\n\n", totalTables)
|
fmt.Fprintf(os.Stderr, " Found: %d table(s)\n\n", totalTables)
|
||||||
|
|
||||||
|
if convertDryRun {
|
||||||
|
if err := validateWriteTarget(db, convertTargetType, convertPackageName, convertSchemaFilter, convertExtraFields); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec convert plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Input: %s database '%s'\n", convertSourceType, db.Name)
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", convertTargetType, convertTargetPath)
|
||||||
|
if convertSchemaFilter != "" {
|
||||||
|
fmt.Fprintf(out, " Schema filter: %s\n", convertSchemaFilter)
|
||||||
|
}
|
||||||
|
printDryRunPlan(out, db)
|
||||||
|
fmt.Fprintf(os.Stderr, "=== Dry run complete: no output written ===\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Write to target format
|
// Write to target format
|
||||||
fmt.Fprintf(os.Stderr, "[2/2] Writing to target format...\n")
|
fmt.Fprintf(os.Stderr, "[2/2] Writing to target format...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", convertTargetType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", convertTargetType)
|
||||||
@@ -380,6 +418,12 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
|
|||||||
}
|
}
|
||||||
reader = mssql.NewReader(newReaderOptions("", connString))
|
reader = mssql.NewReader(newReaderOptions("", connString))
|
||||||
|
|
||||||
|
case "mysql", "mariadb":
|
||||||
|
if connString == "" {
|
||||||
|
return nil, fmt.Errorf("connection string is required for MySQL format")
|
||||||
|
}
|
||||||
|
reader = mysql.NewReader(newReaderOptions("", connString))
|
||||||
|
|
||||||
case "sqlite", "sqlite3":
|
case "sqlite", "sqlite3":
|
||||||
// SQLite can use either file path or connection string
|
// SQLite can use either file path or connection string
|
||||||
dbPath := filePath
|
dbPath := filePath
|
||||||
@@ -408,23 +452,12 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
|
|
||||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
||||||
if extraFields != "" {
|
if extraFields != "" {
|
||||||
if !strings.EqualFold(dbType, "bun") {
|
extraFieldsJSON, err := loadExtraFields(dbType, extraFields)
|
||||||
return fmt.Errorf("--extra-fields is only supported for Bun output")
|
|
||||||
}
|
|
||||||
extraFieldsJSON, err := os.ReadFile(extraFields)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to read --extra-fields file %q: %w", extraFields, err)
|
return err
|
||||||
}
|
|
||||||
|
|
||||||
var parsed []wbun.ExtraFieldConfig
|
|
||||||
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
|
||||||
return fmt.Errorf("invalid --extra-fields JSON in %q: %w", extraFields, err)
|
|
||||||
}
|
|
||||||
if len(parsed) == 0 {
|
|
||||||
return fmt.Errorf("--extra-fields must contain at least one field")
|
|
||||||
}
|
}
|
||||||
writerOpts.Metadata = map[string]interface{}{
|
writerOpts.Metadata = map[string]interface{}{
|
||||||
"extra_fields": string(extraFieldsJSON),
|
"extra_fields": extraFieldsJSON,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -465,6 +498,9 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
case "mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022":
|
case "mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022":
|
||||||
writer = wmssql.NewWriter(writerOpts)
|
writer = wmssql.NewWriter(writerOpts)
|
||||||
|
|
||||||
|
case "mysql", "mariadb":
|
||||||
|
writer = wmysql.NewWriter(writerOpts)
|
||||||
|
|
||||||
case "sqlite", "sqlite3":
|
case "sqlite", "sqlite3":
|
||||||
writer = wsqlite.NewWriter(writerOpts)
|
writer = wsqlite.NewWriter(writerOpts)
|
||||||
|
|
||||||
@@ -529,6 +565,95 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// loadExtraFields reads and validates the --extra-fields JSON file, returning
|
||||||
|
// its raw content.
|
||||||
|
func loadExtraFields(dbType, path string) (string, error) {
|
||||||
|
if !strings.EqualFold(dbType, "bun") {
|
||||||
|
return "", fmt.Errorf("--extra-fields is only supported for Bun output")
|
||||||
|
}
|
||||||
|
extraFieldsJSON, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to read --extra-fields file %q: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var parsed []wbun.ExtraFieldConfig
|
||||||
|
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
||||||
|
return "", fmt.Errorf("invalid --extra-fields JSON in %q: %w", path, err)
|
||||||
|
}
|
||||||
|
if len(parsed) == 0 {
|
||||||
|
return "", fmt.Errorf("--extra-fields must contain at least one field")
|
||||||
|
}
|
||||||
|
return string(extraFieldsJSON), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateWriteTarget performs the checks writeDatabase makes before writing,
|
||||||
|
// without constructing a writer or touching the output path. Used by --dry-run.
|
||||||
|
func validateWriteTarget(db *models.Database, dbType, packageName, schemaFilter, extraFields string) error {
|
||||||
|
if extraFields != "" {
|
||||||
|
if _, err := loadExtraFields(dbType, extraFields); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(dbType) {
|
||||||
|
case "dbml", "dctx", "drawdb", "json", "yaml", "yml", "drizzle",
|
||||||
|
"pgsql", "postgres", "postgresql", "sql",
|
||||||
|
"mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022",
|
||||||
|
"sqlite", "sqlite3", "prisma", "typeorm", "graphql", "gql":
|
||||||
|
case "gorm", "bun":
|
||||||
|
if packageName == "" {
|
||||||
|
return fmt.Errorf("package name is required for %s format (use --package flag)", dbType)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported target format: %s", dbType)
|
||||||
|
}
|
||||||
|
|
||||||
|
if schemaFilter != "" {
|
||||||
|
for _, schema := range db.Schemas {
|
||||||
|
if schema.Name == schemaFilter {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("schema '%s' not found in database. Available schemas: %v",
|
||||||
|
schemaFilter, getSchemaNames(db))
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.EqualFold(dbType, "dctx") {
|
||||||
|
if len(db.Schemas) == 0 {
|
||||||
|
return fmt.Errorf("no schemas found in database")
|
||||||
|
}
|
||||||
|
if len(db.Schemas) > 1 {
|
||||||
|
return fmt.Errorf("multiple schemas found, please specify which schema to export using --schema flag. Available schemas: %v",
|
||||||
|
getSchemaNames(db))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// outWriter returns the command's stdout, falling back to os.Stdout when the
|
||||||
|
// command is nil (tests call the run functions directly).
|
||||||
|
func outWriter(cmd *cobra.Command) io.Writer {
|
||||||
|
if cmd == nil {
|
||||||
|
return os.Stdout
|
||||||
|
}
|
||||||
|
return cmd.OutOrStdout()
|
||||||
|
}
|
||||||
|
|
||||||
|
// printDryRunPlan prints the schemas and tables that would be written.
|
||||||
|
func printDryRunPlan(out io.Writer, db *models.Database) {
|
||||||
|
for _, schema := range db.Schemas {
|
||||||
|
names := make([]string, 0, len(schema.Tables))
|
||||||
|
for _, t := range schema.Tables {
|
||||||
|
names = append(names, t.Name)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(out, " schema %q: %d table(s)", schema.Name, len(schema.Tables))
|
||||||
|
if len(names) > 0 {
|
||||||
|
fmt.Fprintf(out, " [%s]", strings.Join(names, ", "))
|
||||||
|
}
|
||||||
|
fmt.Fprintln(out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// getSchemaNames returns a slice of schema names from a database
|
// getSchemaNames returns a slice of schema names from a database
|
||||||
func getSchemaNames(db *models.Database) []string {
|
func getSchemaNames(db *models.Database) []string {
|
||||||
names := make([]string, len(db.Schemas))
|
names := make([]string, len(db.Schemas))
|
||||||
|
|||||||
@@ -0,0 +1,312 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
const fixturesDir = "../../tests/assets"
|
||||||
|
|
||||||
|
// readableFormats maps each file-based reader format to an existing fixture.
|
||||||
|
var readableFormats = []struct {
|
||||||
|
format string
|
||||||
|
path string
|
||||||
|
}{
|
||||||
|
{"dbml", "dbml/simple.dbml"},
|
||||||
|
{"json", "json/database.json"},
|
||||||
|
{"yaml", "yaml/database.yaml"},
|
||||||
|
{"yml", "yaml/database.yaml"},
|
||||||
|
{"drawdb", "drawdb/simple.json"},
|
||||||
|
{"dctx", "dctx/p1.dctx"},
|
||||||
|
{"graphql", "graphql/simple.graphql"},
|
||||||
|
{"gql", "graphql/simple.graphql"},
|
||||||
|
{"prisma", "prisma/example.prisma"},
|
||||||
|
{"typeorm", "typeorm/example.ts"},
|
||||||
|
{"drizzle", "drizzle/schema.ts"},
|
||||||
|
{"gorm", "gorm/simple.go"},
|
||||||
|
{"bun", "bun/simple.go"},
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadDatabaseForConvert_FileFormats(t *testing.T) {
|
||||||
|
for _, tt := range readableFormats {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
db, err := readDatabaseForConvert(tt.format, filepath.Join(fixturesDir, tt.path), "")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
if db == nil || len(db.Schemas) == 0 {
|
||||||
|
t.Fatalf("no schemas read: %+v", db)
|
||||||
|
}
|
||||||
|
// Uppercase format names are accepted.
|
||||||
|
if _, err := readDatabaseForConvert(strings.ToUpper(tt.format), filepath.Join(fixturesDir, tt.path), ""); err != nil {
|
||||||
|
t.Errorf("uppercase format: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadDatabaseForConvert_Errors(t *testing.T) {
|
||||||
|
filePathFormats := []string{"dbml", "dctx", "drawdb", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm", "graphql"}
|
||||||
|
for _, f := range filePathFormats {
|
||||||
|
t.Run("missing path "+f, func(t *testing.T) {
|
||||||
|
_, err := readDatabaseForConvert(f, "", "")
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
connFormats := []string{"pgsql", "postgres", "postgresql", "mssql", "sqlserver", "mysql", "mariadb"}
|
||||||
|
for _, f := range connFormats {
|
||||||
|
t.Run("missing conn "+f, func(t *testing.T) {
|
||||||
|
_, err := readDatabaseForConvert(f, "", "")
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "connection string is required") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForConvert("sqlite", "", ""); err == nil || !strings.Contains(err.Error(), "required for SQLite") {
|
||||||
|
t.Errorf("sqlite: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForConvert("nope", "x", ""); err == nil || !strings.Contains(err.Error(), "unsupported source format") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForConvert("dbml", filepath.Join(t.TempDir(), "missing.dbml"), ""); err == nil || !strings.Contains(err.Error(), "failed to read database") {
|
||||||
|
t.Errorf("missing file: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadDatabase_DiffReader(t *testing.T) {
|
||||||
|
for _, f := range []string{"dbml", "json", "yaml", "drawdb", "dctx"} {
|
||||||
|
for _, tt := range readableFormats {
|
||||||
|
if tt.format != f {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
t.Run(f, func(t *testing.T) {
|
||||||
|
db, err := readDatabase(f, filepath.Join(fixturesDir, tt.path), "", "source")
|
||||||
|
if err != nil || db == nil || len(db.Schemas) == 0 {
|
||||||
|
t.Fatalf("read: %v %+v", err, db)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, f := range []string{"dbml", "dctx", "drawdb", "json", "yaml", "sqldir"} {
|
||||||
|
if _, err := readDatabase(f, "", "", "src"); err == nil || !strings.Contains(err.Error(), "src: file path is required") {
|
||||||
|
t.Errorf("%s missing path: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := readDatabase("pgsql", "", "", "src"); err == nil || !strings.Contains(err.Error(), "connection string is required") {
|
||||||
|
t.Errorf("pgsql: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabase("sqlite", "", "", "src"); err == nil {
|
||||||
|
t.Error("sqlite without path must fail")
|
||||||
|
}
|
||||||
|
if _, err := readDatabase("nope", "x", "", "src"); err == nil || !strings.Contains(err.Error(), "unsupported database format") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabase("json", filepath.Join(t.TempDir(), "missing.json"), "", "src"); err == nil || !strings.Contains(err.Error(), "src: failed to read database") {
|
||||||
|
t.Errorf("missing file: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMaskPassword(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
in, want string
|
||||||
|
}{
|
||||||
|
{"", ""},
|
||||||
|
{"postgres://user:secret@host:5432/db", "postgres://user:***@host:5432/db"},
|
||||||
|
{"postgres://user@host:5432/db", "postgres://user@host:5432/db"},
|
||||||
|
{"host=h user=u password=secret dbname=d", "host=h user=u password=*** dbname=d"},
|
||||||
|
{"host=h user=u", "host=h user=u"},
|
||||||
|
{"/tmp/file.db", "/tmp/file.db"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := maskPassword(tt.in); got != tt.want {
|
||||||
|
t.Errorf("maskPassword(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
if got := maskPasswordInDiff(tt.in); got != tt.want {
|
||||||
|
t.Errorf("maskPasswordInDiff(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetSchemaNames(t *testing.T) {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
if got := getSchemaNames(db); len(got) != 0 {
|
||||||
|
t.Errorf("empty: %v", got)
|
||||||
|
}
|
||||||
|
db.Schemas = []*models.Schema{{Name: "a"}, {Name: "b"}}
|
||||||
|
if got := strings.Join(getSchemaNames(db), ","); got != "a,b" {
|
||||||
|
t.Errorf("got %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadExtraFields(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
write := func(name, body string) string {
|
||||||
|
p := filepath.Join(dir, name)
|
||||||
|
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
valid := write("valid.json", `[{"name":"extra"}]`)
|
||||||
|
|
||||||
|
if got, err := loadExtraFields("bun", valid); err != nil || !strings.Contains(got, "extra") {
|
||||||
|
t.Errorf("valid: %q %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := loadExtraFields("BUN", valid); err != nil {
|
||||||
|
t.Errorf("case-insensitive format: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := loadExtraFields("gorm", valid); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
|
||||||
|
t.Errorf("non-bun: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := loadExtraFields("bun", filepath.Join(dir, "missing.json")); err == nil || !strings.Contains(err.Error(), "failed to read") {
|
||||||
|
t.Errorf("missing: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := loadExtraFields("bun", write("bad.json", `{not json`)); err == nil || !strings.Contains(err.Error(), "invalid --extra-fields JSON") {
|
||||||
|
t.Errorf("bad json: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := loadExtraFields("bun", write("empty.json", `[]`)); err == nil || !strings.Contains(err.Error(), "at least one field") {
|
||||||
|
t.Errorf("empty: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func multiSchemaDB() *models.Database {
|
||||||
|
db := models.InitDatabase("multi")
|
||||||
|
for _, n := range []string{"a", "b"} {
|
||||||
|
s := models.InitSchema(n)
|
||||||
|
tbl := models.InitTable("t_"+n, n)
|
||||||
|
c := models.InitColumn("id", tbl.Name, n)
|
||||||
|
c.Type = "integer"
|
||||||
|
c.IsPrimaryKey = true
|
||||||
|
tbl.Columns["id"] = c
|
||||||
|
s.Tables = append(s.Tables, tbl)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateWriteTarget(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
single := models.InitDatabase("single")
|
||||||
|
single.Schemas = []*models.Schema{models.InitSchema("only")}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
db *models.Database
|
||||||
|
dbType, pkg, schemaFilter, extraFields, wantErrSubstr string
|
||||||
|
}{
|
||||||
|
{"json ok", db, "json", "", "", "", ""},
|
||||||
|
{"pgsql alias ok", db, "sql", "", "", "", ""},
|
||||||
|
{"gorm needs package", db, "gorm", "", "", "", "package name is required"},
|
||||||
|
{"bun needs package", db, "bun", "", "", "", "package name is required"},
|
||||||
|
{"gorm with package", db, "gorm", "models", "", "", ""},
|
||||||
|
{"unsupported", db, "nope", "", "", "", "unsupported target format"},
|
||||||
|
{"schema filter found", db, "json", "", "a", "", ""},
|
||||||
|
{"schema filter missing", db, "json", "", "zzz", "", "not found in database"},
|
||||||
|
{"dctx multi schema", db, "dctx", "", "", "", "multiple schemas found"},
|
||||||
|
{"dctx multi schema with filter", db, "dctx", "", "a", "", ""},
|
||||||
|
{"dctx single schema", single, "dctx", "", "", "", ""},
|
||||||
|
{"dctx no schemas", models.InitDatabase("e"), "dctx", "", "", "", "no schemas found"},
|
||||||
|
{"extra fields non-bun", db, "json", "", "", "x.json", "only supported for Bun"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
err := validateWriteTarget(tt.db, tt.dbType, tt.pkg, tt.schemaFilter, tt.extraFields)
|
||||||
|
if tt.wantErrSubstr == "" {
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tt.wantErrSubstr) {
|
||||||
|
t.Errorf("got %v, want substring %q", err, tt.wantErrSubstr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Formats(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
formats := []struct{ format, file string }{
|
||||||
|
{"json", "out.json"},
|
||||||
|
{"yaml", "out.yaml"},
|
||||||
|
{"yml", "out.yml"},
|
||||||
|
{"dbml", "out.dbml"},
|
||||||
|
{"drawdb", "out.drawdb.json"},
|
||||||
|
{"pgsql", "out.sql"},
|
||||||
|
{"postgres", "out2.sql"},
|
||||||
|
{"sql", "out3.sql"},
|
||||||
|
{"mssql", "out_ms.sql"},
|
||||||
|
{"mysql", "out_my.sql"},
|
||||||
|
{"sqlite", "out_lite.sql"},
|
||||||
|
{"graphql", "out.graphql"},
|
||||||
|
{"gql", "out2.graphql"},
|
||||||
|
{"prisma", "out.prisma"},
|
||||||
|
{"typeorm", "out.ts"},
|
||||||
|
{"drizzle", "out_drizzle.ts"},
|
||||||
|
}
|
||||||
|
for _, tt := range formats {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), tt.file)
|
||||||
|
if err := writeDatabase(db, tt.format, out, "", "", false, "", "", false, ""); err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
info, err := os.Stat(out)
|
||||||
|
if err != nil || info.Size() == 0 {
|
||||||
|
t.Errorf("output missing or empty: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_GoFormatsWriteIntoDir(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
for _, f := range []string{"gorm", "bun"} {
|
||||||
|
t.Run(f, func(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "models.go")
|
||||||
|
if err := writeDatabase(db, f, out, "models", "", false, "", "", false, ""); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); err != nil {
|
||||||
|
t.Errorf("no output: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, f, out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "package name is required") {
|
||||||
|
t.Errorf("missing package: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_SchemaFilterAndDCTX(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
out := filepath.Join(t.TempDir(), "o.json")
|
||||||
|
|
||||||
|
if err := writeDatabase(db, "json", out, "", "a", false, "", "", false, ""); err != nil {
|
||||||
|
t.Errorf("schema filter: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, "json", out, "", "zzz", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "not found in database") {
|
||||||
|
t.Errorf("missing schema: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "multiple schemas found") {
|
||||||
|
t.Errorf("dctx multi: %v", err)
|
||||||
|
}
|
||||||
|
single := models.InitDatabase("s")
|
||||||
|
single.Schemas = []*models.Schema{db.Schemas[0]}
|
||||||
|
if err := writeDatabase(single, "dctx", filepath.Join(t.TempDir(), "o.dctx"), "", "", false, "", "", false, ""); err != nil {
|
||||||
|
t.Errorf("dctx single: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(models.InitDatabase("e"), "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "no schemas found") {
|
||||||
|
t.Errorf("dctx empty: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, "nope", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "unsupported target format") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, "json", out, "", "", false, "", "", false, filepath.Join(t.TempDir(), "x.json")); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
|
||||||
|
t.Errorf("extra fields with json: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,158 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newDryRunCmd() (*cobra.Command, *bytes.Buffer) {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
cmd := &cobra.Command{}
|
||||||
|
cmd.SetOut(&buf)
|
||||||
|
return cmd, &buf
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvert_DryRunWritesNothing(t *testing.T) {
|
||||||
|
defer func(a, b, c, d string, e bool) {
|
||||||
|
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||||
|
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "out.json")
|
||||||
|
writeTestJSON(t, in, []string{"users", "posts"})
|
||||||
|
|
||||||
|
convertSourceType, convertSourcePath = "json", in
|
||||||
|
convertTargetType, convertTargetPath = "json", out
|
||||||
|
convertDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runConvert(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||||
|
if !strings.Contains(buf.String(), want) {
|
||||||
|
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Normal behavior is unchanged.
|
||||||
|
convertDryRun = false
|
||||||
|
if err := runConvert(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("real run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); err != nil {
|
||||||
|
t.Fatalf("real run should write output: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvert_DryRunValidatesTarget(t *testing.T) {
|
||||||
|
defer func(a, b, c, d string, e bool) {
|
||||||
|
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||||
|
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "models")
|
||||||
|
writeTestJSON(t, in, []string{"users"})
|
||||||
|
|
||||||
|
convertSourceType, convertSourcePath = "json", in
|
||||||
|
convertDryRun = true
|
||||||
|
|
||||||
|
// gorm without --package must fail validation, as a real run would.
|
||||||
|
convertTargetType, convertTargetPath = "gorm", out
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
err := runConvert(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "package name is required") {
|
||||||
|
t.Fatalf("expected package validation error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
convertTargetType = "nope"
|
||||||
|
err = runConvert(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "unsupported target format") {
|
||||||
|
t.Fatalf("expected unsupported format error, got %v", err)
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(out); !os.IsNotExist(statErr) {
|
||||||
|
t.Fatal("dry run must not create the output path")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunSplit_DryRunWritesNothing(t *testing.T) {
|
||||||
|
defer func(a, b, c, d, e string, f bool) {
|
||||||
|
splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun = a, b, c, d, e, f
|
||||||
|
}(splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "subset.json")
|
||||||
|
writeTestJSON(t, in, []string{"users", "posts", "comments"})
|
||||||
|
|
||||||
|
splitSourceType, splitSourcePath = "json", in
|
||||||
|
splitTargetType, splitTargetPath = "json", out
|
||||||
|
splitTables = "users,posts"
|
||||||
|
splitDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runSplit(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
got := buf.String()
|
||||||
|
if !strings.Contains(got, "2 table(s)") || strings.Contains(got, "comments") {
|
||||||
|
t.Errorf("plan should show only the 2 selected tables:\n%s", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A selection that matches nothing fails validation in dry-run too.
|
||||||
|
splitTables = "does_not_exist"
|
||||||
|
if err := runSplit(cmd, nil); err == nil {
|
||||||
|
t.Fatal("expected error for empty selection")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunMerge_DryRunWritesNothing(t *testing.T) {
|
||||||
|
saved := saveMergeState()
|
||||||
|
defer restoreMergeState(saved)
|
||||||
|
defer func(v bool) { mergeDryRun = v }(mergeDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
target := filepath.Join(dir, "target.json")
|
||||||
|
source := filepath.Join(dir, "source.json")
|
||||||
|
out := filepath.Join(dir, "merged.json")
|
||||||
|
writeTestJSON(t, target, []string{"users"})
|
||||||
|
writeTestJSON(t, source, []string{"posts"})
|
||||||
|
|
||||||
|
mergeTargetType, mergeTargetPath, mergeTargetConn = "json", target, ""
|
||||||
|
mergeSourceType, mergeSourcePath, mergeSourceConn = "json", source, ""
|
||||||
|
mergeFromList = nil
|
||||||
|
mergeOutputType, mergeOutputPath, mergeOutputConn = "json", out, ""
|
||||||
|
mergeSkipTables, mergeReportPath = "", ""
|
||||||
|
mergeDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runMerge(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||||
|
if !strings.Contains(buf.String(), want) {
|
||||||
|
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mergeOutputType = "nope"
|
||||||
|
if err := runMerge(cmd, nil); err == nil || !strings.Contains(err.Error(), "unsupported format") {
|
||||||
|
t.Fatalf("expected unsupported output format error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -61,6 +61,7 @@ var (
|
|||||||
mergeReportPath string // Path to write merge report
|
mergeReportPath string // Path to write merge report
|
||||||
mergeFullDDL bool // Execute full DDL instead of diffing the live pgsql database
|
mergeFullDDL bool // Execute full DDL instead of diffing the live pgsql database
|
||||||
mergeFlattenSchema bool
|
mergeFlattenSchema bool
|
||||||
|
mergeDryRun bool
|
||||||
)
|
)
|
||||||
|
|
||||||
var mergeCmd = &cobra.Command{
|
var mergeCmd = &cobra.Command{
|
||||||
@@ -130,6 +131,7 @@ func init() {
|
|||||||
mergeCmd.Flags().BoolVar(&mergeVerbose, "verbose", false, "Show verbose output")
|
mergeCmd.Flags().BoolVar(&mergeVerbose, "verbose", false, "Show verbose output")
|
||||||
mergeCmd.Flags().BoolVar(&mergeFullDDL, "full-ddl", false, "pgsql database output: execute the full DDL instead of only the differences from the live database")
|
mergeCmd.Flags().BoolVar(&mergeFullDDL, "full-ddl", false, "pgsql database output: execute the full DDL instead of only the differences from the live database")
|
||||||
mergeCmd.Flags().StringVar(&mergeReportPath, "merge-report", "", "Path to write merge report (JSON format)")
|
mergeCmd.Flags().StringVar(&mergeReportPath, "merge-report", "", "Path to write merge report (JSON format)")
|
||||||
|
mergeCmd.Flags().BoolVar(&mergeDryRun, "dry-run", false, "Read and merge in memory, then print the merge plan without writing any output")
|
||||||
mergeCmd.Flags().BoolVar(&mergeFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
mergeCmd.Flags().BoolVar(&mergeFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -264,6 +266,29 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
|||||||
merge.GetColumnTypeConflictSummary(result, 10))
|
merge.GetColumnTypeConflictSummary(result, 10))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if mergeDryRun {
|
||||||
|
if !isMergeOutputFormat(mergeOutputType) {
|
||||||
|
return fmt.Errorf("dry run validation failed: Output: unsupported format '%s'", mergeOutputType)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec merge plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Target: %s database '%s'\n", mergeTargetType, targetDB.Name)
|
||||||
|
fmt.Fprintf(out, " Source: %s database '%s'\n", mergeSourceType, sourceDB.Name)
|
||||||
|
switch {
|
||||||
|
case mergeOutputPath != "":
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", mergeOutputType, mergeOutputPath)
|
||||||
|
case mergeOutputConn != "":
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", mergeOutputType, maskPassword(mergeOutputConn))
|
||||||
|
default:
|
||||||
|
fmt.Fprintf(out, " Output: %s\n", mergeOutputType)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(out, " Result:\n")
|
||||||
|
printDryRunPlan(out, targetDB)
|
||||||
|
fmt.Fprintf(out, "\n%s\n", merge.GetMergeSummary(result))
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Dry run complete: no output written ===\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Step 4: Write output
|
// Step 4: Write output
|
||||||
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
|
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
|
||||||
@@ -285,6 +310,16 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isMergeOutputFormat reports whether writeDatabaseForMerge supports dbType.
|
||||||
|
func isMergeOutputFormat(dbType string) bool {
|
||||||
|
switch strings.ToLower(dbType) {
|
||||||
|
case "dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun",
|
||||||
|
"drizzle", "prisma", "typeorm", "sqlite", "sqlite3", "pgsql":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func readDatabaseForMerge(dbType, filePath, connString, label string) (*models.Database, error) {
|
func readDatabaseForMerge(dbType, filePath, connString, label string) (*models.Database, error) {
|
||||||
var reader readers.Reader
|
var reader readers.Reader
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,287 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestReadDatabaseForMerge(t *testing.T) {
|
||||||
|
for _, tt := range readableFormats {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
db, err := readDatabaseForMerge(tt.format, filepath.Join(fixturesDir, tt.path), "", "Target")
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("format %s not supported by merge reader: %v", tt.format, err)
|
||||||
|
}
|
||||||
|
if db == nil || len(db.Schemas) == 0 {
|
||||||
|
t.Errorf("no schemas: %+v", db)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
for _, f := range []string{"dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm"} {
|
||||||
|
if _, err := readDatabaseForMerge(f, "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src: file path is required") {
|
||||||
|
t.Errorf("%s missing path: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForMerge("pgsql", "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src:") {
|
||||||
|
t.Errorf("pgsql: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForMerge("sqlite", "", "", "Src"); err == nil || !strings.Contains(err.Error(), "Src:") {
|
||||||
|
t.Errorf("sqlite: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForMerge("nope", "x", "", "Src"); err == nil || !strings.Contains(err.Error(), "unsupported format 'nope'") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabaseForMerge(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
single := multiSchemaDB()
|
||||||
|
single.Schemas = single.Schemas[:1]
|
||||||
|
|
||||||
|
files := map[string]string{
|
||||||
|
"dbml": "o.dbml", "dctx": "o.dctx", "drawdb": "o.drawdb.json", "graphql": "o.graphql",
|
||||||
|
"json": "o.json", "yaml": "o.yaml", "gorm": "gorm.go", "bun": "bun.go",
|
||||||
|
"drizzle": "o.ts", "prisma": "o.prisma", "typeorm": "te.ts",
|
||||||
|
}
|
||||||
|
for f, name := range files {
|
||||||
|
t.Run(f, func(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), name)
|
||||||
|
if f == "dctx" {
|
||||||
|
// DCTX cannot write a full database.
|
||||||
|
if err := writeDatabaseForMerge(f, out, "", single, "Output", false); err == nil || !strings.Contains(err.Error(), "not supported for DCTX") {
|
||||||
|
t.Errorf("dctx: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabaseForMerge(f, "", "", single, "Output", false); err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||||
|
t.Errorf("dctx missing path: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
src := db
|
||||||
|
if err := writeDatabaseForMerge(f, out, "", src, "Output", false); err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); err != nil {
|
||||||
|
t.Errorf("no output: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabaseForMerge(f, "", "", src, "Output", false); err == nil || !strings.Contains(err.Error(), "Output: file path is required") {
|
||||||
|
t.Errorf("missing path: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
for _, f := range []string{"pgsql", "sqlite"} {
|
||||||
|
out := filepath.Join(t.TempDir(), "o.sql")
|
||||||
|
if err := writeDatabaseForMerge(f, out, "", db, "Output", false); err != nil {
|
||||||
|
t.Errorf("%s script write: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := writeDatabaseForMerge("pgsql", "", "postgres://u:p@127.0.0.1:1/none?connect_timeout=1", db, "Output", false); err == nil {
|
||||||
|
t.Error("pgsql with unreachable conn must fail")
|
||||||
|
}
|
||||||
|
if err := writeDatabaseForMerge("nope", "x", "", db, "Output", false); err == nil || !strings.Contains(err.Error(), "unsupported") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsMergeOutputFormat(t *testing.T) {
|
||||||
|
for _, f := range []string{"dbml", "JSON", "pgsql", "sqlite3", "prisma"} {
|
||||||
|
if !isMergeOutputFormat(f) {
|
||||||
|
t.Errorf("%s should be supported", f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, f := range []string{"", "nope", "mssql"} {
|
||||||
|
if isMergeOutputFormat(f) {
|
||||||
|
t.Errorf("%s should not be supported", f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExpandPath(t *testing.T) {
|
||||||
|
home, err := os.UserHomeDir()
|
||||||
|
if err != nil {
|
||||||
|
t.Skip("no home dir")
|
||||||
|
}
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"/abs/path", "/abs/path"},
|
||||||
|
{"rel/path", "rel/path"},
|
||||||
|
{"~/x/y", filepath.Join(home, "/x/y")},
|
||||||
|
{"~", home},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := expandPath(tt.in); got != tt.want {
|
||||||
|
t.Errorf("expandPath(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseSkipTables(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
in string
|
||||||
|
want []string
|
||||||
|
}{
|
||||||
|
{"", nil},
|
||||||
|
{" , ,", nil},
|
||||||
|
{"Users", []string{"users"}},
|
||||||
|
{" Users , ORDERS,,items ", []string{"users", "orders", "items"}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
got := parseSkipTables(tt.in)
|
||||||
|
if len(got) != len(tt.want) {
|
||||||
|
t.Errorf("parseSkipTables(%q) = %v", tt.in, got)
|
||||||
|
}
|
||||||
|
for _, w := range tt.want {
|
||||||
|
if !got[w] {
|
||||||
|
t.Errorf("parseSkipTables(%q) missing %q", tt.in, w)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadDatabaseForInspect(t *testing.T) {
|
||||||
|
for _, tt := range readableFormats {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
db, err := readDatabaseForInspect(tt.format, filepath.Join(fixturesDir, tt.path), "")
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("format %s not supported by inspect reader: %v", tt.format, err)
|
||||||
|
}
|
||||||
|
if db == nil || len(db.Schemas) == 0 {
|
||||||
|
t.Errorf("no schemas: %+v", db)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
for _, f := range []string{"dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm"} {
|
||||||
|
if _, err := readDatabaseForInspect(f, "", ""); err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||||
|
t.Errorf("%s missing path: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForInspect("pgsql", "", ""); err == nil {
|
||||||
|
t.Error("pgsql without conn must fail")
|
||||||
|
}
|
||||||
|
if _, err := readDatabaseForInspect("nope", "x", ""); err == nil || !strings.Contains(err.Error(), "unsupported database type") {
|
||||||
|
t.Errorf("unsupported: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterDatabaseBySchema(t *testing.T) {
|
||||||
|
db := multiSchemaDB()
|
||||||
|
db.Description = "desc"
|
||||||
|
got := filterDatabaseBySchema(db, "b")
|
||||||
|
if len(got.Schemas) != 1 || got.Schemas[0].Name != "b" || got.Name != db.Name || got.Description != "desc" {
|
||||||
|
t.Errorf("filtered: %+v", got)
|
||||||
|
}
|
||||||
|
if got := filterDatabaseBySchema(db, "zzz"); len(got.Schemas) != 0 {
|
||||||
|
t.Errorf("missing schema should yield no schemas: %+v", got.Schemas)
|
||||||
|
}
|
||||||
|
if len(db.Schemas) != 2 {
|
||||||
|
t.Error("input mutated")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHasSilentFlag(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
args []string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{nil, false},
|
||||||
|
{[]string{"convert"}, false},
|
||||||
|
{[]string{"convert", "--silent"}, true},
|
||||||
|
{[]string{"--silent=true"}, true},
|
||||||
|
{[]string{"--silent=false"}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := hasSilentFlag(tt.args); got != tt.want {
|
||||||
|
t.Errorf("hasSilentFlag(%v) = %v", tt.args, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrintVersionHeader(t *testing.T) {
|
||||||
|
capture := func(args []string) string {
|
||||||
|
old := os.Stdout
|
||||||
|
r, w, _ := os.Pipe()
|
||||||
|
os.Stdout = w
|
||||||
|
printVersionHeader(args)
|
||||||
|
w.Close()
|
||||||
|
os.Stdout = old
|
||||||
|
b := make([]byte, 4096)
|
||||||
|
n, _ := r.Read(b)
|
||||||
|
return string(b[:n])
|
||||||
|
}
|
||||||
|
if out := capture([]string{"convert"}); !strings.HasPrefix(out, "RelSpec ") {
|
||||||
|
t.Errorf("header: %q", out)
|
||||||
|
}
|
||||||
|
if out := capture([]string{"convert", "--no-version"}); out != "" {
|
||||||
|
t.Errorf("--no-version: %q", out)
|
||||||
|
}
|
||||||
|
if out := capture([]string{"version"}); out != "" {
|
||||||
|
t.Errorf("version cmd: %q", out)
|
||||||
|
}
|
||||||
|
if out := capture(nil); !strings.HasPrefix(out, "RelSpec ") {
|
||||||
|
t.Errorf("no args: %q", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReportState(t *testing.T) {
|
||||||
|
cfg := t.TempDir()
|
||||||
|
t.Setenv("XDG_CONFIG_HOME", cfg)
|
||||||
|
t.Setenv("HOME", cfg)
|
||||||
|
|
||||||
|
dir, err := reportStateDir()
|
||||||
|
if err != nil || !strings.HasPrefix(dir, cfg) {
|
||||||
|
t.Fatalf("dir: %q %v", dir, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
state, path, err := loadReportState()
|
||||||
|
if err != nil || !state.LastReport.IsZero() || state.MachineID != "" {
|
||||||
|
t.Fatalf("fresh state: %+v %v", state, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
want := reportState{LastReport: time.Now().UTC().Truncate(time.Second), MachineID: "abc"}
|
||||||
|
if err := saveReportState(path, want); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got, _, err := loadReportState()
|
||||||
|
if err != nil || !got.LastReport.Equal(want.LastReport) || got.MachineID != "abc" {
|
||||||
|
t.Errorf("round trip: %+v %v", got, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Corrupt state is ignored.
|
||||||
|
if err := os.WriteFile(path, []byte("{bad"), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got, _, err := loadReportState(); err != nil || got.MachineID != "" {
|
||||||
|
t.Errorf("corrupt: %+v %v", got, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSystemUniqueID_NonEmpty(t *testing.T) {
|
||||||
|
cfg := t.TempDir()
|
||||||
|
t.Setenv("XDG_CONFIG_HOME", cfg)
|
||||||
|
state, path, _ := loadReportState()
|
||||||
|
id, err := systemUniqueID(state, path)
|
||||||
|
if err != nil || id == "" {
|
||||||
|
t.Errorf("id: %q %v", id, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReportToken_Decodes(t *testing.T) {
|
||||||
|
if _, err := reportToken(); err != nil {
|
||||||
|
t.Errorf("token must decode: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSubmitReport_RateLimited(t *testing.T) {
|
||||||
|
cfg := t.TempDir()
|
||||||
|
t.Setenv("XDG_CONFIG_HOME", cfg)
|
||||||
|
_, path, _ := loadReportState()
|
||||||
|
if err := saveReportState(path, reportState{LastReport: time.Now()}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// Rate limit rejects before any network call is made.
|
||||||
|
if err := submitReport("bug", "t", "b", "", ""); err == nil || !strings.Contains(err.Error(), "please wait") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -27,6 +27,7 @@ func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullab
|
|||||||
FlattenSchema: flattenSchema,
|
FlattenSchema: flattenSchema,
|
||||||
NullableTypes: nullableTypes,
|
NullableTypes: nullableTypes,
|
||||||
NullableArrays: nullableArrays,
|
NullableArrays: nullableArrays,
|
||||||
|
TypeMappings: typeMappings,
|
||||||
Prisma7: prisma7,
|
Prisma7: prisma7,
|
||||||
ContinueOnError: continueOnError,
|
ContinueOnError: continueOnError,
|
||||||
StrictDirectives: strictDirectives,
|
StrictDirectives: strictDirectives,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
)
|
)
|
||||||
|
|
||||||
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
||||||
@@ -17,6 +18,8 @@ var (
|
|||||||
noVersion bool
|
noVersion bool
|
||||||
silent bool
|
silent bool
|
||||||
strictDirectives bool
|
strictDirectives bool
|
||||||
|
typeMapFlags []string
|
||||||
|
typeMappings map[string]string
|
||||||
)
|
)
|
||||||
|
|
||||||
var rootCmd = &cobra.Command{
|
var rootCmd = &cobra.Command{
|
||||||
@@ -28,10 +31,16 @@ bidirectional conversion between various database schema formats.
|
|||||||
It reads database schemas from multiple sources (live databases, DBML,
|
It reads database schemas from multiple sources (live databases, DBML,
|
||||||
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
||||||
JSON, YAML, SQL, etc.).`,
|
JSON, YAML, SQL, etc.).`,
|
||||||
|
PersistentPreRunE: func(cmd *cobra.Command, args []string) error {
|
||||||
|
var err error
|
||||||
|
typeMappings, err = writers.ParseTypeMappings(typeMapFlags)
|
||||||
|
return err
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
rootCmd.AddCommand(convertCmd)
|
rootCmd.AddCommand(convertCmd)
|
||||||
|
rootCmd.AddCommand(batchCmd)
|
||||||
rootCmd.AddCommand(diffCmd)
|
rootCmd.AddCommand(diffCmd)
|
||||||
rootCmd.AddCommand(inspectCmd)
|
rootCmd.AddCommand(inspectCmd)
|
||||||
rootCmd.AddCommand(scriptsCmd)
|
rootCmd.AddCommand(scriptsCmd)
|
||||||
@@ -44,6 +53,7 @@ func init() {
|
|||||||
rootCmd.AddCommand(versionCmd)
|
rootCmd.AddCommand(versionCmd)
|
||||||
rootCmd.AddCommand(reportCmd)
|
rootCmd.AddCommand(reportCmd)
|
||||||
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||||
|
rootCmd.PersistentFlags().StringArrayVar(&typeMapFlags, "type-map", nil, "Override a SQL-to-Go type mapping for bun/gorm output as sqltype=gotype (repeatable), e.g. --type-map uuid=uuid.UUID --type-map numeric=decimal.Decimal")
|
||||||
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
||||||
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
||||||
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRunDiff(t *testing.T) {
|
||||||
|
oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO := sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath
|
||||||
|
t.Cleanup(func() {
|
||||||
|
sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath = oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO
|
||||||
|
})
|
||||||
|
|
||||||
|
src := filepath.Join(fixturesDir, "dbml/simple.dbml")
|
||||||
|
cmplx := filepath.Join(fixturesDir, "dbml/complex.dbml")
|
||||||
|
|
||||||
|
for _, format := range []string{"summary", "json", "html"} {
|
||||||
|
t.Run(format, func(t *testing.T) {
|
||||||
|
sourceType, sourcePath, sourceConn = "dbml", src, ""
|
||||||
|
targetType, targetPath, targetConn = "dbml", cmplx, ""
|
||||||
|
outputFormat = format
|
||||||
|
outputPath = filepath.Join(t.TempDir(), "diff.out")
|
||||||
|
if format == "summary" {
|
||||||
|
outputPath = ""
|
||||||
|
}
|
||||||
|
if err := runDiff(nil, nil); err != nil {
|
||||||
|
t.Fatalf("runDiff: %v", err)
|
||||||
|
}
|
||||||
|
if outputPath != "" {
|
||||||
|
if b, err := os.ReadFile(outputPath); err != nil || len(b) == 0 {
|
||||||
|
t.Errorf("empty output: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("bad source", func(t *testing.T) {
|
||||||
|
sourceType, sourcePath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
|
||||||
|
targetType, targetPath = "dbml", src
|
||||||
|
outputFormat, outputPath = "summary", ""
|
||||||
|
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source database") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("bad target", func(t *testing.T) {
|
||||||
|
sourceType, sourcePath = "dbml", src
|
||||||
|
targetType, targetPath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
|
||||||
|
outputFormat, outputPath = "summary", ""
|
||||||
|
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read target database") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunInspect(t *testing.T) {
|
||||||
|
oldT, oldP, oldC, oldR, oldF, oldO, oldS := inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter
|
||||||
|
t.Cleanup(func() {
|
||||||
|
inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter = oldT, oldP, oldC, oldR, oldF, oldO, oldS
|
||||||
|
})
|
||||||
|
|
||||||
|
inspectSourceType = "dbml"
|
||||||
|
inspectSourcePath = filepath.Join(fixturesDir, "dbml/simple.dbml")
|
||||||
|
inspectSourceConn = ""
|
||||||
|
inspectRulesPath = filepath.Join(t.TempDir(), "no-rules.yaml") // missing: defaults used or error
|
||||||
|
inspectSchemaFilter = ""
|
||||||
|
|
||||||
|
// Whatever the rules outcome, the run must not panic; formats are exercised.
|
||||||
|
for _, format := range []string{"markdown", "json"} {
|
||||||
|
inspectOutputFormat = format
|
||||||
|
inspectOutputPath = filepath.Join(t.TempDir(), "report."+format)
|
||||||
|
_ = runInspect(nil, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
inspectOutputFormat = "bogus"
|
||||||
|
inspectOutputPath = ""
|
||||||
|
if err := runInspect(nil, nil); err == nil {
|
||||||
|
t.Error("bogus output format must fail")
|
||||||
|
}
|
||||||
|
|
||||||
|
inspectSourcePath = filepath.Join(t.TempDir(), "missing.dbml")
|
||||||
|
if err := runInspect(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source") {
|
||||||
|
t.Errorf("missing source: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -24,6 +24,7 @@ var (
|
|||||||
splitExcludeTables string
|
splitExcludeTables string
|
||||||
splitNullableTypes string
|
splitNullableTypes string
|
||||||
splitNullableArrays string
|
splitNullableArrays string
|
||||||
|
splitDryRun bool
|
||||||
)
|
)
|
||||||
|
|
||||||
var splitCmd = &cobra.Command{
|
var splitCmd = &cobra.Command{
|
||||||
@@ -115,6 +116,8 @@ func init() {
|
|||||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||||
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||||
|
|
||||||
|
splitCmd.Flags().BoolVar(&splitDryRun, "dry-run", false, "Read, filter and validate the selection and print the plan without writing any output")
|
||||||
|
|
||||||
err := splitCmd.MarkFlagRequired("from")
|
err := splitCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
||||||
@@ -174,6 +177,26 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
fmt.Fprintf(os.Stderr, " ✓ Filtered to: %d schema(s), %d table(s)\n\n", len(filteredDB.Schemas), filteredTables)
|
fmt.Fprintf(os.Stderr, " ✓ Filtered to: %d schema(s), %d table(s)\n\n", len(filteredDB.Schemas), filteredTables)
|
||||||
|
|
||||||
|
if splitDryRun {
|
||||||
|
if err := validateWriteTarget(filteredDB, splitTargetType, splitPackageName, "", ""); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec split plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Input: %s database '%s'\n", splitSourceType, db.Name)
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", splitTargetType, splitTargetPath)
|
||||||
|
fmt.Fprintf(out, " Selection: %s\n", splitSelection{
|
||||||
|
Schemas: parseCommaSeparated(splitSchemas),
|
||||||
|
Tables: parseCommaSeparated(splitTables),
|
||||||
|
ExcludeSchemas: parseCommaSeparated(splitExcludeSchema),
|
||||||
|
ExcludeTables: parseCommaSeparated(splitExcludeTables),
|
||||||
|
DatabaseName: splitDatabaseName,
|
||||||
|
}.summary())
|
||||||
|
printDryRunPlan(out, filteredDB)
|
||||||
|
fmt.Fprintf(os.Stderr, "=== Dry run complete: no output written ===\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Write to target format
|
// Write to target format
|
||||||
fmt.Fprintf(os.Stderr, "[3/3] Writing to target format...\n")
|
fmt.Fprintf(os.Stderr, "[3/3] Writing to target format...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", splitTargetType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", splitTargetType)
|
||||||
|
|||||||
@@ -0,0 +1,152 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchSnapshot maps a file path to its modification time and size.
|
||||||
|
type watchSnapshot map[string]string
|
||||||
|
|
||||||
|
// takeWatchSnapshot records the state of every file under the given paths.
|
||||||
|
// Directories are walked recursively. Anything at or below the excluded path
|
||||||
|
// (typically the output path) is skipped so regenerating output does not
|
||||||
|
// retrigger the watcher. Missing paths are simply absent from the snapshot, so
|
||||||
|
// creating them later counts as a change.
|
||||||
|
func takeWatchSnapshot(paths []string, exclude string) watchSnapshot {
|
||||||
|
snap := watchSnapshot{}
|
||||||
|
exclude = absPathOrSelf(exclude)
|
||||||
|
for _, root := range paths {
|
||||||
|
_ = filepath.WalkDir(root, func(p string, d fs.DirEntry, err error) error {
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if exclude != "" && isWithin(absPathOrSelf(p), exclude) {
|
||||||
|
if d.IsDir() {
|
||||||
|
return filepath.SkipDir
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if d.IsDir() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, err := d.Info()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
snap[p] = fmt.Sprintf("%d-%d", info.ModTime().UnixNano(), info.Size())
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return snap
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s watchSnapshot) equal(o watchSnapshot) bool {
|
||||||
|
if len(s) != len(o) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for k, v := range s {
|
||||||
|
if o[k] != v {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func absPathOrSelf(p string) string {
|
||||||
|
if p == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if abs, err := filepath.Abs(p); err == nil {
|
||||||
|
return abs
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// isWithin reports whether path equals dir or is located below it.
|
||||||
|
func isWithin(path, dir string) bool {
|
||||||
|
if path == dir {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return strings.HasPrefix(path, dir+string(filepath.Separator))
|
||||||
|
}
|
||||||
|
|
||||||
|
// watchLoop runs fn once immediately and again whenever the watched paths
|
||||||
|
// change, until ctx is cancelled. Changes are debounced: fn runs only after
|
||||||
|
// the snapshot has stayed unchanged for one poll interval. Errors from fn are
|
||||||
|
// reported to w and do not stop the loop.
|
||||||
|
func watchLoop(ctx context.Context, w io.Writer, paths []string, exclude string, interval time.Duration, fn func() error) {
|
||||||
|
run := func() {
|
||||||
|
if err := fn(); err != nil {
|
||||||
|
fmt.Fprintf(w, "Error: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, "Watching for changes (Ctrl-C to stop)...\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
last := takeWatchSnapshot(paths, exclude)
|
||||||
|
run()
|
||||||
|
|
||||||
|
ticker := time.NewTicker(interval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
|
cur := takeWatchSnapshot(paths, exclude)
|
||||||
|
if cur.equal(last) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Debounce: wait until writes settle.
|
||||||
|
for settled := false; !settled; {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-time.After(interval):
|
||||||
|
}
|
||||||
|
next := takeWatchSnapshot(paths, exclude)
|
||||||
|
settled = next.equal(cur)
|
||||||
|
cur = next
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, "\nChange detected, regenerating...\n")
|
||||||
|
last = cur
|
||||||
|
run()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// runConvertWatch runs the conversion once and then again whenever the source
|
||||||
|
// files change, until interrupted.
|
||||||
|
func runConvertWatch(parent context.Context, w io.Writer, run func() error) error {
|
||||||
|
var paths []string
|
||||||
|
switch {
|
||||||
|
case len(convertFromList) > 0:
|
||||||
|
paths = convertFromList
|
||||||
|
case convertSourcePath != "":
|
||||||
|
paths = []string{convertSourcePath}
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("--watch requires --from-path or --from-list (live database connections cannot be watched)")
|
||||||
|
}
|
||||||
|
if convertDryRun {
|
||||||
|
return fmt.Errorf("--watch cannot be combined with --dry-run")
|
||||||
|
}
|
||||||
|
if convertWatchInterval <= 0 {
|
||||||
|
return fmt.Errorf("--watch-interval must be positive")
|
||||||
|
}
|
||||||
|
if parent == nil {
|
||||||
|
parent = context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, stop := signal.NotifyContext(parent, os.Interrupt, syscall.SIGTERM)
|
||||||
|
defer stop()
|
||||||
|
watchLoop(ctx, w, paths, convertTargetPath, convertWatchInterval, run)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,92 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWatchSnapshotExcludesOutput(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
out := filepath.Join(dir, "out")
|
||||||
|
if err := os.MkdirAll(out, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
src := filepath.Join(dir, "schema.dbml")
|
||||||
|
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
before := takeWatchSnapshot([]string{dir}, out)
|
||||||
|
if err := os.WriteFile(filepath.Join(out, "gen.go"), []byte("x"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if after := takeWatchSnapshot([]string{dir}, out); !before.equal(after) {
|
||||||
|
t.Errorf("writing into the excluded output path changed the snapshot")
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(src, []byte("changed"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if after := takeWatchSnapshot([]string{dir}, out); before.equal(after) {
|
||||||
|
t.Errorf("modifying a source file did not change the snapshot")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchLoopRerunsOnChange(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
src := filepath.Join(dir, "schema.dbml")
|
||||||
|
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var runs atomic.Int32
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
watchLoop(ctx, io.Discard, []string{src}, "", 10*time.Millisecond, func() error {
|
||||||
|
runs.Add(1)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
|
||||||
|
waitFor(t, func() bool { return runs.Load() == 1 })
|
||||||
|
if err := os.WriteFile(src, []byte("changed content"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
waitFor(t, func() bool { return runs.Load() == 2 })
|
||||||
|
cancel()
|
||||||
|
<-done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvertWatchValidation(t *testing.T) {
|
||||||
|
oldPath, oldList, oldDry, oldInt := convertSourcePath, convertFromList, convertDryRun, convertWatchInterval
|
||||||
|
defer func() {
|
||||||
|
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = oldPath, oldList, oldDry, oldInt
|
||||||
|
}()
|
||||||
|
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = "", nil, false, time.Second
|
||||||
|
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||||
|
t.Error("expected error without --from-path/--from-list")
|
||||||
|
}
|
||||||
|
convertSourcePath, convertDryRun = "x.dbml", true
|
||||||
|
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||||
|
t.Error("expected error combining --watch with --dry-run")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitFor(t *testing.T, cond func() bool) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(5 * time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if cond() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatal("condition not met in time")
|
||||||
|
}
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
# Format Usage Examples
|
||||||
|
|
||||||
|
Examples for `relspec convert` covering the file-based reader and writer
|
||||||
|
formats. The "Writers" and "Readers" sections below were run against
|
||||||
|
`examples/test_schema.dbml`. The cross-format and live-database examples were
|
||||||
|
not run; they follow the flags shown in `relspec convert --help` and require
|
||||||
|
matching input files or reachable databases.
|
||||||
|
|
||||||
|
Any reader can be combined with any writer: pick `--from`/`--from-path` for the
|
||||||
|
source and `--to`/`--to-path` for the target. Add `--silent` to suppress progress
|
||||||
|
output.
|
||||||
|
|
||||||
|
## Writers: DBML to every format
|
||||||
|
|
||||||
|
```bash
|
||||||
|
S="--from dbml --from-path examples/test_schema.dbml"
|
||||||
|
|
||||||
|
relspec convert $S --to json --to-path schema.json
|
||||||
|
relspec convert $S --to yaml --to-path schema.yaml
|
||||||
|
relspec convert $S --to dctx --to-path schema.dctx
|
||||||
|
relspec convert $S --to drawdb --to-path schema.drawdb.json
|
||||||
|
relspec convert $S --to graphql --to-path schema.graphql
|
||||||
|
relspec convert $S --to prisma --to-path schema.prisma
|
||||||
|
relspec convert $S --to pgsql --to-path schema.pg.sql
|
||||||
|
relspec convert $S --to mssql --to-path schema.mssql.sql
|
||||||
|
relspec convert $S --to sqlite --to-path schema.sqlite.sql
|
||||||
|
relspec convert $S --to drizzle --to-path schema.ts
|
||||||
|
relspec convert $S --to typeorm --to-path entities.ts
|
||||||
|
relspec convert $S --to gorm --to-path models.go --package models
|
||||||
|
relspec convert $S --to bun --to-path models.go --package models
|
||||||
|
```
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
|
||||||
|
- Code-generation writers (`gorm`, `bun`) take `--package`. They also accept
|
||||||
|
`--types baselib|stdlib|sqltypes` to choose the nullable type package.
|
||||||
|
- When `--to-path` is a directory it must already exist.
|
||||||
|
- `sqlite` output automatically flattens `schema.table` names. Use
|
||||||
|
`--flatten-schema` for other formats if the target has no schema support.
|
||||||
|
- `dctx` supports a single schema only; use `--schema <name>` to select one.
|
||||||
|
|
||||||
|
## Readers: file-based formats into DBML (or JSON where noted)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec convert --from json --from-path schema.json --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from yaml --from-path schema.yaml --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from dctx --from-path schema.dctx --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from drawdb --from-path schema.drawdb.json --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from graphql --from-path schema.graphql --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from prisma --from-path schema.prisma --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from drizzle --from-path schema.ts --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from typeorm --from-path entities.ts --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from bun --from-path models.go --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from gorm --from-path models.go --to json --to-path out.json
|
||||||
|
```
|
||||||
|
|
||||||
|
Code-first readers (`gorm`, `bun`, `drizzle`, `typeorm`) accept a single file or a
|
||||||
|
directory of model files.
|
||||||
|
|
||||||
|
> Known issue: reading GORM models and writing DBML currently panics in the DBML
|
||||||
|
> writer (`pkg/writers/dbml/writer.go`, `constraintToDBML`). Use another target
|
||||||
|
> such as JSON until this is fixed.
|
||||||
|
|
||||||
|
## Cross-format combinations
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# ORM models to SQL DDL
|
||||||
|
relspec convert --from gorm --from-path models.go --to pgsql --to-path schema.sql
|
||||||
|
|
||||||
|
# Prisma to Drizzle
|
||||||
|
relspec convert --from prisma --from-path schema.prisma --to drizzle --to-path schema.ts
|
||||||
|
|
||||||
|
# DrawDB diagram to GraphQL
|
||||||
|
relspec convert --from drawdb --from-path diagram.json --to graphql --to-path schema.graphql
|
||||||
|
|
||||||
|
# Merge several files while converting
|
||||||
|
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
||||||
|
```
|
||||||
|
|
||||||
|
## Live databases
|
||||||
|
|
||||||
|
These need a reachable database:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# PostgreSQL
|
||||||
|
relspec convert --from pgsql --from-conn "postgres://user:pass@localhost:5432/mydb" \
|
||||||
|
--to dbml --to-path schema.dbml
|
||||||
|
|
||||||
|
# SQL Server
|
||||||
|
relspec convert --from mssql --from-conn "<mssql connection string>" \
|
||||||
|
--to json --to-path schema.json
|
||||||
|
|
||||||
|
# SQLite database file (--from-conn takes the file path)
|
||||||
|
relspec convert --from sqlite --from-conn ./app.db --to dbml --to-path schema.dbml
|
||||||
|
```
|
||||||
|
|
||||||
|
## Formats outside `convert`
|
||||||
|
|
||||||
|
- `sqldir` (SQL script directory reader) and `sqlexec` (SQL execution writer) are
|
||||||
|
used by `relspec scripts` and `relspec job`, and `sqldir` by `relspec diff`.
|
||||||
|
See [SCRIPTS_COMMAND.md](SCRIPTS_COMMAND.md) and [JOB_FILES.md](JOB_FILES.md).
|
||||||
|
- The `template` writer is exposed through `relspec templ`. See
|
||||||
|
[TEMPLATE_MODE.md](TEMPLATE_MODE.md).
|
||||||
@@ -0,0 +1,212 @@
|
|||||||
|
# TUI mouse support plan
|
||||||
|
|
||||||
|
Issue: #46
|
||||||
|
Status: design only; this document does not implement mouse input.
|
||||||
|
|
||||||
|
## 1. Current implementation and scope
|
||||||
|
|
||||||
|
The editor is created in `cmd/relspec/edit.go` by
|
||||||
|
`ui.NewSchemaEditorWithConfigs(...).Run()`. `pkg/ui/editor.go` owns the
|
||||||
|
`tview.Application`, `tview.Pages`, and application lifecycle. The current
|
||||||
|
code never calls `Application.EnableMouse`, so tcell mouse reporting is off.
|
||||||
|
The module uses tview v0.42.0 and tcell/v2 v2.13.9.
|
||||||
|
|
||||||
|
The first implementation should add `--no-mouse` to the `edit` Cobra command
|
||||||
|
only. The flag is a local boolean, defaulting to false, and should be passed
|
||||||
|
explicitly into the editor (prefer an options/config field rather than a
|
||||||
|
package-global or environment variable). It must not affect convert, inspect,
|
||||||
|
merge, or other commands. There is no environment-variable or persistent
|
||||||
|
configuration setting in this issue: a command-line opt-out is predictable,
|
||||||
|
visible in `edit --help`, and avoids adding configuration precedence rules.
|
||||||
|
|
||||||
|
At startup, the editor should call `app.EnableMouse(!noMouse)` before
|
||||||
|
`Run()`. `--no-mouse` must mean that the application does not enable terminal
|
||||||
|
mouse reporting and that no custom mouse handlers are relied upon. Keyboard
|
||||||
|
behavior must remain identical in both modes.
|
||||||
|
|
||||||
|
Likely implementation files are `cmd/relspec/edit.go`,
|
||||||
|
`pkg/ui/editor.go`, focused TUI mouse helpers/tests under `pkg/ui`, and a
|
||||||
|
short user-facing note in the command help or TUI documentation. Do not
|
||||||
|
refactor unrelated screens or data operations.
|
||||||
|
|
||||||
|
## 2. Widget and screen coverage
|
||||||
|
|
||||||
|
The application composes `Pages`, `Flex`, `TextView`, `List`, `Table`, `Form`,
|
||||||
|
`Button`, `InputField`, `DropDown`, `TextArea`, `CheckBox`, and `Modal`.
|
||||||
|
Vendored tview confirms mouse handlers exist for all of those relevant
|
||||||
|
primitives, including focus on left-down, list/table selection, button clicks,
|
||||||
|
form child dispatch, dropdown opening/drag selection, text-area cursor and
|
||||||
|
scrolling, and modal button dispatch. `Pages`, `Flex`, and `Form` forward events
|
||||||
|
to their children.
|
||||||
|
|
||||||
|
The coverage plan is:
|
||||||
|
|
||||||
|
* Main menu (`pkg/ui/main_menu.go`): left click focuses/selects a list entry;
|
||||||
|
second activation opens it; buttons and exit confirmation remain reachable.
|
||||||
|
* Schema, table, domain, object, relation, and database screens: click a row
|
||||||
|
to select it; double-click the row to perform the same action as the
|
||||||
|
keyboard Enter/selected callback where opening is meaningful; scroll lists
|
||||||
|
and tables; click each action button.
|
||||||
|
* Tables (`schema_screens.go`, `table_screens.go`, and object/relation tables):
|
||||||
|
tview's table handler provides selection and scrolling, but it does not
|
||||||
|
provide application-specific double-click activation. Add a small reusable
|
||||||
|
wrapper/helper for the table instances that need it. It must preserve the
|
||||||
|
existing selected row/column behavior and invoke the same callback as Enter,
|
||||||
|
not duplicate mutation logic.
|
||||||
|
* Forms (`load_save_screens.go` and the form-building screen files): click an
|
||||||
|
input to focus it, click buttons to activate them, click a dropdown to open
|
||||||
|
it and choose an option, scroll multiline help/text areas, and retain all
|
||||||
|
existing keyboard Tab/Shift-Tab, shortcut, Enter, and Escape behavior.
|
||||||
|
* Dialogs (`pkg/ui/dialogs.go` plus confirmation/error/success modals): modal
|
||||||
|
buttons are clickable and the modal keeps focus above the underlying page.
|
||||||
|
Clicking outside a modal must not activate the hidden page or dismiss a
|
||||||
|
destructive confirmation. Escape and the existing button-key behavior stay
|
||||||
|
authoritative.
|
||||||
|
* The planned file browser and connection-string builder from issue #44 must
|
||||||
|
use the same contracts: clickable entries/buttons and scrolling, with
|
||||||
|
keyboard navigation and explicit cancel/accept paths. #46 should not
|
||||||
|
implement #44's widgets; it should define the integration point and test
|
||||||
|
them when #44 lands.
|
||||||
|
|
||||||
|
Do not promise drag semantics for every widget. Drag is appropriate for text
|
||||||
|
selection/cursor movement and dropdown selection where tview already supports
|
||||||
|
it. For ordinary list/table navigation, a click selects and the wheel scrolls;
|
||||||
|
row dragging should not mutate data.
|
||||||
|
|
||||||
|
## 3. Exact mouse action contract
|
||||||
|
|
||||||
|
| Widget/type | Left down/click | Double click | Wheel/drag | Keyboard fallback |
|
||||||
|
| --- | --- | --- | --- | --- |
|
||||||
|
| Main/list menu | focus and select row | invoke row selected callback | scroll list | arrows, Enter, shortcuts |
|
||||||
|
| Data table | focus and select cell/row | invoke the screen's existing open/edit action for the selected row | vertical/horizontal scroll as supported by tview | arrows, PageUp/PageDown, Enter, existing shortcuts |
|
||||||
|
| Button | focus | same as one activation, never duplicate the callback | none | Tab/Shift-Tab, Enter/Space and existing shortcut |
|
||||||
|
| Input field | focus; place cursor if supported | no destructive action | text-area behavior if provided by tview | typing, arrows, Home/End, Tab, Escape |
|
||||||
|
| Text area/help | focus and position cursor | select word only where tview supports it; no application action | scroll; drag text selection if supported | arrows, PageUp/PageDown, standard editing keys |
|
||||||
|
| Dropdown | focus/open and choose the hit option | same as click; no duplicate selection | drag through options only while open | arrows, Enter, Escape, Tab |
|
||||||
|
| Checkbox | toggle on click | no second toggle | none | Space and existing form navigation |
|
||||||
|
| Modal | focus/click visible button | same button action once | no underlying-page scrolling | Tab/Shift-Tab, Enter, Escape, existing button keys |
|
||||||
|
| Blank/border/title area | focus containing primitive where useful | none | no mutation | current screen shortcuts |
|
||||||
|
|
||||||
|
Right and middle clicks should have no application action in the first
|
||||||
|
release. Wheel events should be consumed only by the scrollable primitive
|
||||||
|
under the pointer. Double-click timing/translation should come from tview/
|
||||||
|
tcell; custom code must not fire the action once for both the click and the
|
||||||
|
double-click. Any custom table wrapper needs a small state machine or tview
|
||||||
|
mouse action handling that is tested for this property.
|
||||||
|
|
||||||
|
## 4. tview gaps and implementation boundaries
|
||||||
|
|
||||||
|
Enabling mouse support is not sufficient for the desired behavior. tview's
|
||||||
|
built-in Table handler selects cells and scrolls but has no repository-specific
|
||||||
|
row-open callback on double click. Existing screen code also wires keyboard
|
||||||
|
input captures directly on individual widgets, so mouse actions must call the
|
||||||
|
same screen callbacks rather than route through synthetic key events.
|
||||||
|
|
||||||
|
Use tview's `MouseHandler`/`WrapMouseHandler` contracts and `setFocus` rather
|
||||||
|
than reading terminal coordinates in each screen. A reusable table adapter
|
||||||
|
may embed `*tview.Table`, delegate ordinary actions to the original table
|
||||||
|
handler, and add the screen's double-click callback. Keep the adapter in
|
||||||
|
`pkg/ui` and use it only where a row-opening action exists. Do not modify the
|
||||||
|
vendored tview copy.
|
||||||
|
|
||||||
|
The `Pages`/`Modal` dispatch order must be verified: a visible modal consumes
|
||||||
|
its click before the page below it. Page transitions should happen only in the
|
||||||
|
existing callbacks, so a stale hidden page cannot receive a click.
|
||||||
|
|
||||||
|
## 5. Keyboard, terminal, and copy/paste behavior
|
||||||
|
|
||||||
|
Mouse is an enhancement, never a requirement. Every acceptance path must be
|
||||||
|
reachable with the existing keyboard controls, including load/save, navigation,
|
||||||
|
editing, confirmations, cancel, and exit. `--no-mouse` is the regression mode
|
||||||
|
for proving this contract.
|
||||||
|
|
||||||
|
Mouse reporting is terminal capability dependent. On local terminals it is
|
||||||
|
negotiated by tcell; tmux and SSH can suppress, translate, or fail to pass
|
||||||
|
mouse reporting depending on their configuration. The application must still
|
||||||
|
start and remain keyboard usable if mouse reporting is unavailable or broken.
|
||||||
|
Documentation should state that terminal/tmux configuration may be required,
|
||||||
|
and that SSH behavior depends on the remote terminal path. Windows Terminal and
|
||||||
|
other Windows console hosts should be treated as supported only insofar as the
|
||||||
|
selected tcell backend reports mouse events; the CLI must not assume POSIX
|
||||||
|
escape sequences or add platform-specific code in this issue.
|
||||||
|
|
||||||
|
Enabling mouse capture normally prevents terminal-native selection/copy from
|
||||||
|
seeing ordinary button-drag events. Document the standard workaround: hold the
|
||||||
|
terminal's bypass modifier (commonly Shift, terminal-dependent) for selection,
|
||||||
|
or use `--no-mouse` when native copy/paste is the priority. Do not implement a
|
||||||
|
second clipboard protocol. Input-field/text-area copy/paste must continue to
|
||||||
|
use tview/tcell paste handling and keyboard shortcuts; verify that enabling
|
||||||
|
mouse does not intercept paste events.
|
||||||
|
|
||||||
|
## 6. Test strategy using tcell simulation
|
||||||
|
|
||||||
|
Add focused tests rather than attempting a full interactive end-to-end test.
|
||||||
|
Use `tcell.NewSimulationScreen("")`, `screen.Init()`, construct the editor or
|
||||||
|
an isolated primitive tree, and inject events with the actual API:
|
||||||
|
`SimulationScreen.InjectMouse(x, y, buttons, mod)` and `InjectKey(...)`.
|
||||||
|
Coordinates must be derived from the primitive's drawn rectangle or fixed by a
|
||||||
|
small deterministic test layout; do not use arbitrary coordinates without
|
||||||
|
checking the rendered screen.
|
||||||
|
|
||||||
|
Minimum cases:
|
||||||
|
|
||||||
|
1. Default editor configuration enables mouse; the explicit disabled option
|
||||||
|
leaves it disabled. If the Application API is not observable directly,
|
||||||
|
test through the simulation screen's event path plus a constructor-level
|
||||||
|
option assertion.
|
||||||
|
2. A list click changes focus/selection, and double-click invokes the existing
|
||||||
|
selected action exactly once.
|
||||||
|
3. A table click selects the expected row/cell, wheel events change the visible
|
||||||
|
offset, and double-click invokes the row action exactly once.
|
||||||
|
4. Form button, input field, checkbox, and dropdown clicks match their
|
||||||
|
keyboard callbacks.
|
||||||
|
5. A modal button click acts on the modal and cannot activate the underlying
|
||||||
|
page; Escape still cancels.
|
||||||
|
6. `--no-mouse` leaves keyboard selection/activation unchanged and mouse
|
||||||
|
injection has no application effect.
|
||||||
|
7. Existing dialogs and screen transitions do not leave a stale mouse capture
|
||||||
|
after a page is removed.
|
||||||
|
|
||||||
|
Prefer callback counters and selected-index assertions over screen-text-only
|
||||||
|
assertions. Run the relevant `pkg/ui` tests with `go test -race ./pkg/ui` and
|
||||||
|
run the full repository test suite if time/resources permit.
|
||||||
|
|
||||||
|
## 7. Rollout and acceptance criteria
|
||||||
|
|
||||||
|
Implementation is ready for review when:
|
||||||
|
|
||||||
|
* `relspec edit --help` documents `--no-mouse` and mouse is enabled by default.
|
||||||
|
* Only the edit TUI is affected; non-TUI commands have no changed behavior.
|
||||||
|
* Main screens, tables, lists, forms, dropdowns, buttons, text areas, and
|
||||||
|
visible dialogs support the action contract above.
|
||||||
|
* Keyboard-only operation is complete and verified with `--no-mouse`.
|
||||||
|
* Modal clicks cannot fall through to an underlying page.
|
||||||
|
* Table double-click behavior is explicit, tested, and does not duplicate
|
||||||
|
activation.
|
||||||
|
* tcell simulation tests cover default-on, opt-out, selection, scrolling,
|
||||||
|
activation, dialog focus, and keyboard fallback.
|
||||||
|
* `go test -race ./pkg/ui`, appropriate command tests, `go test ./...`,
|
||||||
|
formatting, and `git diff --check` pass (or any limitation is recorded).
|
||||||
|
* User-facing docs explain tmux/SSH/Windows variability and the terminal
|
||||||
|
modifier workaround for native copy/paste.
|
||||||
|
|
||||||
|
Roll out in two implementation slices if needed: first application option,
|
||||||
|
standard tview handlers, tests, and documentation; second only the reusable
|
||||||
|
table double-click adapter and screen wiring. Do not block the first slice on
|
||||||
|
issue #44, but do not claim #44's future widgets are covered until they use the
|
||||||
|
same contract.
|
||||||
|
|
||||||
|
## 8. Open decisions and dependencies
|
||||||
|
|
||||||
|
* Confirm whether the project wants a public editor options type or a small
|
||||||
|
`SetMouseEnabled`/constructor parameter; avoid a global flag.
|
||||||
|
* Confirm the preferred terminal copy modifier in project documentation, since
|
||||||
|
tmux, SSH clients, and Windows Terminal differ.
|
||||||
|
* Decide whether horizontal wheel events should be supported where tview/table
|
||||||
|
exposes them; vertical scrolling is mandatory, horizontal is optional.
|
||||||
|
* Decide whether double-click opens every data table or only tables with an
|
||||||
|
unambiguous row action. The plan recommends the latter.
|
||||||
|
* Confirm #44's file-browser and connection-builder primitive choices before
|
||||||
|
wiring their mouse tests.
|
||||||
|
* Confirm CI has a stable non-terminal environment for simulation-screen tests;
|
||||||
|
no real terminal, tmux session, database, or network should be required.
|
||||||
@@ -4,6 +4,7 @@ go 1.25.13
|
|||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/gdamore/tcell/v2 v2.13.9
|
github.com/gdamore/tcell/v2 v2.13.9
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/jackc/pgx/v5 v5.9.2
|
github.com/jackc/pgx/v5 v5.9.2
|
||||||
github.com/microsoft/go-mssqldb v1.10.0
|
github.com/microsoft/go-mssqldb v1.10.0
|
||||||
@@ -18,6 +19,7 @@ require (
|
|||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
filippo.io/edwards25519 v1.1.0 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/gdamore/encoding v1.0.1 // indirect
|
github.com/gdamore/encoding v1.0.1 // indirect
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
|
||||||
|
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
||||||
@@ -21,6 +23,8 @@ github.com/gdamore/encoding v1.0.1 h1:YzKZckdBL6jVt2Gc+5p82qhrGiqMdG/eNs6Wy0u3Uh
|
|||||||
github.com/gdamore/encoding v1.0.1/go.mod h1:0Z0cMFinngz9kS1QfMjCP8TY7em3bZYeeklsSDPivEo=
|
github.com/gdamore/encoding v1.0.1/go.mod h1:0Z0cMFinngz9kS1QfMjCP8TY7em3bZYeeklsSDPivEo=
|
||||||
github.com/gdamore/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
|
github.com/gdamore/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
|
||||||
github.com/gdamore/tcell/v2 v2.13.9/go.mod h1:+Wfe208WDdB7INEtCsNrAN6O2m+wsTPk1RAovjaILlo=
|
github.com/gdamore/tcell/v2 v2.13.9/go.mod h1:+Wfe208WDdB7INEtCsNrAN6O2m+wsTPk1RAovjaILlo=
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo=
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3/go.mod h1:qn46aNg1333BRMNU69Lq93t8du/dwxI64Gl8i5p1WMU=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||||
|
|||||||
@@ -71,6 +71,13 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
|
|||||||
return diff
|
return diff
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *SchemaChange) addChange(field string, source, target any) {
|
||||||
|
if c.Changes == nil {
|
||||||
|
c.Changes = make(map[string]any)
|
||||||
|
}
|
||||||
|
c.Changes[field] = map[string]any{"source": source, "target": target}
|
||||||
|
}
|
||||||
|
|
||||||
func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
||||||
change := &SchemaChange{
|
change := &SchemaChange{
|
||||||
Name: source.Name,
|
Name: source.Name,
|
||||||
@@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
|
|||||||
|
|
||||||
hasChanges := false
|
hasChanges := false
|
||||||
|
|
||||||
|
// Compare schema attributes
|
||||||
|
if source.Description != target.Description {
|
||||||
|
change.addChange("description", source.Description, target.Description)
|
||||||
|
hasChanges = true
|
||||||
|
}
|
||||||
|
if source.Owner != target.Owner {
|
||||||
|
change.addChange("owner", source.Owner, target.Owner)
|
||||||
|
hasChanges = true
|
||||||
|
}
|
||||||
|
|
||||||
// Compare tables
|
// Compare tables
|
||||||
tableDiff := compareTables(source.Tables, target.Tables)
|
tableDiff := compareTables(source.Tables, target.Tables)
|
||||||
if !isEmpty(tableDiff) {
|
if !isEmpty(tableDiff) {
|
||||||
|
|||||||
@@ -0,0 +1,337 @@
|
|||||||
|
package diff
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCompareSchemaDetails(t *testing.T) {
|
||||||
|
mk := func() *models.Schema {
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
s.Tables = []*models.Table{models.InitTable("t", "public")}
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := compareSchemaDetails(mk(), mk()); got != nil {
|
||||||
|
t.Errorf("identical schemas must yield nil, got %+v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
mutate func(*models.Schema)
|
||||||
|
check func(*SchemaChange) bool
|
||||||
|
}{
|
||||||
|
{"table added", func(s *models.Schema) { s.Tables = append(s.Tables, models.InitTable("u", "public")) },
|
||||||
|
func(c *SchemaChange) bool { return c.Tables != nil && len(c.Tables.Extra) == 1 }},
|
||||||
|
{"view added", func(s *models.Schema) { s.Views = []*models.View{models.InitView("v", "public")} },
|
||||||
|
func(c *SchemaChange) bool { return c.Views != nil && len(c.Views.Extra) == 1 }},
|
||||||
|
{"sequence added", func(s *models.Schema) { s.Sequences = []*models.Sequence{models.InitSequence("sq", "public")} },
|
||||||
|
func(c *SchemaChange) bool { return c.Sequences != nil && len(c.Sequences.Extra) == 1 }},
|
||||||
|
{"script added", func(s *models.Schema) { s.Scripts = []*models.Script{models.InitScript("sc")} },
|
||||||
|
func(c *SchemaChange) bool { return c.Scripts != nil && len(c.Scripts.Extra) == 1 }},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
target := mk()
|
||||||
|
tt.mutate(target)
|
||||||
|
got := compareSchemaDetails(mk(), target)
|
||||||
|
if got == nil || got.Name != "public" || !tt.check(got) {
|
||||||
|
t.Errorf("unexpected change: %+v", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareConstraintDetails(t *testing.T) {
|
||||||
|
base := func() *models.Constraint {
|
||||||
|
c := models.InitConstraint("fk", models.ForeignKeyConstraint)
|
||||||
|
c.Columns = []string{"a"}
|
||||||
|
c.ReferencedTable = "users"
|
||||||
|
c.ReferencedColumns = []string{"id"}
|
||||||
|
c.OnDelete = "CASCADE"
|
||||||
|
c.OnUpdate = "NO ACTION"
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
if got := compareConstraintDetails(base(), base()); len(got) != 0 {
|
||||||
|
t.Errorf("identical: %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
mutate func(*models.Constraint)
|
||||||
|
wantKey string
|
||||||
|
}{
|
||||||
|
{"type", func(c *models.Constraint) { c.Type = models.UniqueConstraint }, "type"},
|
||||||
|
{"columns", func(c *models.Constraint) { c.Columns = []string{"b"} }, "columns"},
|
||||||
|
{"referenced table", func(c *models.Constraint) { c.ReferencedTable = "other" }, "referenced_table"},
|
||||||
|
{"referenced columns", func(c *models.Constraint) { c.ReferencedColumns = []string{"x"} }, "referenced_columns"},
|
||||||
|
{"on delete", func(c *models.Constraint) { c.OnDelete = "SET NULL" }, "on_delete"},
|
||||||
|
{"on update", func(c *models.Constraint) { c.OnUpdate = "CASCADE" }, "on_update"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
target := base()
|
||||||
|
tt.mutate(target)
|
||||||
|
got := compareConstraintDetails(base(), target)
|
||||||
|
if _, ok := got[tt.wantKey]; !ok || len(got) != 1 {
|
||||||
|
t.Errorf("got %v, want only %q", got, tt.wantKey)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Action spelling variants that mean the same thing are not changes.
|
||||||
|
a, b := base(), base()
|
||||||
|
a.OnDelete, b.OnDelete = "cascade", " CASCADE "
|
||||||
|
a.OnUpdate, b.OnUpdate = "", "no action"
|
||||||
|
if got := compareConstraintDetails(a, b); len(got) != 0 {
|
||||||
|
t.Errorf("equivalent actions reported as changes: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeConstraintAction(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"NO ACTION", ""},
|
||||||
|
{"no action", ""},
|
||||||
|
{" No Action ", ""},
|
||||||
|
{"cascade", "CASCADE"},
|
||||||
|
{" set null ", "SET NULL"},
|
||||||
|
{"RESTRICT", "RESTRICT"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := normalizeConstraintAction(tt.in); got != tt.want {
|
||||||
|
t.Errorf("normalizeConstraintAction(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConstraintCompareKey(t *testing.T) {
|
||||||
|
uq := &models.Constraint{Name: "UQ_Name", Type: models.UniqueConstraint}
|
||||||
|
if got := constraintCompareKey(uq); got != "uq_name" {
|
||||||
|
t.Errorf("non-FK key: %q", got)
|
||||||
|
}
|
||||||
|
fk := func(name string) *models.Constraint {
|
||||||
|
return &models.Constraint{
|
||||||
|
Name: name, Type: models.ForeignKeyConstraint, Schema: "Public", Table: "Orders",
|
||||||
|
Columns: []string{"user_id"}, ReferencedSchema: "Public", ReferencedTable: "Users", ReferencedColumns: []string{"id"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if constraintCompareKey(fk("a")) != constraintCompareKey(fk("b")) {
|
||||||
|
t.Error("FK key must ignore the constraint name")
|
||||||
|
}
|
||||||
|
other := fk("a")
|
||||||
|
other.ReferencedColumns = []string{"uid"}
|
||||||
|
if constraintCompareKey(fk("a")) == constraintCompareKey(other) {
|
||||||
|
t.Error("FK key must include referenced columns")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterPrimaryKeyConstraints(t *testing.T) {
|
||||||
|
in := map[string]*models.Constraint{
|
||||||
|
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
|
||||||
|
"uq": {Name: "uq", Type: models.UniqueConstraint},
|
||||||
|
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
|
||||||
|
}
|
||||||
|
got := filterPrimaryKeyConstraints(in)
|
||||||
|
if len(got) != 2 || got["pk"] != nil || got["uq"] == nil || got["fk"] == nil {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
if len(in) != 3 {
|
||||||
|
t.Error("input must not be modified")
|
||||||
|
}
|
||||||
|
if got := filterPrimaryKeyConstraints(nil); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("nil: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareRelationshipDetails(t *testing.T) {
|
||||||
|
base := func() *models.Relationship {
|
||||||
|
r := models.InitRelationship("r", models.RelationType("one_to_many"))
|
||||||
|
r.FromTable, r.ToTable = "orders", "users"
|
||||||
|
r.FromColumns, r.ToColumns = []string{"user_id"}, []string{"id"}
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
if got := compareRelationshipDetails(base(), base()); len(got) != 0 {
|
||||||
|
t.Errorf("identical: %v", got)
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
mutate func(*models.Relationship)
|
||||||
|
wantKey string
|
||||||
|
}{
|
||||||
|
{"type", func(r *models.Relationship) { r.Type = "one_to_one" }, "type"},
|
||||||
|
{"from table", func(r *models.Relationship) { r.FromTable = "x" }, "from_table"},
|
||||||
|
{"to table", func(r *models.Relationship) { r.ToTable = "x" }, "to_table"},
|
||||||
|
{"from columns", func(r *models.Relationship) { r.FromColumns = []string{"x"} }, "from_columns"},
|
||||||
|
{"to columns", func(r *models.Relationship) { r.ToColumns = []string{"x"} }, "to_columns"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
target := base()
|
||||||
|
tt.mutate(target)
|
||||||
|
got := compareRelationshipDetails(base(), target)
|
||||||
|
if _, ok := got[tt.wantKey]; !ok || len(got) != 1 {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareRelationshipsModified(t *testing.T) {
|
||||||
|
src := map[string]*models.Relationship{
|
||||||
|
"same": {Name: "same", Type: "one_to_many"},
|
||||||
|
"changed": {Name: "changed", Type: "one_to_many"},
|
||||||
|
"missing": {Name: "missing"},
|
||||||
|
}
|
||||||
|
tgt := map[string]*models.Relationship{
|
||||||
|
"same": {Name: "same", Type: "one_to_many"},
|
||||||
|
"changed": {Name: "changed", Type: "many_to_many"},
|
||||||
|
"extra": {Name: "extra"},
|
||||||
|
}
|
||||||
|
d := compareRelationships(src, tgt)
|
||||||
|
if len(d.Missing) != 1 || d.Missing[0].Name != "missing" || len(d.Extra) != 1 || d.Extra[0].Name != "extra" ||
|
||||||
|
len(d.Modified) != 1 || d.Modified[0].Name != "changed" {
|
||||||
|
t.Errorf("got %+v", d)
|
||||||
|
}
|
||||||
|
if _, ok := d.Modified[0].Changes["type"]; !ok {
|
||||||
|
t.Errorf("changes: %v", d.Modified[0].Changes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareViews(t *testing.T) {
|
||||||
|
v := func(name, def string) *models.View { return &models.View{Name: name, Definition: def} }
|
||||||
|
src := []*models.View{v("Keep", "select 1"), v("Changed", "select 1"), v("Gone", "select 1")}
|
||||||
|
tgt := []*models.View{v("keep", "select 1"), v("changed", "select 2"), v("New", "select 1")}
|
||||||
|
|
||||||
|
d := compareViews(src, tgt)
|
||||||
|
if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" {
|
||||||
|
t.Errorf("missing: %+v", d.Missing)
|
||||||
|
}
|
||||||
|
if len(d.Extra) != 1 || d.Extra[0].Name != "New" {
|
||||||
|
t.Errorf("extra: %+v", d.Extra)
|
||||||
|
}
|
||||||
|
if len(d.Modified) != 1 || d.Modified[0].Name != "changed" || d.Modified[0].Source.Definition != "select 1" || d.Modified[0].Target.Definition != "select 2" {
|
||||||
|
t.Errorf("modified: %+v", d.Modified)
|
||||||
|
}
|
||||||
|
want := map[string]any{"definition": map[string]string{"source": "select 1", "target": "select 2"}}
|
||||||
|
if !reflect.DeepEqual(d.Modified[0].Changes, want) {
|
||||||
|
t.Errorf("changes: %v", d.Modified[0].Changes)
|
||||||
|
}
|
||||||
|
if !isEmpty(compareViews(nil, nil)) {
|
||||||
|
t.Error("nil views must be empty")
|
||||||
|
}
|
||||||
|
if got := compareViewDetails(v("a", "x"), v("a", "x")); len(got) != 0 {
|
||||||
|
t.Errorf("same definition: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareSequences(t *testing.T) {
|
||||||
|
seq := func(name string, start, inc, min, max int64, cycle bool) *models.Sequence {
|
||||||
|
return &models.Sequence{Name: name, StartValue: start, IncrementBy: inc, MinValue: min, MaxValue: max, Cycle: cycle}
|
||||||
|
}
|
||||||
|
src := []*models.Sequence{seq("Same", 1, 1, 1, 100, false), seq("Diff", 1, 1, 1, 100, false), seq("Gone", 1, 1, 1, 1, false)}
|
||||||
|
tgt := []*models.Sequence{seq("same", 1, 1, 1, 100, false), seq("diff", 5, 2, 3, 200, true), seq("New", 1, 1, 1, 1, false)}
|
||||||
|
|
||||||
|
d := compareSequences(src, tgt)
|
||||||
|
if len(d.Missing) != 1 || d.Missing[0].Name != "Gone" || len(d.Extra) != 1 || d.Extra[0].Name != "New" || len(d.Modified) != 1 {
|
||||||
|
t.Fatalf("got %+v", d)
|
||||||
|
}
|
||||||
|
ch := d.Modified[0].Changes
|
||||||
|
for _, key := range []string{"start_value", "increment_by", "min_value", "max_value", "cycle"} {
|
||||||
|
if _, ok := ch[key]; !ok {
|
||||||
|
t.Errorf("missing change key %q in %v", key, ch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := ch["increment_by"].(map[string]int64); got["source"] != 1 || got["target"] != 2 {
|
||||||
|
t.Errorf("increment_by: %v", got)
|
||||||
|
}
|
||||||
|
if got := ch["cycle"].(map[string]bool); got["source"] || !got["target"] {
|
||||||
|
t.Errorf("cycle: %v", got)
|
||||||
|
}
|
||||||
|
if got := compareSequenceDetails(seq("a", 1, 1, 1, 1, false), seq("a", 1, 1, 1, 1, false)); len(got) != 0 {
|
||||||
|
t.Errorf("identical: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareScriptDetailsAllFields(t *testing.T) {
|
||||||
|
a := &models.Script{Name: "s", SQL: "a", Rollback: "ra", RunAfter: []string{"x"}, Schema: "p", Version: "1", Priority: 1, Sequence: 1}
|
||||||
|
b := &models.Script{Name: "s", SQL: "b", Rollback: "rb", RunAfter: []string{"y"}, Schema: "q", Version: "2", Priority: 2, Sequence: 2}
|
||||||
|
got := compareScriptDetails(a, b)
|
||||||
|
for _, key := range []string{"sql", "rollback", "run_after", "schema", "version", "priority", "sequence"} {
|
||||||
|
if _, ok := got[key]; !ok {
|
||||||
|
t.Errorf("missing %q in %v", key, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := compareScriptDetails(a, a); len(got) != 0 {
|
||||||
|
t.Errorf("identical: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsEmptyAllTypes(t *testing.T) {
|
||||||
|
if !isEmpty(&ViewDiff{}) || !isEmpty(&SequenceDiff{}) {
|
||||||
|
t.Error("empty view/sequence diffs must be empty")
|
||||||
|
}
|
||||||
|
if isEmpty(&ViewDiff{Extra: []*models.View{{Name: "v"}}}) || isEmpty(&SequenceDiff{Modified: []*SequenceChange{{Name: "s"}}}) {
|
||||||
|
t.Error("non-empty diffs reported as empty")
|
||||||
|
}
|
||||||
|
if isEmpty(&ConstraintDiff{Modified: []*ConstraintChange{{Name: "c"}}}) || isEmpty(&RelationshipDiff{Missing: []*models.Relationship{{Name: "r"}}}) {
|
||||||
|
t.Error("non-empty diffs reported as empty")
|
||||||
|
}
|
||||||
|
if isEmpty(&IndexDiff{Modified: []*IndexChange{{Name: "i"}}}) || isEmpty(&TableDiff{Modified: []*TableChange{{Name: "t"}}}) {
|
||||||
|
t.Error("non-empty diffs reported as empty")
|
||||||
|
}
|
||||||
|
if isEmpty("something else") || isEmpty(nil) {
|
||||||
|
t.Error("unknown types must not be treated as empty")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComputeSummaryFullTree(t *testing.T) {
|
||||||
|
res := &DiffResult{Schemas: &SchemaDiff{
|
||||||
|
Missing: []*models.Schema{{Name: "m"}},
|
||||||
|
Extra: []*models.Schema{{Name: "e"}},
|
||||||
|
Modified: []*SchemaChange{{
|
||||||
|
Name: "public",
|
||||||
|
Tables: &TableDiff{
|
||||||
|
Missing: []*models.Table{{Name: "a"}},
|
||||||
|
Extra: []*models.Table{{Name: "b"}, {Name: "c"}},
|
||||||
|
Modified: []*TableChange{{
|
||||||
|
Name: "t",
|
||||||
|
Columns: &ColumnDiff{Missing: []*models.Column{{}}, Extra: []*models.Column{{}, {}}, Modified: []*ColumnChange{{}}},
|
||||||
|
Indexes: &IndexDiff{Missing: []*models.Index{{}}, Extra: []*models.Index{{}}, Modified: []*IndexChange{{}, {}}},
|
||||||
|
Constraints: &ConstraintDiff{Missing: []*models.Constraint{{}}, Modified: []*ConstraintChange{{}}},
|
||||||
|
Relationships: &RelationshipDiff{Extra: []*models.Relationship{{}}},
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
Views: &ViewDiff{Missing: []*models.View{{}}, Extra: []*models.View{{}}, Modified: []*ViewChange{{}}},
|
||||||
|
Sequences: &SequenceDiff{Missing: []*models.Sequence{{}}, Extra: []*models.Sequence{{}, {}}},
|
||||||
|
Scripts: &ScriptDiff{Modified: []*ScriptChange{{}}},
|
||||||
|
}},
|
||||||
|
}}
|
||||||
|
s := ComputeSummary(res)
|
||||||
|
checks := []struct {
|
||||||
|
name string
|
||||||
|
got [3]int
|
||||||
|
want [3]int
|
||||||
|
}{
|
||||||
|
{"schemas", [3]int{s.Schemas.Missing, s.Schemas.Extra, s.Schemas.Modified}, [3]int{1, 1, 1}},
|
||||||
|
{"tables", [3]int{s.Tables.Missing, s.Tables.Extra, s.Tables.Modified}, [3]int{1, 2, 1}},
|
||||||
|
{"columns", [3]int{s.Columns.Missing, s.Columns.Extra, s.Columns.Modified}, [3]int{1, 2, 1}},
|
||||||
|
{"indexes", [3]int{s.Indexes.Missing, s.Indexes.Extra, s.Indexes.Modified}, [3]int{1, 1, 2}},
|
||||||
|
{"constraints", [3]int{s.Constraints.Missing, s.Constraints.Extra, s.Constraints.Modified}, [3]int{1, 0, 1}},
|
||||||
|
{"relationships", [3]int{s.Relationships.Missing, s.Relationships.Extra, s.Relationships.Modified}, [3]int{0, 1, 0}},
|
||||||
|
{"views", [3]int{s.Views.Missing, s.Views.Extra, s.Views.Modified}, [3]int{1, 1, 1}},
|
||||||
|
{"sequences", [3]int{s.Sequences.Missing, s.Sequences.Extra, s.Sequences.Modified}, [3]int{1, 2, 0}},
|
||||||
|
{"scripts", [3]int{s.Scripts.Missing, s.Scripts.Extra, s.Scripts.Modified}, [3]int{0, 0, 1}},
|
||||||
|
}
|
||||||
|
for _, c := range checks {
|
||||||
|
if c.got != c.want {
|
||||||
|
t.Errorf("%s: got %v, want %v", c.name, c.got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := ComputeSummary(&DiffResult{}); got == nil || got.Schemas != (SchemaSummary{}) {
|
||||||
|
t.Errorf("nil Schemas: %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package diff
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCompareSchemaDetails_DescriptionAndOwner(t *testing.T) {
|
||||||
|
mk := func(desc, owner string) *models.Schema {
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
s.Description, s.Owner = desc, owner
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
src, tgt *models.Schema
|
||||||
|
wantFields []string
|
||||||
|
}{
|
||||||
|
{"identical", mk("d", "o"), mk("d", "o"), nil},
|
||||||
|
{"description", mk("a", "o"), mk("b", "o"), []string{"description"}},
|
||||||
|
{"owner", mk("d", "x"), mk("d", "y"), []string{"owner"}},
|
||||||
|
{"both", mk("a", "x"), mk("b", "y"), []string{"description", "owner"}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := compareSchemaDetails(tt.src, tt.tgt)
|
||||||
|
if len(tt.wantFields) == 0 {
|
||||||
|
if got != nil {
|
||||||
|
t.Fatalf("expected no change, got %+v", got)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if got == nil || len(got.Changes) != len(tt.wantFields) {
|
||||||
|
t.Fatalf("changes: %+v", got)
|
||||||
|
}
|
||||||
|
for _, f := range tt.wantFields {
|
||||||
|
c, ok := got.Changes[f].(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("missing %s: %+v", f, got.Changes)
|
||||||
|
}
|
||||||
|
if c["source"] == c["target"] {
|
||||||
|
t.Errorf("%s source and target equal: %v", f, c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareDatabases_SchemaAttrsCounted(t *testing.T) {
|
||||||
|
src, tgt := models.InitDatabase("a"), models.InitDatabase("b")
|
||||||
|
s1, s2 := models.InitSchema("public"), models.InitSchema("public")
|
||||||
|
s1.Owner, s2.Owner = "alice", "bob"
|
||||||
|
src.Schemas, tgt.Schemas = append(src.Schemas, s1), append(tgt.Schemas, s2)
|
||||||
|
|
||||||
|
res := CompareDatabases(src, tgt)
|
||||||
|
if res.Schemas == nil || len(res.Schemas.Modified) != 1 {
|
||||||
|
t.Fatalf("schema owner change not reported: %+v", res.Schemas)
|
||||||
|
}
|
||||||
|
if ComputeSummary(res).Schemas.Modified != 1 {
|
||||||
|
t.Error("summary must count the modified schema")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -19,6 +19,7 @@ type SchemaDiff struct {
|
|||||||
// SchemaChange represents changes within a schema
|
// SchemaChange represents changes within a schema
|
||||||
type SchemaChange struct {
|
type SchemaChange struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
|
Changes map[string]any `json:"changes,omitempty"` // Schema attributes that differ (description, owner), keyed by field name
|
||||||
Tables *TableDiff `json:"tables,omitempty"`
|
Tables *TableDiff `json:"tables,omitempty"`
|
||||||
Views *ViewDiff `json:"views,omitempty"`
|
Views *ViewDiff `json:"views,omitempty"`
|
||||||
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
Sequences *SequenceDiff `json:"sequences,omitempty"`
|
||||||
|
|||||||
@@ -0,0 +1,262 @@
|
|||||||
|
package jobs
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// validateYAML loads one job file and returns the Validate error text ("" when valid).
|
||||||
|
func validateYAML(t *testing.T, body string) string {
|
||||||
|
t.Helper()
|
||||||
|
set := loadOne(t, "version: 1\njobs:\n"+body)
|
||||||
|
if err := set.Validate(); err != nil {
|
||||||
|
return err.Error()
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateJobTable(t *testing.T) {
|
||||||
|
in := " inputs:\n - path: a.dbml\n format: dbml\n"
|
||||||
|
out := " output:\n format: json\n path: out.json\n"
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
job string
|
||||||
|
want string // substring of the error, "" for valid
|
||||||
|
}{
|
||||||
|
{"missing command", " x:\n description: d\n", "missing command"},
|
||||||
|
{"convert valid", " x:\n command: convert\n" + in + out, ""},
|
||||||
|
{"convert script dirs", " x:\n command: convert\n script_dirs: [s]\n" + in + out, "script_dirs is not valid"},
|
||||||
|
{"convert missing output", " x:\n command: convert\n" + in, "missing output"},
|
||||||
|
{"convert output missing format", " x:\n command: convert\n" + in + " output:\n path: o\n", "output: missing format"},
|
||||||
|
{"convert output unsupported format", " x:\n command: convert\n" + in + " output:\n format: nope\n path: o\n", "unsupported output format"},
|
||||||
|
{"convert output missing path", " x:\n command: convert\n" + in + " output:\n format: json\n", "output: missing path"},
|
||||||
|
{"convert output conn_env on non-exec format", " x:\n command: convert\n" + in + " output:\n format: json\n conn_env: DB\n", "not supported for format"},
|
||||||
|
{"convert output path and conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n path: o.sql\n", "either path or conn_env"},
|
||||||
|
{"convert output conn_env ok", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n", ""},
|
||||||
|
{"output secret conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: postgres://u:p@h/db\n", "environment variable name"},
|
||||||
|
{"merge needs two inputs", " x:\n command: merge\n" + in + out, "at least 2 input"},
|
||||||
|
{"input missing format", " x:\n command: convert\n inputs:\n - path: a\n" + out, "missing format"},
|
||||||
|
{"input unsupported format", " x:\n command: convert\n inputs:\n - path: a\n format: nope\n" + out, "unsupported input format"},
|
||||||
|
{"input file missing path", " x:\n command: convert\n inputs:\n - format: dbml\n" + out, "missing path"},
|
||||||
|
{"input file with conn_env", " x:\n command: convert\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n" + out, "does not use conn_env"},
|
||||||
|
{"input db missing conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n" + out, "requires conn_env"},
|
||||||
|
{"input db with path", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n" + out, "takes conn_env, not path"},
|
||||||
|
{"input db ok", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n" + out, ""},
|
||||||
|
{"input secret conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: \"host=h password=p\"\n" + out, "environment variable name"},
|
||||||
|
{"bad log size", " x:\n command: convert\n log_max_size: lots\n" + in + out, "log_max_size"},
|
||||||
|
{"absolute logfile", " x:\n command: convert\n logfile: /var/log/x.log\n" + in + out, "absolute paths"},
|
||||||
|
{"home path", " x:\n command: convert\n template: ~/t\n" + in + out, "home-relative"},
|
||||||
|
{"report path traversal", " x:\n command: inspect\n" + in + " report:\n format: json\n path: ../r.json\n", "escapes"},
|
||||||
|
{"script_dir traversal", " x:\n command: scripts-list\n script_dirs: [../x]\n", "escapes"},
|
||||||
|
|
||||||
|
{"templ valid", " x:\n command: templ\n" + in + " template: t.tmpl\n mode: table\n output:\n format: text\n path: o\n", ""},
|
||||||
|
{"templ pgsql input valid", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n template: t.tmpl\n", ""},
|
||||||
|
{"templ no inputs", " x:\n command: templ\n template: t.tmpl\n", "at least 1 input"},
|
||||||
|
{"templ no template", " x:\n command: templ\n" + in, "requires template"},
|
||||||
|
{"templ bad mode", " x:\n command: templ\n" + in + " template: t\n mode: weird\n", "unsupported mode"},
|
||||||
|
{"templ script dirs", " x:\n command: templ\n" + in + " template: t\n script_dirs: [s]\n", "script_dirs is not valid"},
|
||||||
|
{"templ db output", " x:\n command: templ\n" + in + " template: t\n output:\n conn_env: DB\n", "does not support database output"},
|
||||||
|
{"templ non-text output", " x:\n command: templ\n" + in + " template: t\n output:\n format: json\n path: o\n", "only output.format: text"},
|
||||||
|
{"templ input missing format", " x:\n command: templ\n inputs:\n - path: a\n template: t\n", "missing format"},
|
||||||
|
{"templ pgsql input without conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n template: t\n", "requires conn_env"},
|
||||||
|
{"templ pgsql input with path", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n template: t\n", "takes conn_env, not path"},
|
||||||
|
{"templ file input without path", " x:\n command: templ\n inputs:\n - format: dbml\n template: t\n", "missing path"},
|
||||||
|
{"templ file input with conn_env", " x:\n command: templ\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n template: t\n", "does not use conn_env"},
|
||||||
|
{"templ unsupported input format", " x:\n command: templ\n inputs:\n - path: a\n format: nope\n template: t\n", "unsupported templ input format"},
|
||||||
|
{"templ secret conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: a/b\n template: t\n", "environment variable name"},
|
||||||
|
|
||||||
|
{"split needs input", " x:\n command: split\n" + out, "at least 1 input"},
|
||||||
|
{"split script dirs", " x:\n command: split\n" + in + " script_dirs: [s]\n" + out, "script_dirs is not valid"},
|
||||||
|
{"split report", " x:\n command: split\n" + in + " report:\n format: json\n path: r\n" + out, "report is not valid"},
|
||||||
|
{"split db output", " x:\n command: split\n" + in + " output:\n format: pgsql\n conn_env: DB\n", "writes a file"},
|
||||||
|
|
||||||
|
{"inspect script dirs", " x:\n command: inspect\n" + in + " script_dirs: [s]\n report:\n path: r\n", "script_dirs is not valid"},
|
||||||
|
{"inspect output", " x:\n command: inspect\n" + in + out + " report:\n path: r\n", "output is not valid"},
|
||||||
|
{"inspect bad report format", " x:\n command: inspect\n" + in + " report:\n format: html\n path: r\n", "not supported"},
|
||||||
|
{"inspect report without path", " x:\n command: inspect\n" + in + " report:\n format: json\n", "requires report.path"},
|
||||||
|
{"inspect default format ok", " x:\n command: inspect\n" + in + " report:\n path: r.md\n", ""},
|
||||||
|
{"diff summary without path ok", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: summary\n", ""},
|
||||||
|
{"diff json needs path", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: json\n", "requires report.path"},
|
||||||
|
{"diff output", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n" + out + " report:\n format: summary\n", "output is not valid"},
|
||||||
|
{"diff script dirs", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n script_dirs: [s]\n report:\n format: summary\n", "script_dirs is not valid"},
|
||||||
|
{"diff no report", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n", "requires a report block"},
|
||||||
|
|
||||||
|
{"scripts-list inputs", " x:\n command: scripts-list\n script_dirs: [s]\n" + in, "inputs is not valid"},
|
||||||
|
{"scripts-list output", " x:\n command: scripts-list\n script_dirs: [s]\n" + out, "output is not valid"},
|
||||||
|
{"scripts-exec inputs", " x:\n command: scripts-exec\n script_dirs: [s]\n" + in + " output:\n conn_env: DB\n", "inputs is not valid"},
|
||||||
|
{"scripts-exec report", " x:\n command: scripts-exec\n script_dirs: [s]\n report:\n path: r\n output:\n conn_env: DB\n", "report is not valid"},
|
||||||
|
{"scripts-exec output path", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n path: p\n", "output.path is not supported"},
|
||||||
|
{"scripts-exec non-pgsql", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: mssql\n", "only supports pgsql"},
|
||||||
|
{"scripts-exec secret conn_env", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: \"postgres://u@h/d\"\n", "environment variable name"},
|
||||||
|
{"scripts-exec no script dirs", " x:\n command: scripts-exec\n output:\n conn_env: DB\n", "requires at least one script_dir"},
|
||||||
|
{"scripts-exec pgsql format ok", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: pgsql\n", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := validateYAML(t, tt.job)
|
||||||
|
if tt.want == "" {
|
||||||
|
if got != "" {
|
||||||
|
t.Errorf("expected valid, got: %s", got)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, tt.want) {
|
||||||
|
t.Errorf("error %q does not contain %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFromJobInputShape(t *testing.T) {
|
||||||
|
producer := " p:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: out.json\n"
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"path", " - from_job: p\n path: x\n", "takes no path"},
|
||||||
|
{"format", " - from_job: p\n format: json\n", "drop format"},
|
||||||
|
{"conn_env", " - from_job: p\n conn_env: DB\n", "takes no conn_env"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
for _, cmd := range []string{"convert", "templ"} {
|
||||||
|
t.Run(cmd+"/"+tt.name, func(t *testing.T) {
|
||||||
|
extra := " output:\n format: json\n path: o.json\n"
|
||||||
|
if cmd == "templ" {
|
||||||
|
extra = " template: t.tmpl\n"
|
||||||
|
}
|
||||||
|
got := validateYAML(t, producer+" c:\n command: "+cmd+"\n inputs:\n"+tt.input+extra)
|
||||||
|
if !strings.Contains(got, tt.want) {
|
||||||
|
t.Errorf("error %q does not contain %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolvedLogPolicy(t *testing.T) {
|
||||||
|
keep2 := 2
|
||||||
|
keep0 := 0
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
job Job
|
||||||
|
want LogPolicy
|
||||||
|
}{
|
||||||
|
{"built-in defaults", Job{}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||||
|
{"file defaults", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}}, LogPolicy{MaxSizeBytes: 1 << 20, Keep: 7}},
|
||||||
|
{"file defaults invalid size falls back", Job{fileDefaults: &Defaults{LogMaxSize: "junk", LogKeep: 0}}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||||
|
{"job overrides file", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}, LogMaxSize: "2kb", LogKeep: &keep2}, LogPolicy{MaxSizeBytes: 2 << 10, Keep: 2}},
|
||||||
|
{"job keep zero is honoured", Job{LogKeep: &keep0}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: 0}},
|
||||||
|
{"job invalid size ignored", Job{LogMaxSize: "junk"}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := tt.job.ResolvedLogPolicy(); got != tt.want {
|
||||||
|
t.Errorf("got %+v, want %+v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadAppliesFileDefaultsAndDir(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
p := filepath.Join(dir, "relspec.yml")
|
||||||
|
write(t, p, "version: 1\ndefaults:\n log_max_size: 1MB\n log_keep: 9\n"+"jobs:\n a:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n")
|
||||||
|
set, err := Load([]string{p})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
job := set.Jobs["a"]
|
||||||
|
if job.Dir() != dir {
|
||||||
|
t.Errorf("Dir = %q, want %q", job.Dir(), dir)
|
||||||
|
}
|
||||||
|
if pol := job.ResolvedLogPolicy(); pol.MaxSizeBytes != 1<<20 || pol.Keep != 9 {
|
||||||
|
t.Errorf("policy %+v", pol)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSetNamesSorted(t *testing.T) {
|
||||||
|
set := &Set{Jobs: map[string]*Job{"b": {}, "a": {}, "c": {}}}
|
||||||
|
if got := strings.Join(set.Names(), ","); got != "a,b,c" {
|
||||||
|
t.Errorf("got %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPlanErrors(t *testing.T) {
|
||||||
|
set := loadOne(t, "version: 1\njobs:\n"+
|
||||||
|
" a:\n command: convert\n depends_on: [ghost]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n"+
|
||||||
|
" b:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o2.json\n")
|
||||||
|
|
||||||
|
if _, err := set.Plan("nope", true); err == nil || !strings.Contains(err.Error(), "unknown job") || !strings.Contains(err.Error(), "a, b") {
|
||||||
|
t.Errorf("unknown job: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "unknown job \"ghost\"") {
|
||||||
|
t.Errorf("unknown dependency: %v", err)
|
||||||
|
}
|
||||||
|
// Without dependencies the declared dependency is not walked.
|
||||||
|
if got, err := set.Plan("a", false); err != nil || len(got) != 1 || got[0].Name != "a" {
|
||||||
|
t.Errorf("no-deps plan: %v %v", got, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPlanCycleAtRuntime(t *testing.T) {
|
||||||
|
set := &Set{Jobs: map[string]*Job{
|
||||||
|
"a": {Name: "a", DependsOn: []string{"b"}},
|
||||||
|
"b": {Name: "b", DependsOn: []string{"a"}},
|
||||||
|
}}
|
||||||
|
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "cycle") {
|
||||||
|
t.Errorf("want cycle error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDiscoverErrorsAndFiltering(t *testing.T) {
|
||||||
|
if _, err := Discover(filepath.Join(t.TempDir(), "missing")); err == nil {
|
||||||
|
t.Error("missing dir must fail")
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
for _, f := range []string{"relspec.yaml", "relspec.b.yml", "relspec.a.yaml", "relspec.txt", "other.yml", "relspec"} {
|
||||||
|
write(t, filepath.Join(dir, f), "")
|
||||||
|
}
|
||||||
|
if err := os.Mkdir(filepath.Join(dir, "relspec.dir.yml"), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got, err := Discover(dir)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var names []string
|
||||||
|
for _, p := range got {
|
||||||
|
names = append(names, filepath.Base(p))
|
||||||
|
}
|
||||||
|
if strings.Join(names, ",") != "relspec.yaml,relspec.a.yaml,relspec.b.yml" {
|
||||||
|
t.Errorf("got %v", names)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeJoinCases(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
if got, err := SafeJoin(root, "sub/file.sql"); err != nil || !strings.HasSuffix(got, filepath.Join("sub", "file.sql")) {
|
||||||
|
t.Errorf("nested: %q %v", got, err)
|
||||||
|
}
|
||||||
|
for _, bad := range []string{"", "/etc/passwd", "~/x", "..", "../x", "a/../../x"} {
|
||||||
|
if _, err := SafeJoin(root, bad); err == nil {
|
||||||
|
t.Errorf("SafeJoin(%q) must fail", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := SafeJoin(filepath.Join(root, "does", "not", "exist"), "x"); err == nil {
|
||||||
|
t.Error("unresolvable root must fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLooksLikeSecret(t *testing.T) {
|
||||||
|
for in, want := range map[string]bool{
|
||||||
|
"": false, "DB_URL": false, "MY_DB": false,
|
||||||
|
"postgres://u:p@h/db": true, "host=h": true, "a b": true, "a/b": true, "u@h": true, "k:v": true,
|
||||||
|
} {
|
||||||
|
if got := looksLikeSecret(in); got != want {
|
||||||
|
t.Errorf("looksLikeSecret(%q) = %v, want %v", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
|
|||||||
} else {
|
} else {
|
||||||
// Schema doesn't exist, add it
|
// Schema doesn't exist, add it
|
||||||
newSchema := cloneSchema(srcSchema)
|
newSchema := cloneSchema(srcSchema)
|
||||||
|
if len(opts.SkipTableNames) > 0 {
|
||||||
|
kept := newSchema.Tables[:0]
|
||||||
|
for _, t := range newSchema.Tables {
|
||||||
|
if !opts.SkipTableNames[strings.ToLower(t.SQLName())] {
|
||||||
|
kept = append(kept, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
newSchema.Tables = kept
|
||||||
|
}
|
||||||
target.Schemas = append(target.Schemas, newSchema)
|
target.Schemas = append(target.Schemas, newSchema)
|
||||||
r.SchemasAdded++
|
r.SchemasAdded++
|
||||||
}
|
}
|
||||||
@@ -440,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table {
|
|||||||
Description: table.Description,
|
Description: table.Description,
|
||||||
Schema: table.Schema,
|
Schema: table.Schema,
|
||||||
Comment: table.Comment,
|
Comment: table.Comment,
|
||||||
|
Tablespace: table.Tablespace,
|
||||||
|
GUID: table.GUID,
|
||||||
Sequence: table.Sequence,
|
Sequence: table.Sequence,
|
||||||
UpdatedAt: table.UpdatedAt,
|
UpdatedAt: table.UpdatedAt,
|
||||||
Columns: make(map[string]*models.Column),
|
Columns: make(map[string]*models.Column),
|
||||||
@@ -469,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table {
|
|||||||
newTable.Indexes[idxName] = cloneIndex(index)
|
newTable.Indexes[idxName] = cloneIndex(index)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Clone relationships
|
||||||
|
if table.Relationships != nil {
|
||||||
|
newTable.Relationships = make(map[string]*models.Relationship, len(table.Relationships))
|
||||||
|
for relName, rel := range table.Relationships {
|
||||||
|
newTable.Relationships[relName] = cloneRelation(rel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return newTable
|
return newTable
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,277 @@
|
|||||||
|
package merge
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMergeSequences(t *testing.T) {
|
||||||
|
target := models.InitSchema("public")
|
||||||
|
target.Sequences = []*models.Sequence{{Name: "Existing", StartValue: 1, IncrementBy: 1}}
|
||||||
|
|
||||||
|
source := models.InitSchema("public")
|
||||||
|
source.Sequences = []*models.Sequence{
|
||||||
|
{Name: "existing", StartValue: 100, IncrementBy: 10}, // conflicting: must not overwrite
|
||||||
|
{Name: "fresh", StartValue: 5, IncrementBy: 2, MinValue: 1, MaxValue: 99, CacheSize: 3, Cycle: true, OwnedByTable: "t", OwnedByColumn: "id", Comment: "c", Description: "d"},
|
||||||
|
}
|
||||||
|
|
||||||
|
res := &MergeResult{}
|
||||||
|
res.mergeSequences(target, source)
|
||||||
|
|
||||||
|
if res.SequencesAdded != 1 || len(target.Sequences) != 2 {
|
||||||
|
t.Fatalf("added=%d len=%d", res.SequencesAdded, len(target.Sequences))
|
||||||
|
}
|
||||||
|
if target.Sequences[0].StartValue != 1 || target.Sequences[0].IncrementBy != 1 {
|
||||||
|
t.Errorf("existing sequence was modified: %+v", target.Sequences[0])
|
||||||
|
}
|
||||||
|
added := target.Sequences[1]
|
||||||
|
if added.Name != "fresh" || added.StartValue != 5 || added.IncrementBy != 2 || added.MinValue != 1 || added.MaxValue != 99 ||
|
||||||
|
added.CacheSize != 3 || !added.Cycle || added.OwnedByTable != "t" || added.OwnedByColumn != "id" || added.Comment != "c" || added.Description != "d" {
|
||||||
|
t.Errorf("clone lost fields: %+v", added)
|
||||||
|
}
|
||||||
|
if added == source.Sequences[1] {
|
||||||
|
t.Error("sequence must be cloned, not shared")
|
||||||
|
}
|
||||||
|
source.Sequences[1].StartValue = 777
|
||||||
|
if added.StartValue != 5 {
|
||||||
|
t.Error("clone must be independent of source")
|
||||||
|
}
|
||||||
|
if cloneSequence(nil) != nil {
|
||||||
|
t.Error("cloneSequence(nil) must be nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloneSchemaIsIndependent(t *testing.T) {
|
||||||
|
src := models.InitSchema("public")
|
||||||
|
src.Description, src.Owner, src.Comment, src.Sequence = "d", "o", "c", 4
|
||||||
|
src.Permissions["r"] = "all"
|
||||||
|
src.Metadata["k"] = "v"
|
||||||
|
src.Scripts = []*models.Script{{Name: "s"}}
|
||||||
|
|
||||||
|
tbl := models.InitTable("t", "public")
|
||||||
|
col := models.InitColumn("id", "t", "public")
|
||||||
|
col.Type = "integer"
|
||||||
|
tbl.Columns["id"] = col
|
||||||
|
tbl.Constraints["pk"] = &models.Constraint{Name: "pk", Type: models.PrimaryKeyConstraint, Columns: []string{"id"}}
|
||||||
|
tbl.Indexes["i"] = &models.Index{Name: "i", Columns: []string{"id"}, Include: []string{"x"}}
|
||||||
|
tbl.Metadata["tm"] = 1
|
||||||
|
src.Tables = []*models.Table{tbl}
|
||||||
|
|
||||||
|
v := models.InitView("v", "public")
|
||||||
|
v.Definition = "select 1"
|
||||||
|
v.Columns["c"] = &models.Column{Name: "c"}
|
||||||
|
v.Metadata["vm"] = 1
|
||||||
|
src.Views = []*models.View{v}
|
||||||
|
src.Sequences = []*models.Sequence{{Name: "sq", StartValue: 3}}
|
||||||
|
src.Enums = []*models.Enum{{Name: "e", Values: []string{"a", "b"}}}
|
||||||
|
src.Relations = []*models.Relationship{{Name: "r", FromColumns: []string{"a"}, ToColumns: []string{"b"}, Properties: map[string]string{"p": "q"}}}
|
||||||
|
|
||||||
|
got := cloneSchema(src)
|
||||||
|
if got == src || got.Name != "public" || got.Description != "d" || got.Owner != "o" || got.Comment != "c" || got.Sequence != 4 {
|
||||||
|
t.Fatalf("scalar fields: %+v", got)
|
||||||
|
}
|
||||||
|
if got.Permissions["r"] != "all" || got.Metadata["k"] != "v" || len(got.Scripts) != 1 {
|
||||||
|
t.Errorf("maps/scripts: %+v", got)
|
||||||
|
}
|
||||||
|
if len(got.Tables) != 1 || got.Tables[0] == tbl || got.Tables[0].Columns["id"] == col || got.Tables[0].Columns["id"].Type != "integer" {
|
||||||
|
t.Errorf("tables not deep cloned: %+v", got.Tables)
|
||||||
|
}
|
||||||
|
if len(got.Views) != 1 || got.Views[0] == v || got.Views[0].Definition != "select 1" || got.Views[0].Columns["c"] == v.Columns["c"] || got.Views[0].Metadata["vm"] != 1 {
|
||||||
|
t.Errorf("views not deep cloned: %+v", got.Views)
|
||||||
|
}
|
||||||
|
if len(got.Sequences) != 1 || got.Sequences[0] == src.Sequences[0] || got.Sequences[0].StartValue != 3 {
|
||||||
|
t.Errorf("sequences: %+v", got.Sequences)
|
||||||
|
}
|
||||||
|
if len(got.Enums) != 1 || got.Enums[0] == src.Enums[0] || strings.Join(got.Enums[0].Values, ",") != "a,b" {
|
||||||
|
t.Errorf("enums: %+v", got.Enums)
|
||||||
|
}
|
||||||
|
if len(got.Relations) != 1 || got.Relations[0] == src.Relations[0] || got.Relations[0].Properties["p"] != "q" {
|
||||||
|
t.Errorf("relations: %+v", got.Relations)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mutating the clone must not touch the source.
|
||||||
|
got.Permissions["r"] = "none"
|
||||||
|
got.Metadata["k"] = "changed"
|
||||||
|
got.Tables[0].Columns["id"].Type = "text"
|
||||||
|
got.Tables[0].Constraints["pk"].Columns[0] = "zzz"
|
||||||
|
got.Tables[0].Indexes["i"].Columns[0] = "zzz"
|
||||||
|
got.Tables[0].Metadata["tm"] = 2
|
||||||
|
got.Enums[0].Values[0] = "zzz"
|
||||||
|
got.Relations[0].FromColumns[0] = "zzz"
|
||||||
|
got.Relations[0].Properties["p"] = "zzz"
|
||||||
|
got.Views[0].Columns["c"].Name = "zzz"
|
||||||
|
if src.Permissions["r"] != "all" || src.Metadata["k"] != "v" || col.Type != "integer" ||
|
||||||
|
tbl.Constraints["pk"].Columns[0] != "id" || tbl.Indexes["i"].Columns[0] != "id" || tbl.Metadata["tm"] != 1 ||
|
||||||
|
src.Enums[0].Values[0] != "a" || src.Relations[0].FromColumns[0] != "a" || src.Relations[0].Properties["p"] != "q" ||
|
||||||
|
v.Columns["c"].Name != "c" {
|
||||||
|
t.Error("clone shares state with the source")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cloneSchema(nil) != nil {
|
||||||
|
t.Error("cloneSchema(nil) must be nil")
|
||||||
|
}
|
||||||
|
bare := cloneSchema(&models.Schema{Name: "bare"})
|
||||||
|
if bare.Permissions != nil || bare.Metadata != nil {
|
||||||
|
t.Errorf("nil maps must stay nil: %+v", bare)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloneNilInputs(t *testing.T) {
|
||||||
|
if cloneTable(nil) != nil || cloneColumn(nil) != nil || cloneConstraint(nil) != nil || cloneIndex(nil) != nil ||
|
||||||
|
cloneView(nil) != nil || cloneEnum(nil) != nil || cloneRelation(nil) != nil || cloneDomain(nil) != nil {
|
||||||
|
t.Error("clone of nil must be nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloneDomainAndRelation(t *testing.T) {
|
||||||
|
d := &models.Domain{Name: "d", Description: "x", Comment: "c", Sequence: 2, Metadata: map[string]any{"k": 1}, Tables: []*models.DomainTable{{TableName: "t", SchemaName: "s"}}}
|
||||||
|
cd := cloneDomain(d)
|
||||||
|
if cd == d || cd.Name != "d" || cd.Description != "x" || cd.Comment != "c" || cd.Sequence != 2 || cd.Metadata["k"] != 1 || len(cd.Tables) != 1 {
|
||||||
|
t.Errorf("domain clone: %+v", cd)
|
||||||
|
}
|
||||||
|
cd.Metadata["k"] = 2
|
||||||
|
if d.Metadata["k"] != 1 {
|
||||||
|
t.Error("domain metadata shared")
|
||||||
|
}
|
||||||
|
|
||||||
|
r := &models.Relationship{Name: "r", Type: "one_to_many", FromTable: "a", FromSchema: "s", ToTable: "b", ToSchema: "s", ForeignKey: "fk", ThroughTable: "l", ThroughSchema: "s", Description: "d", Sequence: 3}
|
||||||
|
cr := cloneRelation(r)
|
||||||
|
if cr == r || cr.Name != "r" || cr.Type != "one_to_many" || cr.FromTable != "a" || cr.ToTable != "b" || cr.ForeignKey != "fk" || cr.ThroughTable != "l" || cr.Description != "d" || cr.Sequence != 3 {
|
||||||
|
t.Errorf("relation clone: %+v", cr)
|
||||||
|
}
|
||||||
|
if cr.Properties != nil {
|
||||||
|
t.Errorf("nil properties must stay nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractTypeParts(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
col models.Column
|
||||||
|
wantType string
|
||||||
|
wantLen, wantPrec, wantScale int
|
||||||
|
}{
|
||||||
|
{"plain", models.Column{Type: "TEXT"}, "text", 0, 0, 0},
|
||||||
|
{"trim and lower", models.Column{Type: " Integer "}, "integer", 0, 0, 0},
|
||||||
|
{"embedded length", models.Column{Type: "varchar(50)"}, "varchar", 50, 0, 0},
|
||||||
|
{"embedded precision and scale", models.Column{Type: "numeric(10,2)"}, "numeric", 0, 10, 2},
|
||||||
|
{"embedded with spaces", models.Column{Type: "numeric( 10 , 2 )"}, "numeric", 0, 10, 2},
|
||||||
|
{"fields win over embedded precision", models.Column{Type: "numeric(10,2)", Precision: 12, Scale: 4}, "numeric", 0, 12, 4},
|
||||||
|
{"fields win over embedded length", models.Column{Type: "varchar(50)", Length: 80}, "varchar", 80, 0, 0},
|
||||||
|
{"precision field blocks embedded length", models.Column{Type: "varchar(50)", Precision: 5}, "varchar", 0, 5, 0},
|
||||||
|
{"non-numeric modifier", models.Column{Type: "varchar(max)"}, "varchar", 0, 0, 0},
|
||||||
|
{"zero modifier", models.Column{Type: "char(0)"}, "char", 0, 0, 0},
|
||||||
|
{"serial sugar", models.Column{Type: "bigserial"}, "bigint", 0, 0, 0},
|
||||||
|
{"smallserial sugar", models.Column{Type: "smallserial"}, "smallint", 0, 0, 0},
|
||||||
|
{"three modifiers ignored", models.Column{Type: "x(1,2,3)"}, "x", 0, 0, 0},
|
||||||
|
{"empty", models.Column{}, "", 0, 0, 0},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
col := tt.col
|
||||||
|
gt, gl, gp, gs := extractTypeParts(&col)
|
||||||
|
if gt != tt.wantType || gl != tt.wantLen || gp != tt.wantPrec || gs != tt.wantScale {
|
||||||
|
t.Errorf("got (%q,%d,%d,%d), want (%q,%d,%d,%d)", gt, gl, gp, gs, tt.wantType, tt.wantLen, tt.wantPrec, tt.wantScale)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnTypeConflict(t *testing.T) {
|
||||||
|
c := func(typ string, l, p, s int) *models.Column {
|
||||||
|
return &models.Column{Type: typ, Length: l, Precision: p, Scale: s}
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
a, b *models.Column
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"nil target", nil, c("text", 0, 0, 0), false},
|
||||||
|
{"nil source", c("text", 0, 0, 0), nil, false},
|
||||||
|
{"same", c("text", 0, 0, 0), c("TEXT", 0, 0, 0), false},
|
||||||
|
{"different base", c("text", 0, 0, 0), c("integer", 0, 0, 0), true},
|
||||||
|
{"embedded equals field", c("varchar(50)", 0, 0, 0), c("varchar", 50, 0, 0), false},
|
||||||
|
{"different length", c("varchar", 50, 0, 0), c("varchar", 80, 0, 0), true},
|
||||||
|
{"different scale", c("numeric", 0, 10, 2), c("numeric", 0, 10, 3), true},
|
||||||
|
{"serial vs int", c("bigserial", 0, 0, 0), c("bigint", 0, 0, 0), false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := columnTypeConflict(tt.a, tt.b); got != tt.want {
|
||||||
|
t.Errorf("got %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDescribeColumnType(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
col *models.Column
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{nil, ""},
|
||||||
|
{&models.Column{}, ""},
|
||||||
|
{&models.Column{Type: " "}, ""},
|
||||||
|
{&models.Column{Type: "text"}, "text"},
|
||||||
|
{&models.Column{Type: " numeric ", Precision: 10, Scale: 2}, "numeric(10,2)"},
|
||||||
|
{&models.Column{Type: "numeric", Precision: 10}, "numeric(10)"},
|
||||||
|
{&models.Column{Type: "varchar", Length: 50}, "varchar(50)"},
|
||||||
|
{&models.Column{Type: "varchar", Length: 50, Precision: 7}, "varchar(7)"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := describeColumnType(tt.col); got != tt.want {
|
||||||
|
t.Errorf("describeColumnType(%+v) = %q, want %q", tt.col, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFirstNonEmpty(t *testing.T) {
|
||||||
|
if got := firstNonEmpty("", " ", "x", "y"); got != "x" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := firstNonEmpty(); got != "" {
|
||||||
|
t.Errorf("none: %q", got)
|
||||||
|
}
|
||||||
|
if got := firstNonEmpty("", " "); got != "" {
|
||||||
|
t.Errorf("all blank: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetColumnTypeConflictSummary(t *testing.T) {
|
||||||
|
conflicts := []ColumnTypeConflict{
|
||||||
|
{Schema: "s", Table: "t", Column: "a", TargetType: "text", SourceType: "integer"},
|
||||||
|
{Schema: "s", Table: "t", Column: "b", TargetType: "int", SourceType: "text"},
|
||||||
|
{Schema: "s", Table: "u", Column: "c", TargetType: "x", SourceType: "y"},
|
||||||
|
}
|
||||||
|
res := &MergeResult{TypeConflicts: conflicts}
|
||||||
|
|
||||||
|
if GetColumnTypeConflictSummary(nil, 5) != "" || GetColumnTypeConflictSummary(&MergeResult{}, 5) != "" {
|
||||||
|
t.Error("no conflicts must yield empty summary")
|
||||||
|
}
|
||||||
|
|
||||||
|
all := GetColumnTypeConflictSummary(res, 0)
|
||||||
|
if !strings.Contains(all, "column type conflicts detected:") || !strings.Contains(all, "s.t.a: target=text source=integer") ||
|
||||||
|
!strings.Contains(all, "s.u.c: target=x source=y") || strings.Contains(all, "more") {
|
||||||
|
t.Errorf("unlimited summary:\n%s", all)
|
||||||
|
}
|
||||||
|
if neg := GetColumnTypeConflictSummary(res, -1); neg != all {
|
||||||
|
t.Error("negative limit must behave as unlimited")
|
||||||
|
}
|
||||||
|
|
||||||
|
limited := GetColumnTypeConflictSummary(res, 2)
|
||||||
|
if !strings.Contains(limited, "s.t.b") || strings.Contains(limited, "s.u.c") || !strings.HasSuffix(limited, "... and 1 more") {
|
||||||
|
t.Errorf("limited summary:\n%s", limited)
|
||||||
|
}
|
||||||
|
exact := GetColumnTypeConflictSummary(res, 3)
|
||||||
|
if strings.Contains(exact, "more") {
|
||||||
|
t.Errorf("limit == len must not truncate:\n%s", exact)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMinHelper(t *testing.T) {
|
||||||
|
if min(1, 2) != 1 || min(2, 1) != 1 || min(3, 3) != 3 {
|
||||||
|
t.Error("min")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
package merge
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func sourceWithRelationship() *models.Database {
|
||||||
|
db := models.InitDatabase("src")
|
||||||
|
s := models.InitSchema("sales")
|
||||||
|
orders := models.InitTable("orders", "sales")
|
||||||
|
orders.Tablespace = "fast"
|
||||||
|
orders.GUID = "guid-1"
|
||||||
|
orders.Relationships["fk_cust"] = &models.Relationship{
|
||||||
|
Name: "fk_cust", FromTable: "orders", ToTable: "customers",
|
||||||
|
FromColumns: []string{"cust_id"}, ToColumns: []string{"id"},
|
||||||
|
}
|
||||||
|
s.Tables = append(s.Tables, orders, models.InitTable("Audit", "sales"))
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloneTable_CopiesRelationshipsTablespaceGUID(t *testing.T) {
|
||||||
|
src := sourceWithRelationship()
|
||||||
|
target := models.InitDatabase("tgt")
|
||||||
|
MergeDatabases(target, src, nil)
|
||||||
|
|
||||||
|
got := target.Schemas[0].Tables[0]
|
||||||
|
if got.Tablespace != "fast" || got.GUID != "guid-1" {
|
||||||
|
t.Errorf("tablespace/guid lost: %+v", got)
|
||||||
|
}
|
||||||
|
rel := got.Relationships["fk_cust"]
|
||||||
|
if rel == nil || rel.ToTable != "customers" {
|
||||||
|
t.Fatalf("relationship lost: %+v", got.Relationships)
|
||||||
|
}
|
||||||
|
if rel == src.Schemas[0].Tables[0].Relationships["fk_cust"] {
|
||||||
|
t.Error("relationship must be deep-copied")
|
||||||
|
}
|
||||||
|
rel.FromColumns[0] = "changed"
|
||||||
|
if src.Schemas[0].Tables[0].Relationships["fk_cust"].FromColumns[0] != "cust_id" {
|
||||||
|
t.Error("relationship columns shared with source")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMerge_SkipTablesAppliesToNewSchemas(t *testing.T) {
|
||||||
|
src := sourceWithRelationship()
|
||||||
|
target := models.InitDatabase("tgt")
|
||||||
|
MergeDatabases(target, src, &MergeOptions{SkipTableNames: map[string]bool{"audit": true}})
|
||||||
|
|
||||||
|
tables := target.Schemas[0].Tables
|
||||||
|
if len(tables) != 1 || tables[0].Name != "orders" {
|
||||||
|
t.Errorf("skipped table copied into new schema: %+v", tables)
|
||||||
|
}
|
||||||
|
if len(src.Schemas[0].Tables) != 2 {
|
||||||
|
t.Error("source must not be modified")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -20,6 +20,7 @@ const (
|
|||||||
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
|
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
|
||||||
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
|
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
|
||||||
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
|
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
|
||||||
|
MySQLDatabaseType DatabaseType = "mysql" // MySQL/MariaDB database
|
||||||
)
|
)
|
||||||
|
|
||||||
// Database represents the complete database schema
|
// Database represents the complete database schema
|
||||||
|
|||||||
@@ -0,0 +1,232 @@
|
|||||||
|
package models
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSQLNameLowercases(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
got string
|
||||||
|
}{
|
||||||
|
{"database", (&Database{Name: "MyDB"}).SQLName()},
|
||||||
|
{"domain", (&Domain{Name: "MyDomain"}).SQLName()},
|
||||||
|
{"schema", (&Schema{Name: "MySchema"}).SQLName()},
|
||||||
|
{"table", (&Table{Name: "MyTable"}).SQLName()},
|
||||||
|
{"view", (&View{Name: "MyView"}).SQLName()},
|
||||||
|
{"sequence", (&Sequence{Name: "MySeq"}).SQLName()},
|
||||||
|
{"column", (&Column{Name: "MyCol"}).SQLName()},
|
||||||
|
{"index", (&Index{Name: "MyIdx"}).SQLName()},
|
||||||
|
{"relationship", (&Relationship{Name: "MyRel"}).SQLName()},
|
||||||
|
{"constraint", (&Constraint{Name: "MyCon"}).SQLName()},
|
||||||
|
{"enum", (&Enum{Name: "MyEnum"}).SQLName()},
|
||||||
|
{"script", (&Script{Name: "MyScript"}).SQLName()},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if tt.got == "" || tt.got != lower(tt.got) {
|
||||||
|
t.Errorf("SQLName not lowercase: %q", tt.got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if got := (&Table{}).SQLName(); got != "" {
|
||||||
|
t.Errorf("empty name: %q", got)
|
||||||
|
}
|
||||||
|
if got := (&Table{Name: "MyTable"}).SQLName(); got != "mytable" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func lower(s string) string {
|
||||||
|
b := []byte(s)
|
||||||
|
for i, c := range b {
|
||||||
|
if c >= 'A' && c <= 'Z' {
|
||||||
|
b[i] = c + 32
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateDatePropagates(t *testing.T) {
|
||||||
|
db := InitDatabase("d")
|
||||||
|
schema := InitSchema("s")
|
||||||
|
schema.RefDatabase = db
|
||||||
|
table := InitTable("t", "s")
|
||||||
|
table.RefSchema = schema
|
||||||
|
|
||||||
|
table.UpdateDate()
|
||||||
|
for name, v := range map[string]string{"table": table.UpdatedAt, "schema": schema.UpdatedAt, "database": db.UpdatedAt} {
|
||||||
|
ts, err := time.Parse(time.RFC3339, v)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("%s UpdatedAt %q: %v", name, v, err)
|
||||||
|
}
|
||||||
|
if time.Since(ts) > time.Minute {
|
||||||
|
t.Errorf("%s UpdatedAt too old: %v", name, ts)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Without references only the receiver is updated.
|
||||||
|
lone := InitTable("lone", "s")
|
||||||
|
lone.UpdateDate()
|
||||||
|
if lone.UpdatedAt == "" {
|
||||||
|
t.Error("lone table not updated")
|
||||||
|
}
|
||||||
|
loneSchema := InitSchema("x")
|
||||||
|
loneSchema.UpdateDate()
|
||||||
|
if loneSchema.UpdatedAt == "" {
|
||||||
|
t.Error("lone schema not updated")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetPrimaryKey(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
cols []*Column
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"none", []*Column{{Name: "a"}}, ""},
|
||||||
|
{"single", []*Column{{Name: "a"}, {Name: "id", IsPrimaryKey: true}}, "id"},
|
||||||
|
{"composite ordered by sequence", []*Column{
|
||||||
|
{Name: "a", IsPrimaryKey: true, Sequence: 2},
|
||||||
|
{Name: "b", IsPrimaryKey: true, Sequence: 1},
|
||||||
|
}, "b"},
|
||||||
|
{"composite without sequence falls back to name", []*Column{
|
||||||
|
{Name: "z", IsPrimaryKey: true},
|
||||||
|
{Name: "m", IsPrimaryKey: true},
|
||||||
|
}, "m"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
tbl := InitTable("t", "s")
|
||||||
|
for _, c := range tt.cols {
|
||||||
|
tbl.Columns[c.Name] = c
|
||||||
|
}
|
||||||
|
got := tbl.GetPrimaryKey()
|
||||||
|
if tt.want == "" {
|
||||||
|
if got != nil {
|
||||||
|
t.Errorf("expected nil, got %s", got.Name)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if got == nil || got.Name != tt.want {
|
||||||
|
t.Errorf("got %v, want %s", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if InitTable("empty", "s").GetPrimaryKey() != nil {
|
||||||
|
t.Error("empty table must have no PK")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnLess(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
a, b *Column
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b", Sequence: 2}, true},
|
||||||
|
{&Column{Name: "a", Sequence: 2}, &Column{Name: "b", Sequence: 1}, false},
|
||||||
|
{&Column{Name: "a"}, &Column{Name: "b"}, true},
|
||||||
|
{&Column{Name: "b"}, &Column{Name: "a"}, false},
|
||||||
|
{&Column{Name: "b", Sequence: 1}, &Column{Name: "a"}, false}, // one side unsequenced: by name
|
||||||
|
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b"}, true},
|
||||||
|
}
|
||||||
|
for i, tt := range tests {
|
||||||
|
if got := columnLess(tt.a, tt.b); got != tt.want {
|
||||||
|
t.Errorf("case %d: got %v, want %v", i, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetForeignKeys(t *testing.T) {
|
||||||
|
tbl := InitTable("t", "s")
|
||||||
|
add := func(name string, typ ConstraintType, seq uint) {
|
||||||
|
c := InitConstraint(name, typ)
|
||||||
|
c.Sequence = seq
|
||||||
|
tbl.Constraints[name] = c
|
||||||
|
}
|
||||||
|
add("pk", PrimaryKeyConstraint, 0)
|
||||||
|
add("fk_b", ForeignKeyConstraint, 0)
|
||||||
|
add("fk_a", ForeignKeyConstraint, 0)
|
||||||
|
add("uq", UniqueConstraint, 0)
|
||||||
|
|
||||||
|
got := tbl.GetForeignKeys()
|
||||||
|
if len(got) != 2 || got[0].Name != "fk_a" || got[1].Name != "fk_b" {
|
||||||
|
t.Errorf("by name: %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
tbl.Constraints["fk_a"].Sequence = 5
|
||||||
|
tbl.Constraints["fk_b"].Sequence = 2
|
||||||
|
got = tbl.GetForeignKeys()
|
||||||
|
if got[0].Name != "fk_b" || got[1].Name != "fk_a" {
|
||||||
|
t.Errorf("by sequence: %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := InitTable("e", "s").GetForeignKeys(); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("empty table must give non-nil empty slice, got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInitConstructors(t *testing.T) {
|
||||||
|
db := InitDatabase("db")
|
||||||
|
if db.Name != "db" || db.Schemas == nil || db.Domains == nil || db.Metadata == nil || db.GUID == "" {
|
||||||
|
t.Errorf("InitDatabase: %+v", db)
|
||||||
|
}
|
||||||
|
s := InitSchema("s")
|
||||||
|
if s.Name != "s" || s.Tables == nil || s.Views == nil || s.Sequences == nil || s.Permissions == nil || s.Metadata == nil || s.Scripts == nil || s.GUID == "" {
|
||||||
|
t.Errorf("InitSchema: %+v", s)
|
||||||
|
}
|
||||||
|
tb := InitTable("t", "s")
|
||||||
|
if tb.Name != "t" || tb.Schema != "s" || tb.Columns == nil || tb.Constraints == nil || tb.Indexes == nil || tb.Relationships == nil || tb.Metadata == nil || tb.GUID == "" {
|
||||||
|
t.Errorf("InitTable: %+v", tb)
|
||||||
|
}
|
||||||
|
c := InitColumn("c", "t", "s")
|
||||||
|
if c.Name != "c" || c.Table != "t" || c.Schema != "s" || c.Metadata == nil || c.GUID == "" {
|
||||||
|
t.Errorf("InitColumn: %+v", c)
|
||||||
|
}
|
||||||
|
ix := InitIndex("i", "t", "s")
|
||||||
|
if ix.Name != "i" || ix.Table != "t" || ix.Schema != "s" || ix.Columns == nil || ix.Include == nil || ix.Metadata == nil || ix.GUID == "" {
|
||||||
|
t.Errorf("InitIndex: %+v", ix)
|
||||||
|
}
|
||||||
|
r := InitRelation("r", "s")
|
||||||
|
if r.Name != "r" || r.FromSchema != "s" || r.ToSchema != "s" || r.Properties == nil || r.FromColumns == nil || r.ToColumns == nil || r.GUID == "" {
|
||||||
|
t.Errorf("InitRelation: %+v", r)
|
||||||
|
}
|
||||||
|
rel := InitRelationship("rel", RelationType("one_to_many"))
|
||||||
|
if rel.Name != "rel" || rel.Type != "one_to_many" || rel.Properties == nil || rel.GUID == "" {
|
||||||
|
t.Errorf("InitRelationship: %+v", rel)
|
||||||
|
}
|
||||||
|
con := InitConstraint("k", UniqueConstraint)
|
||||||
|
if con.Name != "k" || con.Type != UniqueConstraint || con.Columns == nil || con.ReferencedColumns == nil || con.GUID == "" {
|
||||||
|
t.Errorf("InitConstraint: %+v", con)
|
||||||
|
}
|
||||||
|
sc := InitScript("sc")
|
||||||
|
if sc.Name != "sc" || sc.RunAfter == nil || sc.Metadata == nil || sc.GUID == "" {
|
||||||
|
t.Errorf("InitScript: %+v", sc)
|
||||||
|
}
|
||||||
|
v := InitView("v", "s")
|
||||||
|
if v.Name != "v" || v.Schema != "s" || v.Columns == nil || v.Metadata == nil || v.GUID == "" {
|
||||||
|
t.Errorf("InitView: %+v", v)
|
||||||
|
}
|
||||||
|
sq := InitSequence("sq", "s")
|
||||||
|
if sq.Name != "sq" || sq.Schema != "s" || sq.IncrementBy != 1 || sq.StartValue != 1 || sq.GUID == "" {
|
||||||
|
t.Errorf("InitSequence: %+v", sq)
|
||||||
|
}
|
||||||
|
d := InitDomain("d")
|
||||||
|
if d.Name != "d" || d.Tables == nil || d.Metadata == nil || d.GUID == "" {
|
||||||
|
t.Errorf("InitDomain: %+v", d)
|
||||||
|
}
|
||||||
|
dt := InitDomainTable("t", "s")
|
||||||
|
if dt.TableName != "t" || dt.SchemaName != "s" || dt.GUID == "" {
|
||||||
|
t.Errorf("InitDomainTable: %+v", dt)
|
||||||
|
}
|
||||||
|
e := InitEnum("e", "s")
|
||||||
|
if e.Name != "e" || e.Schema != "s" || e.Values == nil || e.GUID == "" {
|
||||||
|
t.Errorf("InitEnum: %+v", e)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GUIDs are unique per call.
|
||||||
|
if InitTable("t", "s").GUID == InitTable("t", "s").GUID {
|
||||||
|
t.Error("GUIDs must be unique")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
package models
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
type sortCase struct {
|
||||||
|
name string
|
||||||
|
seq uint
|
||||||
|
}
|
||||||
|
|
||||||
|
var sortFixture = []sortCase{{"Banana", 3}, {"apple", 1}, {"Cherry", 2}}
|
||||||
|
|
||||||
|
var (
|
||||||
|
wantNameAsc = []string{"apple", "Banana", "Cherry"}
|
||||||
|
wantNameDesc = []string{"Cherry", "Banana", "apple"}
|
||||||
|
wantSeqAsc = []string{"apple", "Cherry", "Banana"}
|
||||||
|
wantSeqDesc = []string{"Banana", "Cherry", "apple"}
|
||||||
|
)
|
||||||
|
|
||||||
|
func checkNames(t *testing.T, label string, got, want []string) {
|
||||||
|
t.Helper()
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("%s: got %v, want %v", label, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// runSortSuite exercises a by-name and by-sequence sorter pair over the shared fixture.
|
||||||
|
func runSortSuite[T any](t *testing.T, build func(sortCase) T, name func(T) string,
|
||||||
|
byName func([]T, bool) error, bySeq func([]T, bool) error) {
|
||||||
|
t.Helper()
|
||||||
|
mk := func() []T {
|
||||||
|
out := make([]T, 0, len(sortFixture))
|
||||||
|
for _, c := range sortFixture {
|
||||||
|
out = append(out, build(c))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
names := func(items []T) []string {
|
||||||
|
out := make([]string, 0, len(items))
|
||||||
|
for _, it := range items {
|
||||||
|
out = append(out, name(it))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
if byName != nil {
|
||||||
|
items := mk()
|
||||||
|
_ = byName(items, false)
|
||||||
|
checkNames(t, "name asc", names(items), wantNameAsc)
|
||||||
|
_ = byName(items, true)
|
||||||
|
checkNames(t, "name desc", names(items), wantNameDesc)
|
||||||
|
_ = byName(nil, false)
|
||||||
|
_ = byName([]T{}, true)
|
||||||
|
}
|
||||||
|
if bySeq != nil {
|
||||||
|
items := mk()
|
||||||
|
_ = bySeq(items, false)
|
||||||
|
checkNames(t, "seq asc", names(items), wantSeqAsc)
|
||||||
|
_ = bySeq(items, true)
|
||||||
|
checkNames(t, "seq desc", names(items), wantSeqDesc)
|
||||||
|
_ = bySeq(nil, false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortSchemas(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Schema { return &Schema{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *Schema) string { return s.Name }, SortSchemasByName, SortSchemasBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortTables(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Table { return &Table{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *Table) string { return s.Name }, SortTablesByName, SortTablesBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortColumns(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Column { return &Column{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *Column) string { return s.Name }, SortColumnsByName, SortColumnsBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortViews(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *View { return &View{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *View) string { return s.Name }, SortViewsByName, SortViewsBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortSequences(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Sequence { return &Sequence{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *Sequence) string { return s.Name }, SortSequencesByName, SortSequencesBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortIndexes(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Index { return &Index{Name: c.name, Sequence: c.seq} },
|
||||||
|
func(s *Index) string { return s.Name }, SortIndexesByName, SortIndexesBySequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortNameOnly(t *testing.T) {
|
||||||
|
runSortSuite(t, func(c sortCase) *Constraint { return &Constraint{Name: c.name} },
|
||||||
|
func(s *Constraint) string { return s.Name }, SortConstraintsByName, nil)
|
||||||
|
runSortSuite(t, func(c sortCase) *Relationship { return &Relationship{Name: c.name} },
|
||||||
|
func(s *Relationship) string { return s.Name }, SortRelationshipsByName, nil)
|
||||||
|
runSortSuite(t, func(c sortCase) *Script { return &Script{Name: c.name} },
|
||||||
|
func(s *Script) string { return s.Name }, SortScriptsByName, nil)
|
||||||
|
runSortSuite(t, func(c sortCase) *Enum { return &Enum{Name: c.name} },
|
||||||
|
func(s *Enum) string { return s.Name }, SortEnumsByName, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortStableForTies(t *testing.T) {
|
||||||
|
cols := []*Column{{Name: "x", Description: "first"}, {Name: "X", Description: "second"}, {Name: "x", Description: "third"}}
|
||||||
|
_ = SortColumnsByName(cols, false)
|
||||||
|
if cols[0].Description != "first" || cols[1].Description != "second" || cols[2].Description != "third" {
|
||||||
|
t.Errorf("ties must keep input order: %v %v %v", cols[0].Description, cols[1].Description, cols[2].Description)
|
||||||
|
}
|
||||||
|
_ = SortColumnsBySequence(cols, true)
|
||||||
|
if cols[0].Description != "first" || cols[2].Description != "third" {
|
||||||
|
t.Errorf("sequence ties must keep input order")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortMapVariants(t *testing.T) {
|
||||||
|
cols := map[string]*Column{}
|
||||||
|
idx := map[string]*Index{}
|
||||||
|
cons := map[string]*Constraint{}
|
||||||
|
rels := map[string]*Relationship{}
|
||||||
|
for _, c := range sortFixture {
|
||||||
|
cols[c.name] = &Column{Name: c.name, Sequence: c.seq}
|
||||||
|
idx[c.name] = &Index{Name: c.name, Sequence: c.seq}
|
||||||
|
cons[c.name] = &Constraint{Name: c.name}
|
||||||
|
rels[c.name] = &Relationship{Name: c.name}
|
||||||
|
}
|
||||||
|
colNames := func(l []*Column) (o []string) {
|
||||||
|
for _, x := range l {
|
||||||
|
o = append(o, x.Name)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
idxNames := func(l []*Index) (o []string) {
|
||||||
|
for _, x := range l {
|
||||||
|
o = append(o, x.Name)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
conNames := func(l []*Constraint) (o []string) {
|
||||||
|
for _, x := range l {
|
||||||
|
o = append(o, x.Name)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
relNames := func(l []*Relationship) (o []string) {
|
||||||
|
for _, x := range l {
|
||||||
|
o = append(o, x.Name)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
checkNames(t, "cols name", colNames(SortColumnsMapByName(cols, false)), wantNameAsc)
|
||||||
|
checkNames(t, "cols name desc", colNames(SortColumnsMapByName(cols, true)), wantNameDesc)
|
||||||
|
checkNames(t, "cols seq", colNames(SortColumnsMapBySequence(cols, false)), wantSeqAsc)
|
||||||
|
checkNames(t, "cols seq desc", colNames(SortColumnsMapBySequence(cols, true)), wantSeqDesc)
|
||||||
|
checkNames(t, "idx name", idxNames(SortIndexesMapByName(idx, false)), wantNameAsc)
|
||||||
|
checkNames(t, "idx seq", idxNames(SortIndexesMapBySequence(idx, true)), wantSeqDesc)
|
||||||
|
checkNames(t, "con name", conNames(SortConstraintsMapByName(cons, false)), wantNameAsc)
|
||||||
|
checkNames(t, "rel name", relNames(SortRelationshipsMapByName(rels, true)), wantNameDesc)
|
||||||
|
|
||||||
|
if got := SortColumnsMapByName(nil, false); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("nil map must give non-nil empty slice")
|
||||||
|
}
|
||||||
|
if len(cols) != 3 {
|
||||||
|
t.Error("input map must not be modified")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,249 @@
|
|||||||
|
package models
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// viewFixture builds a two-schema database whose map contents would randomise output order.
|
||||||
|
func viewFixture() *Database {
|
||||||
|
db := InitDatabase("shop")
|
||||||
|
db.Description = "desc"
|
||||||
|
db.DatabaseType = PostgresqlDatabaseType
|
||||||
|
db.DatabaseVersion = "16"
|
||||||
|
|
||||||
|
for _, sn := range []string{"sales", "public"} {
|
||||||
|
s := InitSchema(sn)
|
||||||
|
s.Owner = "owner_" + sn
|
||||||
|
s.Scripts = append(s.Scripts, InitScript("seed"))
|
||||||
|
|
||||||
|
users := InitTable("users", sn)
|
||||||
|
for _, cn := range []string{"id", "email", "name"} {
|
||||||
|
c := InitColumn(cn, "users", sn)
|
||||||
|
c.Type = "text"
|
||||||
|
users.Columns[cn] = c
|
||||||
|
}
|
||||||
|
users.Columns["id"].IsPrimaryKey = true
|
||||||
|
users.Columns["id"].NotNull = true
|
||||||
|
|
||||||
|
pk := InitConstraint("users_pkey", PrimaryKeyConstraint)
|
||||||
|
pk.Columns = []string{"id"}
|
||||||
|
users.Constraints["users_pkey"] = pk
|
||||||
|
ck := InitConstraint("users_ck", CheckConstraint)
|
||||||
|
ck.Expression = "id > 0"
|
||||||
|
users.Constraints["users_ck"] = ck
|
||||||
|
users.Indexes["users_idx"] = InitIndex("users_idx", "users", sn)
|
||||||
|
|
||||||
|
orders := InitTable("orders", sn)
|
||||||
|
oid := InitColumn("id", "orders", sn)
|
||||||
|
orders.Columns["id"] = oid
|
||||||
|
uid := InitColumn("user_id", "orders", sn)
|
||||||
|
orders.Columns["user_id"] = uid
|
||||||
|
fk := InitConstraint("orders_user_fk", ForeignKeyConstraint)
|
||||||
|
fk.Columns = []string{"user_id"}
|
||||||
|
fk.ReferencedSchema = sn
|
||||||
|
fk.ReferencedTable = "users"
|
||||||
|
fk.ReferencedColumns = []string{"id"}
|
||||||
|
fk.OnDelete = "CASCADE"
|
||||||
|
orders.Constraints["orders_user_fk"] = fk
|
||||||
|
rel := InitRelationship("orders_users", RelationType("one_to_many"))
|
||||||
|
rel.FromTable, rel.FromSchema = "orders", sn
|
||||||
|
rel.ToTable, rel.ToSchema = "users", sn
|
||||||
|
rel.ForeignKey = "orders_user_fk"
|
||||||
|
rel.ThroughTable, rel.ThroughSchema = "link", sn
|
||||||
|
orders.Relationships["orders_users"] = rel
|
||||||
|
plain := InitRelationship("plain", RelationType("one_to_one"))
|
||||||
|
plain.FromTable, plain.FromSchema = "orders", sn
|
||||||
|
plain.ToTable, plain.ToSchema = "users", sn
|
||||||
|
orders.Relationships["plain"] = plain
|
||||||
|
|
||||||
|
s.Tables = append(s.Tables, users, orders)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToFlatColumns(t *testing.T) {
|
||||||
|
db := viewFixture()
|
||||||
|
first := db.ToFlatColumns()
|
||||||
|
if len(first) != 2*(3+2) {
|
||||||
|
t.Fatalf("got %d columns", len(first))
|
||||||
|
}
|
||||||
|
for i := 1; i < len(first); i++ {
|
||||||
|
if first[i-1].FullyQualifiedName >= first[i].FullyQualifiedName {
|
||||||
|
t.Fatalf("not sorted at %d: %s >= %s", i, first[i-1].FullyQualifiedName, first[i].FullyQualifiedName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if first[0].FullyQualifiedName != "shop.public.orders.id" {
|
||||||
|
t.Errorf("first: %s", first[0].FullyQualifiedName)
|
||||||
|
}
|
||||||
|
var id *FlatColumn
|
||||||
|
for _, c := range first {
|
||||||
|
if c.FullyQualifiedName == "shop.sales.users.id" {
|
||||||
|
id = c
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if id == nil || !id.IsPrimaryKey || !id.NotNull || id.Type != "text" || id.DatabaseName != "shop" || id.SchemaName != "sales" || id.TableName != "users" || id.ColumnName != "id" {
|
||||||
|
t.Errorf("flat id column: %+v", id)
|
||||||
|
}
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
if !reflect.DeepEqual(first, db.ToFlatColumns()) {
|
||||||
|
t.Fatal("ToFlatColumns not deterministic")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := InitDatabase("e").ToFlatColumns(); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("empty db: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToFlatTables(t *testing.T) {
|
||||||
|
got := viewFixture().ToFlatTables()
|
||||||
|
if len(got) != 4 {
|
||||||
|
t.Fatalf("got %d tables", len(got))
|
||||||
|
}
|
||||||
|
// schema order follows the database slice: sales first
|
||||||
|
if got[0].FullyQualifiedName != "shop.sales.users" || got[0].ColumnCount != 3 || got[0].ConstraintCount != 2 || got[0].IndexCount != 1 {
|
||||||
|
t.Errorf("first: %+v", got[0])
|
||||||
|
}
|
||||||
|
if got[1].FullyQualifiedName != "shop.sales.orders" || got[1].ColumnCount != 2 || got[1].ConstraintCount != 1 {
|
||||||
|
t.Errorf("second: %+v", got[1])
|
||||||
|
}
|
||||||
|
if got := InitDatabase("e").ToFlatTables(); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("empty db: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToFlatConstraints(t *testing.T) {
|
||||||
|
db := viewFixture()
|
||||||
|
got := db.ToFlatConstraints()
|
||||||
|
if len(got) != 6 {
|
||||||
|
t.Fatalf("got %d constraints", len(got))
|
||||||
|
}
|
||||||
|
for i := 1; i < len(got); i++ {
|
||||||
|
if got[i-1].FullyQualifiedName >= got[i].FullyQualifiedName {
|
||||||
|
t.Fatalf("not sorted: %s >= %s", got[i-1].FullyQualifiedName, got[i].FullyQualifiedName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var fk, ck *FlatConstraint
|
||||||
|
for _, c := range got {
|
||||||
|
switch c.FullyQualifiedName {
|
||||||
|
case "shop.sales.orders.orders_user_fk":
|
||||||
|
fk = c
|
||||||
|
case "shop.sales.users.users_ck":
|
||||||
|
ck = c
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if fk == nil || fk.ReferencedFQN != "shop.sales.users" || fk.OnDelete != "CASCADE" || fk.Type != ForeignKeyConstraint {
|
||||||
|
t.Errorf("fk: %+v", fk)
|
||||||
|
}
|
||||||
|
if ck == nil || ck.ReferencedFQN != "" || ck.Expression != "id > 0" {
|
||||||
|
t.Errorf("check: %+v", ck)
|
||||||
|
}
|
||||||
|
|
||||||
|
// FK without a referenced table gets no FQN.
|
||||||
|
db2 := InitDatabase("d")
|
||||||
|
s := InitSchema("s")
|
||||||
|
tb := InitTable("t", "s")
|
||||||
|
tb.Constraints["fk"] = InitConstraint("fk", ForeignKeyConstraint)
|
||||||
|
s.Tables = append(s.Tables, tb)
|
||||||
|
db2.Schemas = append(db2.Schemas, s)
|
||||||
|
if out := db2.ToFlatConstraints(); len(out) != 1 || out[0].ReferencedFQN != "" {
|
||||||
|
t.Errorf("unreferenced fk: %+v", out)
|
||||||
|
}
|
||||||
|
if got := InitDatabase("e").ToFlatConstraints(); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("empty db: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToFlatRelationships(t *testing.T) {
|
||||||
|
db := viewFixture()
|
||||||
|
got := db.ToFlatRelationships()
|
||||||
|
if len(got) != 4 {
|
||||||
|
t.Fatalf("got %d relationships", len(got))
|
||||||
|
}
|
||||||
|
for i := 1; i < len(got); i++ {
|
||||||
|
a, b := got[i-1], got[i]
|
||||||
|
if a.FromFQN > b.FromFQN || (a.FromFQN == b.FromFQN && a.RelationshipName > b.RelationshipName) {
|
||||||
|
t.Fatalf("not sorted at %d", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var through, plain *FlatRelationship
|
||||||
|
for _, r := range got {
|
||||||
|
if r.FromSchema == "sales" && r.RelationshipName == "orders_users" {
|
||||||
|
through = r
|
||||||
|
}
|
||||||
|
if r.FromSchema == "sales" && r.RelationshipName == "plain" {
|
||||||
|
plain = r
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if through == nil || through.ThroughTableFQN != "shop.sales.link" || through.FromFQN != "shop.sales.orders" || through.ToFQN != "shop.sales.users" || through.ForeignKey != "orders_user_fk" {
|
||||||
|
t.Errorf("through: %+v", through)
|
||||||
|
}
|
||||||
|
if plain == nil || plain.ThroughTableFQN != "" {
|
||||||
|
t.Errorf("plain: %+v", plain)
|
||||||
|
}
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
if !reflect.DeepEqual(got, db.ToFlatRelationships()) {
|
||||||
|
t.Fatal("ToFlatRelationships not deterministic")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := InitDatabase("e").ToFlatRelationships(); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("empty db: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSummaries(t *testing.T) {
|
||||||
|
db := viewFixture()
|
||||||
|
ds := db.ToSummary()
|
||||||
|
if ds.Name != "shop" || ds.Description != "desc" || ds.DatabaseType != PostgresqlDatabaseType || ds.DatabaseVersion != "16" ||
|
||||||
|
ds.SchemaCount != 2 || ds.TotalTables != 4 || ds.TotalColumns != 10 {
|
||||||
|
t.Errorf("database summary: %+v", ds)
|
||||||
|
}
|
||||||
|
if es := InitDatabase("e").ToSummary(); es.SchemaCount != 0 || es.TotalTables != 0 || es.TotalColumns != 0 {
|
||||||
|
t.Errorf("empty summary: %+v", es)
|
||||||
|
}
|
||||||
|
|
||||||
|
ss := db.Schemas[0].ToSummary()
|
||||||
|
if ss.Name != "sales" || ss.Owner != "owner_sales" || ss.TableCount != 2 || ss.ScriptCount != 1 || ss.TotalColumns != 5 || ss.TotalConstraints != 3 {
|
||||||
|
t.Errorf("schema summary: %+v", ss)
|
||||||
|
}
|
||||||
|
|
||||||
|
users := db.Schemas[0].Tables[0].ToSummary()
|
||||||
|
if users.Name != "users" || users.Schema != "sales" || users.ColumnCount != 3 || users.ConstraintCount != 2 || users.IndexCount != 1 ||
|
||||||
|
users.RelationshipCount != 0 || !users.HasPrimaryKey || users.ForeignKeyCount != 0 {
|
||||||
|
t.Errorf("users summary: %+v", users)
|
||||||
|
}
|
||||||
|
orders := db.Schemas[0].Tables[1].ToSummary()
|
||||||
|
if orders.HasPrimaryKey || orders.ForeignKeyCount != 1 || orders.RelationshipCount != 2 {
|
||||||
|
t.Errorf("orders summary: %+v", orders)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDirectiveFromAny(t *testing.T) {
|
||||||
|
want := Directive{Namespace: "postgres", Key: "partition", Args: "partition by RANGE (x)", Line: 7}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in any
|
||||||
|
want Directive
|
||||||
|
ok bool
|
||||||
|
}{
|
||||||
|
{"directive", want, want, true},
|
||||||
|
{"string map int line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7}, want, true},
|
||||||
|
{"string map int64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": int64(7)}, want, true},
|
||||||
|
{"string map float64 line", map[string]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": float64(7)}, want, true},
|
||||||
|
{"any map ignores non-string keys", map[any]any{"namespace": "postgres", "key": "partition", "args": "partition by RANGE (x)", "line": 7, 5: "x"}, want, true},
|
||||||
|
{"key derived from args", map[string]any{"namespace": "sqlite", "args": "WITHOUT ROWID extra"}, Directive{Namespace: "sqlite", Key: "without", Args: "WITHOUT ROWID extra"}, true},
|
||||||
|
{"wrong field types ignored", map[string]any{"namespace": 1, "key": 2, "args": 3, "line": "x"}, Directive{}, true},
|
||||||
|
{"unsupported type", "nope", Directive{}, false},
|
||||||
|
{"nil", nil, Directive{}, false},
|
||||||
|
{"int", 5, Directive{}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, ok := directiveFromAny(tt.in)
|
||||||
|
if ok != tt.ok || got != tt.want {
|
||||||
|
t.Errorf("got (%+v,%v), want (%+v,%v)", got, ok, tt.want, tt.ok)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,105 @@
|
|||||||
|
package bun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"go/ast"
|
||||||
|
"go/parser"
|
||||||
|
"go/token"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestReader() *Reader { return NewReader(&readers.ReaderOptions{}) }
|
||||||
|
|
||||||
|
func mustExpr(t *testing.T, src string) ast.Expr {
|
||||||
|
t.Helper()
|
||||||
|
e, err := parser.ParseExpr(src)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGoTypeToSQL(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
tests := []struct{ src, want string }{
|
||||||
|
{"int", "integer"}, {"int32", "integer"}, {"int64", "bigint"},
|
||||||
|
{"string", "text"}, {"bool", "boolean"}, {"float32", "real"},
|
||||||
|
{"float64", "double precision"}, {"uint8", "text"},
|
||||||
|
{"time.Time", "timestamp"}, {"time.Duration", "text"},
|
||||||
|
{"sql_types.SqlString", "text"}, {"sql_types.SqlInt", "integer"},
|
||||||
|
{"sql_types.SqlInt64", "bigint"}, {"sql_types.SqlFloat", "double precision"},
|
||||||
|
{"sql_types.SqlBool", "boolean"}, {"sql_types.SqlTime", "timestamp"},
|
||||||
|
{"sql_types.Other", "text"}, {"other.Thing", "text"},
|
||||||
|
{"*int64", "bigint"}, {"*time.Time", "timestamp"}, {"[]byte", "text"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.src, func(t *testing.T) {
|
||||||
|
if got := r.goTypeToSQL(mustExpr(t, tt.src)); got != tt.want {
|
||||||
|
t.Errorf("got %q want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDeriveTableName(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
for in, want := range map[string]string{
|
||||||
|
"ModelUser": "user",
|
||||||
|
"ModelUserRole": "user_role",
|
||||||
|
"Account": "account",
|
||||||
|
"OrderItem": "order_item",
|
||||||
|
} {
|
||||||
|
if got := r.deriveTableName(in); got != want {
|
||||||
|
t.Errorf("%q: got %q want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetReceiverType(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
for src, want := range map[string]string{"User": "User", "*User": "User", "*pkg.User": "", "[]User": ""} {
|
||||||
|
if got := r.getReceiverType(mustExpr(t, src)); got != want {
|
||||||
|
t.Errorf("%s: got %q want %q", src, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetRelationType(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
for tag, want := range map[string]string{
|
||||||
|
`bun:"rel:has-many,join:id=user_id"`: "has-many",
|
||||||
|
`bun:"rel:belongs-to,join:user_id=id"`: "belongs-to",
|
||||||
|
`bun:"rel:has-one,join:id=user_id"`: "has-one",
|
||||||
|
`bun:"rel:many-to-many,join_table:x"`: "many-to-many",
|
||||||
|
`bun:"rel:unknown"`: "",
|
||||||
|
`bun:"id,pk"`: "",
|
||||||
|
} {
|
||||||
|
if got := r.getRelationType(tag); got != want {
|
||||||
|
t.Errorf("%s: got %q want %q", tag, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTableNameMethod(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
parse := func(src string) *ast.FuncDecl {
|
||||||
|
f, err := parser.ParseFile(token.NewFileSet(), "x.go", "package p\n"+src, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return f.Decls[0].(*ast.FuncDecl)
|
||||||
|
}
|
||||||
|
if tbl, sch := r.parseTableNameMethod(parse(`func (User) TableName() string { return "public.users" }`)); tbl != "users" || sch != "public" {
|
||||||
|
t.Errorf("qualified: %q %q", tbl, sch)
|
||||||
|
}
|
||||||
|
if tbl, sch := r.parseTableNameMethod(parse(`func (User) TableName() string { return "users" }`)); tbl != "users" || sch != "public" {
|
||||||
|
t.Errorf("plain: %q %q", tbl, sch)
|
||||||
|
}
|
||||||
|
if tbl, _ := r.parseTableNameMethod(parse(`func (User) TableName() string`)); tbl != "" {
|
||||||
|
t.Errorf("no body: %q", tbl)
|
||||||
|
}
|
||||||
|
if tbl, _ := r.parseTableNameMethod(parse(`func (User) TableName() string { x := 1; _ = x; return foo() }`)); tbl != "" {
|
||||||
|
t.Errorf("non-literal: %q", tbl)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -15,6 +15,9 @@ import (
|
|||||||
// Reader implements the readers.Reader interface for Drizzle schema format
|
// Reader implements the readers.Reader interface for Drizzle schema format
|
||||||
type Reader struct {
|
type Reader struct {
|
||||||
options *readers.ReaderOptions
|
options *readers.ReaderOptions
|
||||||
|
// enumVars maps the constant a pgEnum() is assigned to (e.g. "role") to the
|
||||||
|
// enum's SQL name (e.g. "Role"), so columns declared as role('col') resolve.
|
||||||
|
enumVars map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewReader creates a new Drizzle reader with the given options
|
// NewReader creates a new Drizzle reader with the given options
|
||||||
@@ -29,6 +32,7 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
|
|||||||
if r.options.FilePath == "" {
|
if r.options.FilePath == "" {
|
||||||
return nil, fmt.Errorf("file path is required for Drizzle reader")
|
return nil, fmt.Errorf("file path is required for Drizzle reader")
|
||||||
}
|
}
|
||||||
|
r.enumVars = make(map[string]string)
|
||||||
|
|
||||||
// Check if it's a file or directory
|
// Check if it's a file or directory
|
||||||
info, err := os.Stat(r.options.FilePath)
|
info, err := os.Stat(r.options.FilePath)
|
||||||
@@ -100,6 +104,13 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
|||||||
return nil, fmt.Errorf("failed to glob directory: %w", err)
|
return nil, fmt.Errorf("failed to glob directory: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Enums may be declared in a different file than the tables using them
|
||||||
|
for _, file := range files {
|
||||||
|
if content, err := os.ReadFile(file); err == nil {
|
||||||
|
r.collectEnumVars(string(content))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Parse each file
|
// Parse each file
|
||||||
for _, file := range files {
|
for _, file := range files {
|
||||||
content, err := os.ReadFile(file)
|
content, err := os.ReadFile(file)
|
||||||
@@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
|
|||||||
return db, nil
|
return db, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var enumVarRegex = regexp.MustCompile(`export\s+const\s+(\w+)\s*=\s*pgEnum\s*\(\s*['"](\w+)['"]`)
|
||||||
|
|
||||||
|
// collectEnumVars records every pgEnum() constant declared in content.
|
||||||
|
func (r *Reader) collectEnumVars(content string) {
|
||||||
|
if r.enumVars == nil {
|
||||||
|
r.enumVars = make(map[string]string)
|
||||||
|
}
|
||||||
|
for _, m := range enumVarRegex.FindAllStringSubmatch(content, -1) {
|
||||||
|
r.enumVars[m[1]] = m[2]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// parseDrizzle parses Drizzle schema content and returns a Database model
|
// parseDrizzle parses Drizzle schema content and returns a Database model
|
||||||
func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
|
func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
|
||||||
db := models.InitDatabase("database")
|
db := models.InitDatabase("database")
|
||||||
|
r.collectEnumVars(content)
|
||||||
|
|
||||||
if r.options.Metadata != nil {
|
if r.options.Metadata != nil {
|
||||||
if name, ok := r.options.Metadata["name"].(string); ok {
|
if name, ok := r.options.Metadata["name"].(string); ok {
|
||||||
@@ -375,6 +399,9 @@ func (r *Reader) parseColumnDefinition(line, fieldName, drizzleType string, tabl
|
|||||||
|
|
||||||
// Map Drizzle type to SQL type
|
// Map Drizzle type to SQL type
|
||||||
column.Type = r.drizzleTypeToSQL(drizzleType)
|
column.Type = r.drizzleTypeToSQL(drizzleType)
|
||||||
|
if enumName, ok := r.enumVars[drizzleType]; ok {
|
||||||
|
column.Type = enumName
|
||||||
|
}
|
||||||
|
|
||||||
// Default: columns are nullable unless specified
|
// Default: columns are nullable unless specified
|
||||||
column.NotNull = false
|
column.NotNull = false
|
||||||
|
|||||||
@@ -0,0 +1,114 @@
|
|||||||
|
package drizzle
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const fixture = "../../../tests/assets/drizzle/schema.ts"
|
||||||
|
|
||||||
|
func readFile(t *testing.T, path string) *models.Database {
|
||||||
|
t.Helper()
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func findTable(db *models.Database, name string) *models.Table {
|
||||||
|
for _, s := range db.Schemas {
|
||||||
|
for _, tb := range s.Tables {
|
||||||
|
if tb.Name == name {
|
||||||
|
return tb
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadFixture(t *testing.T) {
|
||||||
|
db := readFile(t, fixture)
|
||||||
|
if len(db.Schemas) == 0 || len(db.Schemas[0].Tables) == 0 {
|
||||||
|
t.Fatal("expected tables")
|
||||||
|
}
|
||||||
|
if len(db.Schemas[0].Enums) != 1 || db.Schemas[0].Enums[0].Name != "Role" {
|
||||||
|
t.Fatalf("enums = %+v", db.Schemas[0].Enums)
|
||||||
|
}
|
||||||
|
var found bool
|
||||||
|
for _, tb := range db.Schemas[0].Tables {
|
||||||
|
if c, ok := tb.Columns["role"]; ok {
|
||||||
|
found = true
|
||||||
|
if c.Type != "Role" {
|
||||||
|
t.Errorf("role type = %q", c.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for n := range tb.Columns {
|
||||||
|
if n == "profile" {
|
||||||
|
t.Errorf("relation field leaked as column in %s", tb.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
t.Error("no role column")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnumColumnSyntax(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
src string
|
||||||
|
}{
|
||||||
|
{"enum constant", "export const role = pgEnum('Role', ['A','B']);\nexport const users = pgTable('users', {\n role: role('role').notNull(),\n});\n"},
|
||||||
|
{"legacy", "export const role = pgEnum('Role', ['A','B']);\nexport const users = pgTable('users', {\n role: pgEnum('Role')('role').notNull(),\n});\n"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
p := filepath.Join(t.TempDir(), "s.ts")
|
||||||
|
if err := os.WriteFile(p, []byte(tt.src), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tb := findTable(readFile(t, p), "users")
|
||||||
|
if tb == nil {
|
||||||
|
t.Fatal("users missing")
|
||||||
|
}
|
||||||
|
c := tb.Columns["role"]
|
||||||
|
if c == nil || c.Type != "Role" || !c.NotNull {
|
||||||
|
t.Errorf("column = %+v", c)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadDirectorySeparateEnums(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
files := map[string]string{
|
||||||
|
"enums.ts": "export const status = pgEnum('Status', ['on','off']);\n",
|
||||||
|
"tables.ts": "export const items = pgTable('items', {\n status: status('status'),\n});\n",
|
||||||
|
}
|
||||||
|
for n, c := range files {
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, n), []byte(c), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tb := findTable(readFile(t, dir), "items")
|
||||||
|
if tb == nil {
|
||||||
|
t.Fatal("items missing")
|
||||||
|
}
|
||||||
|
if c := tb.Columns["status"]; c == nil || c.Type != "Status" {
|
||||||
|
t.Errorf("column = %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReaderErrors(t *testing.T) {
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
|
||||||
|
t.Error("expected error for empty path")
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{FilePath: "/nonexistent.ts"}).ReadDatabase(); err == nil {
|
||||||
|
t.Error("expected error for missing file")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,152 @@
|
|||||||
|
package gorm
|
||||||
|
|
||||||
|
import (
|
||||||
|
"go/ast"
|
||||||
|
"go/parser"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestReader() *Reader { return NewReader(&readers.ReaderOptions{}) }
|
||||||
|
|
||||||
|
func mustExpr(t *testing.T, src string) ast.Expr {
|
||||||
|
t.Helper()
|
||||||
|
e, err := parser.ParseExpr(src)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGoTypeToSQL(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
tests := []struct{ src, want string }{
|
||||||
|
{"int", "integer"}, {"int32", "integer"}, {"int64", "bigint"},
|
||||||
|
{"string", "text"}, {"bool", "boolean"}, {"float32", "real"},
|
||||||
|
{"float64", "double precision"}, {"uint8", "text"},
|
||||||
|
{"time.Time", "timestamp"}, {"time.Duration", "text"},
|
||||||
|
{"sql_types.SqlString", "text"}, {"sql_types.SqlInt", "integer"},
|
||||||
|
{"sql_types.SqlInt64", "bigint"}, {"sql_types.SqlFloat", "double precision"},
|
||||||
|
{"sql_types.SqlBool", "boolean"}, {"sql_types.SqlTime", "timestamp"},
|
||||||
|
{"sql_types.Other", "text"}, {"other.Thing", "text"},
|
||||||
|
{"*int64", "bigint"}, {"*time.Time", "timestamp"}, {"[]byte", "text"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.src, func(t *testing.T) {
|
||||||
|
if got := r.goTypeToSQL(mustExpr(t, tt.src)); got != tt.want {
|
||||||
|
t.Errorf("got %q want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFieldNameToColumnName(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
for in, want := range map[string]string{"ID": "i_d", "UserName": "user_name", "name": "name", "": ""} {
|
||||||
|
if got := r.fieldNameToColumnName(in); got != want {
|
||||||
|
t.Errorf("%q: got %q want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetReceiverType(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
tests := []struct{ src, want string }{
|
||||||
|
{"User", "User"}, {"*User", "User"}, {"*pkg.User", ""}, {"[]User", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := r.getReceiverType(mustExpr(t, tt.src)); got != tt.want {
|
||||||
|
t.Errorf("%s: got %q want %q", tt.src, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsGORMModel(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
field *ast.Field
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"embedded gorm.Model", &ast.Field{Type: mustExpr(t, "gorm.Model")}, true},
|
||||||
|
{"named field", &ast.Field{Names: []*ast.Ident{ast.NewIdent("M")}, Type: mustExpr(t, "gorm.Model")}, false},
|
||||||
|
{"plain ident", &ast.Field{Type: mustExpr(t, "Model")}, false},
|
||||||
|
{"other package", &ast.Field{Type: mustExpr(t, "other.Model")}, false},
|
||||||
|
{"gorm other", &ast.Field{Type: mustExpr(t, "gorm.DB")}, false},
|
||||||
|
{"non-ident selector base", &ast.Field{Type: mustExpr(t, "a.b.Model")}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := r.isGORMModel(tt.field); got != tt.want {
|
||||||
|
t.Errorf("got %v want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseTypeWithReferences(t *testing.T) {
|
||||||
|
r := newTestReader()
|
||||||
|
tests := []struct {
|
||||||
|
in string
|
||||||
|
base string
|
||||||
|
length int
|
||||||
|
refInfo string
|
||||||
|
}{
|
||||||
|
{"bigint", "bigint", 0, ""},
|
||||||
|
{"varchar(50)", "varchar", 50, ""},
|
||||||
|
{"bigint references mainaccount(id) ON DELETE CASCADE", "bigint", 0, "mainaccount(id) ON DELETE CASCADE"},
|
||||||
|
{"varchar(20) REFERENCES t(c)", "varchar", 20, "t(c)"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
base, length, ref := r.parseTypeWithReferences(tt.in)
|
||||||
|
if base != tt.base || length != tt.length || ref != tt.refInfo {
|
||||||
|
t.Errorf("%q: got (%q,%d,%q)", tt.in, base, length, ref)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateInlineReferenceConstraint(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
ref string
|
||||||
|
wantNone bool
|
||||||
|
schema string
|
||||||
|
table string
|
||||||
|
col string
|
||||||
|
onDelete string
|
||||||
|
onUpdate string
|
||||||
|
}{
|
||||||
|
{"simple", "accounts(id)", false, "public", "accounts", "id", "NO ACTION", "NO ACTION"},
|
||||||
|
{"schema qualified", "billing.accounts(id)", false, "billing", "accounts", "id", "NO ACTION", "NO ACTION"},
|
||||||
|
{"cascade restrict", "accounts(id) ON DELETE CASCADE ON UPDATE RESTRICT", false, "public", "accounts", "id", "CASCADE", "RESTRICT"},
|
||||||
|
{"set null no action", "accounts(id) on delete set null on update no action", false, "public", "accounts", "id", "SET NULL", "NO ACTION"},
|
||||||
|
{"restrict delete cascade update", "accounts(id) ON DELETE RESTRICT ON UPDATE CASCADE", false, "public", "accounts", "id", "RESTRICT", "CASCADE"},
|
||||||
|
{"update set null", "accounts(id) ON DELETE NO ACTION ON UPDATE SET NULL", false, "public", "accounts", "id", "NO ACTION", "SET NULL"},
|
||||||
|
{"no parens", "accounts", true, "", "", "", "", ""},
|
||||||
|
{"reversed parens", "accounts)id(", true, "", "", "", "", ""},
|
||||||
|
}
|
||||||
|
r := newTestReader()
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
table := models.InitTable("orders", "public")
|
||||||
|
col := models.InitColumn("account_id", "orders", "public")
|
||||||
|
r.createInlineReferenceConstraint(table, col, tt.ref)
|
||||||
|
if tt.wantNone {
|
||||||
|
if len(table.Constraints) != 0 {
|
||||||
|
t.Fatalf("unexpected constraints: %v", table.Constraints)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c := table.Constraints["fk_orders_account_id"]
|
||||||
|
if c == nil {
|
||||||
|
t.Fatal("constraint missing")
|
||||||
|
}
|
||||||
|
if c.ReferencedSchema != tt.schema || c.ReferencedTable != tt.table ||
|
||||||
|
c.ReferencedColumns[0] != tt.col || c.OnDelete != tt.onDelete || c.OnUpdate != tt.onUpdate {
|
||||||
|
t.Errorf("constraint = %+v", c)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,254 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
_ "github.com/go-sql-driver/mysql"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Reader struct {
|
||||||
|
options *readers.ReaderOptions
|
||||||
|
db *sql.DB
|
||||||
|
ctx context.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewReader(options *readers.ReaderOptions) *Reader {
|
||||||
|
return &Reader{options: options, ctx: context.Background()}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||||
|
if r.options == nil || r.options.ConnectionString == "" {
|
||||||
|
return nil, fmt.Errorf("connection string is required")
|
||||||
|
}
|
||||||
|
if err := r.connect(); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to connect: %w", err)
|
||||||
|
}
|
||||||
|
defer r.close()
|
||||||
|
var name, version string
|
||||||
|
if err := r.db.QueryRowContext(r.ctx, "SELECT DATABASE()").Scan(&name); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get database name: %w", err)
|
||||||
|
}
|
||||||
|
_ = r.db.QueryRowContext(r.ctx, "SELECT VERSION()").Scan(&version)
|
||||||
|
db := models.InitDatabase(name)
|
||||||
|
db.DatabaseType, db.SourceFormat, db.DatabaseVersion = models.MySQLDatabaseType, "mysql", version
|
||||||
|
schemas, err := r.querySchemas(name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to query schemas: %w", err)
|
||||||
|
}
|
||||||
|
for _, schema := range schemas {
|
||||||
|
tables, err := r.queryTables(schema.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
schema.Tables = tables
|
||||||
|
for _, table := range tables {
|
||||||
|
table.Columns, err = r.queryColumns(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.Constraints, err = r.queryConstraints(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.Indexes, err = r.queryIndexes(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.RefSchema = schema
|
||||||
|
for _, c := range table.Constraints {
|
||||||
|
if c.Type == models.ForeignKeyConstraint {
|
||||||
|
r.deriveRelationship(table, c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
schema.RefDatabase = db
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
}
|
||||||
|
return db, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadSchema() (*models.Schema, error) {
|
||||||
|
db, err := r.ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(db.Schemas) == 0 {
|
||||||
|
return nil, fmt.Errorf("no schemas found in database")
|
||||||
|
}
|
||||||
|
return db.Schemas[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadTable() (*models.Table, error) {
|
||||||
|
s, err := r.ReadSchema()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(s.Tables) == 0 {
|
||||||
|
return nil, fmt.Errorf("no tables found in schema")
|
||||||
|
}
|
||||||
|
return s.Tables[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) connect() error {
|
||||||
|
db, err := sql.Open("mysql", r.options.ConnectionString)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err = db.PingContext(r.ctx); err != nil {
|
||||||
|
db.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
r.db = db
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) close() {
|
||||||
|
if r.db != nil {
|
||||||
|
_ = r.db.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func (r *Reader) mapDataType(t string) string { return mariadb.ConvertMariaDBToCanonical(t) }
|
||||||
|
|
||||||
|
func (r *Reader) querySchemas(current string) ([]*models.Schema, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, "SELECT SCHEMA_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME = ?", current)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var out []*models.Schema
|
||||||
|
for rows.Next() {
|
||||||
|
var n string
|
||||||
|
if err := rows.Scan(&n); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out = append(out, models.InitSchema(n))
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryTables(schema string) ([]*models.Table, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, "SELECT TABLE_NAME FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? AND TABLE_TYPE = 'BASE TABLE' ORDER BY TABLE_NAME", schema)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var out []*models.Table
|
||||||
|
for rows.Next() {
|
||||||
|
var n string
|
||||||
|
if err := rows.Scan(&n); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out = append(out, models.InitTable(n, schema))
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryColumns(schema, table string) (map[string]*models.Column, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT COLUMN_NAME, COLUMN_TYPE, IS_NULLABLE, COLUMN_DEFAULT, ORDINAL_POSITION, EXTRA, COLUMN_COMMENT FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? ORDER BY ORDINAL_POSITION`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Column{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, typ, nullable, extra, comment string
|
||||||
|
var def sql.NullString
|
||||||
|
var pos int
|
||||||
|
if err := rows.Scan(&name, &typ, &nullable, &def, &pos, &extra, &comment); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c := models.InitColumn(name, table, schema)
|
||||||
|
c.Type = r.mapDataType(typ)
|
||||||
|
c.NotNull = strings.EqualFold(nullable, "NO")
|
||||||
|
c.Sequence = uint(pos)
|
||||||
|
c.Comment = comment
|
||||||
|
if def.Valid {
|
||||||
|
c.Default = def.String
|
||||||
|
}
|
||||||
|
c.AutoIncrement = strings.Contains(strings.ToLower(extra), "auto_increment")
|
||||||
|
out[name] = c
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryConstraints(schema, table string) (map[string]*models.Constraint, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT CONSTRAINT_NAME, CONSTRAINT_TYPE, COLUMN_NAME, REFERENCED_TABLE_SCHEMA, REFERENCED_TABLE_NAME, REFERENCED_COLUMN_NAME, ORDINAL_POSITION FROM information_schema.KEY_COLUMN_USAGE k JOIN information_schema.TABLE_CONSTRAINTS t USING (CONSTRAINT_SCHEMA, TABLE_NAME, CONSTRAINT_NAME) WHERE k.TABLE_SCHEMA=? AND k.TABLE_NAME=? ORDER BY CONSTRAINT_NAME, ORDINAL_POSITION`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Constraint{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, typ, col, rs, rt, rc string
|
||||||
|
var pos int
|
||||||
|
if err := rows.Scan(&name, &typ, &col, &rs, &rt, &rc, &pos); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c := out[name]
|
||||||
|
if c == nil {
|
||||||
|
ct := models.UniqueConstraint
|
||||||
|
if typ == "PRIMARY KEY" {
|
||||||
|
ct = models.PrimaryKeyConstraint
|
||||||
|
}
|
||||||
|
if typ == "FOREIGN KEY" {
|
||||||
|
ct = models.ForeignKeyConstraint
|
||||||
|
}
|
||||||
|
c = models.InitConstraint(name, ct)
|
||||||
|
c.Schema = schema
|
||||||
|
c.Table = table
|
||||||
|
c.ReferencedSchema = rs
|
||||||
|
c.ReferencedTable = rt
|
||||||
|
out[name] = c
|
||||||
|
}
|
||||||
|
c.Columns = append(c.Columns, col)
|
||||||
|
if rc != "" {
|
||||||
|
c.ReferencedColumns = append(c.ReferencedColumns, rc)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryIndexes(schema, table string) (map[string]*models.Index, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT INDEX_NAME, NON_UNIQUE, COLUMN_NAME, SEQ_IN_INDEX, INDEX_TYPE FROM information_schema.STATISTICS WHERE TABLE_SCHEMA=? AND TABLE_NAME=? AND INDEX_NAME <> 'PRIMARY' ORDER BY INDEX_NAME, SEQ_IN_INDEX`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Index{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, col, typ string
|
||||||
|
var non, seq int
|
||||||
|
if err := rows.Scan(&name, &non, &col, &seq, &typ); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
i := out[name]
|
||||||
|
if i == nil {
|
||||||
|
i = models.InitIndex(name, table, schema)
|
||||||
|
i.Unique = non == 0
|
||||||
|
i.Type = strings.ToLower(typ)
|
||||||
|
out[name] = i
|
||||||
|
}
|
||||||
|
i.Columns = append(i.Columns, col)
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) deriveRelationship(t *models.Table, c *models.Constraint) {
|
||||||
|
n := fmt.Sprintf("%s_to_%s", t.Name, c.ReferencedTable)
|
||||||
|
rel := models.InitRelationship(n, models.OneToMany)
|
||||||
|
rel.FromTable = t.Name
|
||||||
|
rel.FromSchema = t.Schema
|
||||||
|
rel.FromColumns = append([]string(nil), c.Columns...)
|
||||||
|
rel.ToTable = c.ReferencedTable
|
||||||
|
rel.ToSchema = c.ReferencedSchema
|
||||||
|
rel.ToColumns = append([]string(nil), c.ReferencedColumns...)
|
||||||
|
rel.ForeignKey = c.Name
|
||||||
|
t.Relationships[n] = rel
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestReaderMapDataType(t *testing.T) {
|
||||||
|
r := NewReader(&readers.ReaderOptions{})
|
||||||
|
for _, tc := range []struct{ input, want string }{{"varchar(64)", "string"}, {"bigint unsigned", "int64"}, {"datetime", "timestamp"}, {"json", "json"}} {
|
||||||
|
if got := r.mapDataType(tc.input); got != tc.want {
|
||||||
|
t.Errorf("mapDataType(%q) = %q, want %q", tc.input, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReaderRequiresConnectionString(t *testing.T) {
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
|
||||||
|
t.Fatal("expected missing connection string error")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
package pgsql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNormalizePostgresDefault(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"empty", "", ""},
|
||||||
|
{"function", "now()", "now()"},
|
||||||
|
{"nextval passthrough", "nextval('seq'::regclass)", "nextval('seq'::regclass)"},
|
||||||
|
{"number", "42", "42"},
|
||||||
|
{"null cast", "NULL::text", "NULL::text"},
|
||||||
|
{"quoted literal", "'abc'", "abc"},
|
||||||
|
{"quoted with cast", "'abc'::character varying", "abc"},
|
||||||
|
{"escaped quote", "'it''s'::text", "it's"},
|
||||||
|
{"empty literal", "''::text", ""},
|
||||||
|
{"only escaped quotes", "''''", "'"},
|
||||||
|
{"unterminated", "'abc", "abc"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := normalizePostgresDefault(tt.in); got != tt.want {
|
||||||
|
t.Errorf("normalizePostgresDefault(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCountHelpers(t *testing.T) {
|
||||||
|
cols := map[string]map[string]*models.Column{
|
||||||
|
"a": {"x": {}, "y": {}},
|
||||||
|
"b": {"z": {}},
|
||||||
|
"c": {},
|
||||||
|
}
|
||||||
|
if got := countColumns(cols); got != 3 {
|
||||||
|
t.Errorf("countColumns = %d, want 3", got)
|
||||||
|
}
|
||||||
|
if got := countColumns(nil); got != 0 {
|
||||||
|
t.Errorf("countColumns(nil) = %d, want 0", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
cons := map[string][]*models.Constraint{"a": {{}, {}}, "b": {{}}}
|
||||||
|
if got := countConstraints(cons); got != 3 {
|
||||||
|
t.Errorf("countConstraints = %d, want 3", got)
|
||||||
|
}
|
||||||
|
if got := countConstraints(nil); got != 0 {
|
||||||
|
t.Errorf("countConstraints(nil) = %d, want 0", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
idx := map[string][]*models.Index{"a": {{}}, "b": {{}, {}, {}}}
|
||||||
|
if got := countIndexes(idx); got != 4 {
|
||||||
|
t.Errorf("countIndexes = %d, want 4", got)
|
||||||
|
}
|
||||||
|
if got := countIndexes(nil); got != 0 {
|
||||||
|
t.Errorf("countIndexes(nil) = %d, want 0", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractIndexOperatorClass(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in []string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"none", nil, ""},
|
||||||
|
{"sort modifiers only", []string{"DESC", "NULLS", "LAST"}, ""},
|
||||||
|
{"opclass", []string{"", " Vector_Cosine_Ops "}, "vector_cosine_ops"},
|
||||||
|
{"opclass after ordering", []string{"desc", "gin_trgm_ops"}, "gin_trgm_ops"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := extractIndexOperatorClass(tt.in); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildIndexHint(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
opClass, params, want string
|
||||||
|
}{
|
||||||
|
{"", "", ""},
|
||||||
|
{"vector_cosine_ops", "", "opclass=vector_cosine_ops"},
|
||||||
|
{"", "m=16", "with (m=16)"},
|
||||||
|
{"vector_cosine_ops", "m=16", "opclass=vector_cosine_ops; with (m=16)"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := buildIndexHint(tt.opClass, tt.params); got != tt.want {
|
||||||
|
t.Errorf("buildIndexHint(%q,%q) = %q, want %q", tt.opClass, tt.params, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeIndexStorageParams(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"m='16', ef_construction='64'", "m=16, ef_construction=64"},
|
||||||
|
{"key_field='id'", "key_field='id'"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := normalizeIndexStorageParams(tt.in); got != tt.want {
|
||||||
|
t.Errorf("normalizeIndexStorageParams(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
// Reader implements the readers.Reader interface for Prisma schema format
|
// Reader implements the readers.Reader interface for Prisma schema format
|
||||||
type Reader struct {
|
type Reader struct {
|
||||||
options *readers.ReaderOptions
|
options *readers.ReaderOptions
|
||||||
|
enumNames map[string]bool // enum names declared in the schema being parsed
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewReader creates a new Prisma reader with the given options
|
// NewReader creates a new Prisma reader with the given options
|
||||||
@@ -82,6 +83,8 @@ func (r *Reader) parsePrisma(content string) (*models.Database, error) {
|
|||||||
schema := models.InitSchema("public")
|
schema := models.InitSchema("public")
|
||||||
schema.Enums = make([]*models.Enum, 0)
|
schema.Enums = make([]*models.Enum, 0)
|
||||||
|
|
||||||
|
r.enumNames = collectEnumNames(content)
|
||||||
|
|
||||||
scanner := bufio.NewScanner(strings.NewReader(content))
|
scanner := bufio.NewScanner(strings.NewReader(content))
|
||||||
|
|
||||||
// State tracking
|
// State tracking
|
||||||
@@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// isEnumType checks if a type name might be an enum
|
// isEnumType reports whether typeName is an enum declared in the schema.
|
||||||
// Note: We can't definitively check against schema.Enums at parse time
|
// Enum names are collected up front because enums may be declared after the
|
||||||
// because enums might be defined after the model, so we just check
|
// models that use them.
|
||||||
// if it starts with uppercase (Prisma convention for enums)
|
func (r *Reader) isEnumType(typeName string, _ *models.Table) bool {
|
||||||
func (r *Reader) isEnumType(typeName string, table *models.Table) bool {
|
return r.enumNames[typeName]
|
||||||
// Simple heuristic: enum types start with uppercase letter
|
}
|
||||||
// and are not known model names (though we can't check that yet)
|
|
||||||
if len(typeName) > 0 && typeName[0] >= 'A' && typeName[0] <= 'Z' {
|
var enumDeclRegex = regexp.MustCompile(`(?m)^\s*enum\s+(\w+)\s*{`)
|
||||||
// Additional check: primitive types are already handled above
|
|
||||||
// So if it's uppercase and not primitive, it's likely an enum or model
|
func collectEnumNames(content string) map[string]bool {
|
||||||
// We'll assume it's an enum if it's a single word
|
names := make(map[string]bool)
|
||||||
return !strings.Contains(typeName, "_")
|
for _, m := range enumDeclRegex.FindAllStringSubmatch(content, -1) {
|
||||||
|
names[m[1]] = true
|
||||||
}
|
}
|
||||||
return false
|
return names
|
||||||
}
|
}
|
||||||
|
|
||||||
// createConstraintFromRelation creates a FK constraint from a @relation attribute
|
// createConstraintFromRelation creates a FK constraint from a @relation attribute
|
||||||
|
|||||||
@@ -0,0 +1,348 @@
|
|||||||
|
package prisma
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const examplePrisma = "../../../tests/assets/prisma/example.prisma"
|
||||||
|
|
||||||
|
func readFixture(t *testing.T) *models.Schema {
|
||||||
|
t.Helper()
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db.Schemas[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func readSource(t *testing.T, src string) *models.Database {
|
||||||
|
t.Helper()
|
||||||
|
p := filepath.Join(t.TempDir(), "schema.prisma")
|
||||||
|
if err := os.WriteFile(p, []byte(src), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: p}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func table(s *models.Schema, name string) *models.Table {
|
||||||
|
for _, tb := range s.Tables {
|
||||||
|
if tb.Name == name {
|
||||||
|
return tb
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFixture_NoRelationFieldColumns(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
// Relation fields (user, author, posts, profile, categories) are not columns.
|
||||||
|
for tbl, fields := range map[string][]string{
|
||||||
|
"User": {"posts", "profile"}, "Profile": {"user"}, "Post": {"author", "categories"}, "Category": {"posts"},
|
||||||
|
} {
|
||||||
|
for _, f := range fields {
|
||||||
|
if _, ok := table(s, tbl).Columns[f]; ok {
|
||||||
|
t.Errorf("%s.%s is a relation field and must not be a column", tbl, f)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Enum-typed fields stay columns.
|
||||||
|
if c := table(s, "User").Columns["role"]; c == nil || c.Type != "Role" || c.Default != "USER" {
|
||||||
|
t.Errorf("User.role: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFixture_Structure(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
if len(s.Enums) != 1 || s.Enums[0].Name != "Role" || len(s.Enums[0].Values) != 2 {
|
||||||
|
t.Errorf("enums: %+v", s.Enums)
|
||||||
|
}
|
||||||
|
for _, n := range []string{"User", "Profile", "Post", "Category", "_CategoryToPost"} {
|
||||||
|
if table(s, n) == nil {
|
||||||
|
t.Errorf("table %s missing", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
user := table(s, "User")
|
||||||
|
if id := user.Columns["id"]; id == nil || !id.IsPrimaryKey || !id.AutoIncrement || id.Type != "integer" {
|
||||||
|
t.Errorf("User.id: %+v", id)
|
||||||
|
}
|
||||||
|
if c := user.Columns["name"]; c == nil || c.NotNull {
|
||||||
|
t.Errorf("optional name: %+v", c)
|
||||||
|
}
|
||||||
|
if uq := user.Constraints["uq_email"]; uq == nil || uq.Columns[0] != "email" {
|
||||||
|
t.Errorf("unique: %+v", user.Constraints)
|
||||||
|
}
|
||||||
|
|
||||||
|
post := table(s, "Post")
|
||||||
|
if c := post.Columns["createdAt"]; c == nil || c.Type != "timestamp" || c.Default != "now()" {
|
||||||
|
t.Errorf("createdAt: %+v", c)
|
||||||
|
}
|
||||||
|
if c := post.Columns["updatedAt"]; c == nil || !strings.Contains(c.Comment, "@updatedAt") {
|
||||||
|
t.Errorf("updatedAt: %+v", c)
|
||||||
|
}
|
||||||
|
if c := post.Columns["published"]; c == nil || c.Default != false {
|
||||||
|
t.Errorf("published default: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFixture_Relations(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
|
||||||
|
fk := table(s, "Post").Constraints["fk_Post_authorId"]
|
||||||
|
if fk == nil || fk.Type != models.ForeignKeyConstraint || fk.Columns[0] != "authorId" || fk.ReferencedTable != "User" || fk.ReferencedColumns[0] != "id" {
|
||||||
|
t.Errorf("Post.author fk: %+v", fk)
|
||||||
|
}
|
||||||
|
|
||||||
|
jt := table(s, "_CategoryToPost")
|
||||||
|
if len(jt.Columns) != 2 {
|
||||||
|
t.Fatalf("join columns: %v", jt.Columns)
|
||||||
|
}
|
||||||
|
var pk, fks int
|
||||||
|
for _, c := range jt.Constraints {
|
||||||
|
switch c.Type {
|
||||||
|
case models.PrimaryKeyConstraint:
|
||||||
|
pk++
|
||||||
|
case models.ForeignKeyConstraint:
|
||||||
|
fks++
|
||||||
|
if c.OnDelete != "Cascade" {
|
||||||
|
t.Errorf("join fk on delete: %q", c.OnDelete)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if pk != 1 || fks != 2 {
|
||||||
|
t.Errorf("join constraints: pk=%d fks=%d", pk, fks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBlockAttributesAndDefaults(t *testing.T) {
|
||||||
|
db := readSource(t, `datasource db {
|
||||||
|
provider = "mysql"
|
||||||
|
}
|
||||||
|
|
||||||
|
model Membership {
|
||||||
|
userId Int
|
||||||
|
groupId Int
|
||||||
|
role String @default("member")
|
||||||
|
alias String @default('x')
|
||||||
|
score Float @default(1.5)
|
||||||
|
tag String @default(cuid())
|
||||||
|
token String @default(uuid())
|
||||||
|
user User @relation(fields: [userId], references: [id], onDelete: Cascade, onUpdate: Restrict)
|
||||||
|
@@id([userId, groupId])
|
||||||
|
@@unique([userId, role])
|
||||||
|
@@index([groupId])
|
||||||
|
@@map("memberships")
|
||||||
|
}
|
||||||
|
|
||||||
|
model User {
|
||||||
|
id Int @id
|
||||||
|
memberships Membership[]
|
||||||
|
slug String @unique @default(dbgenerated("abc(1)"))
|
||||||
|
}
|
||||||
|
`)
|
||||||
|
if db.DatabaseType != "mysql" {
|
||||||
|
t.Errorf("db type: %q", db.DatabaseType)
|
||||||
|
}
|
||||||
|
m := table(db.Schemas[0], "Membership")
|
||||||
|
|
||||||
|
pk := m.Constraints["pk_Membership"]
|
||||||
|
if pk == nil || len(pk.Columns) != 2 || !m.Columns["userId"].IsPrimaryKey || !m.Columns["groupId"].NotNull {
|
||||||
|
t.Errorf("composite pk: %+v", pk)
|
||||||
|
}
|
||||||
|
if uq := m.Constraints["uq_Membership_userId_role"]; uq == nil || len(uq.Columns) != 2 {
|
||||||
|
t.Errorf("composite unique: %+v", m.Constraints)
|
||||||
|
}
|
||||||
|
if ix := m.Indexes["idx_Membership_groupId"]; ix == nil || ix.Columns[0] != "groupId" {
|
||||||
|
t.Errorf("index: %+v", m.Indexes)
|
||||||
|
}
|
||||||
|
|
||||||
|
checks := map[string]any{"role": "member", "alias": "x", "score": "1.5"}
|
||||||
|
for col, want := range checks {
|
||||||
|
if got := m.Columns[col].Default; got != want {
|
||||||
|
t.Errorf("%s default = %#v, want %#v", col, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if m.Columns["tag"].Comment != "default(cuid())" {
|
||||||
|
t.Errorf("cuid comment: %q", m.Columns["tag"].Comment)
|
||||||
|
}
|
||||||
|
if m.Columns["token"].Default != "gen_random_uuid()" {
|
||||||
|
t.Errorf("uuid default: %v", m.Columns["token"].Default)
|
||||||
|
}
|
||||||
|
if m.Columns["score"].Type != "double precision" {
|
||||||
|
t.Errorf("score type: %s", m.Columns["score"].Type)
|
||||||
|
}
|
||||||
|
|
||||||
|
fk := m.Constraints["fk_Membership_userId"]
|
||||||
|
if fk == nil || fk.OnDelete != "Cascade" || fk.OnUpdate != "Restrict" {
|
||||||
|
t.Errorf("fk actions: %+v", fk)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Default with nested parentheses is extracted whole.
|
||||||
|
if got := table(db.Schemas[0], "User").Columns["slug"].Default; got != `dbgenerated("abc(1)")` {
|
||||||
|
t.Errorf("nested default: %#v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnumDeclaredAfterModel(t *testing.T) {
|
||||||
|
db := readSource(t, `model Account {
|
||||||
|
id Int @id
|
||||||
|
status Status @default(ACTIVE)
|
||||||
|
owner Owner?
|
||||||
|
}
|
||||||
|
|
||||||
|
model Owner {
|
||||||
|
id Int @id
|
||||||
|
}
|
||||||
|
|
||||||
|
enum Status {
|
||||||
|
ACTIVE
|
||||||
|
CLOSED
|
||||||
|
}
|
||||||
|
`)
|
||||||
|
a := table(db.Schemas[0], "Account")
|
||||||
|
if c := a.Columns["status"]; c == nil || c.Type != "Status" {
|
||||||
|
t.Errorf("enum column declared before enum: %+v", c)
|
||||||
|
}
|
||||||
|
if _, ok := a.Columns["owner"]; ok {
|
||||||
|
t.Error("model-typed field must not be a column")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseDatasourceProviders(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
tests := []struct {
|
||||||
|
provider string
|
||||||
|
want models.DatabaseType
|
||||||
|
}{
|
||||||
|
{`"postgresql"`, models.PostgresqlDatabaseType}, {`"postgres"`, models.PostgresqlDatabaseType},
|
||||||
|
{`"mysql"`, "mysql"}, {`"sqlite"`, models.SqlLiteDatabaseType},
|
||||||
|
{`"sqlserver"`, models.MSSQLDatabaseType}, {`"cockroachdb"`, models.PostgresqlDatabaseType},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
r.parseDatasource([]string{" provider = " + tt.provider}, db)
|
||||||
|
if db.DatabaseType != tt.want {
|
||||||
|
t.Errorf("%s -> %q, want %q", tt.provider, db.DatabaseType, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseGenerator(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
lines []string
|
||||||
|
opts *readers.ReaderOptions
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"js client", []string{`provider = "prisma-client-js"`}, &readers.ReaderOptions{}, "prisma"},
|
||||||
|
{"new client", []string{`provider = "prisma-client"`}, &readers.ReaderOptions{}, "prisma7"},
|
||||||
|
{"no provider, flag", []string{`output = "x"`}, &readers.ReaderOptions{Prisma7: true}, "prisma7"},
|
||||||
|
{"no provider, nil options", nil, nil, ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
db.SourceFormat = ""
|
||||||
|
(&Reader{options: tt.opts}).parseGenerator(tt.lines, db)
|
||||||
|
if db.SourceFormat != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", db.SourceFormat, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrisma7FlagWithoutGeneratorBlock(t *testing.T) {
|
||||||
|
p := filepath.Join(t.TempDir(), "s.prisma")
|
||||||
|
if err := os.WriteFile(p, []byte("model A {\n id Int @id\n}\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: p, Prisma7: true}).ReadDatabase()
|
||||||
|
if err != nil || db.SourceFormat != "prisma7" {
|
||||||
|
t.Errorf("%v %q", err, db.SourceFormat)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMetadataNameAndComments(t *testing.T) {
|
||||||
|
p := filepath.Join(t.TempDir(), "s.prisma")
|
||||||
|
src := "// leading comment\nmodel A {\n // inner comment\n id Int @id\n}\n"
|
||||||
|
if err := os.WriteFile(p, []byte(src), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: p, Metadata: map[string]any{"name": "shop"}}).ReadDatabase()
|
||||||
|
if err != nil || db.Name != "shop" || len(db.Schemas[0].Tables[0].Columns) != 1 {
|
||||||
|
t.Errorf("%v %+v", err, db)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadSchemaAndTable(t *testing.T) {
|
||||||
|
r := NewReader(&readers.ReaderOptions{FilePath: examplePrisma})
|
||||||
|
s, err := r.ReadSchema()
|
||||||
|
if err != nil || s.Name != "public" {
|
||||||
|
t.Fatalf("schema: %v", err)
|
||||||
|
}
|
||||||
|
tbl, err := r.ReadTable()
|
||||||
|
if err != nil || tbl.Name != "User" {
|
||||||
|
t.Fatalf("table: %v %+v", err, tbl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReader_Errors(t *testing.T) {
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||||
|
t.Errorf("empty path: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{FilePath: filepath.Join(t.TempDir(), "x")}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "failed to read file") {
|
||||||
|
t.Errorf("missing file: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadSchema(); err == nil {
|
||||||
|
t.Error("ReadSchema without path")
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadTable(); err == nil {
|
||||||
|
t.Error("ReadTable without path")
|
||||||
|
}
|
||||||
|
empty := filepath.Join(t.TempDir(), "e.prisma")
|
||||||
|
if err := os.WriteFile(empty, []byte("// nothing\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{FilePath: empty}).ReadTable(); err == nil || !strings.Contains(err.Error(), "no tables found") {
|
||||||
|
t.Errorf("ReadTable on empty: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractDefaultValue(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"@id @default(autoincrement())", "autoincrement()"},
|
||||||
|
{`@default("a(b)")`, `"a(b)"`},
|
||||||
|
{"@unique", ""},
|
||||||
|
{"@default(unclosed(", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := r.extractDefaultValue(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrismaTypeToSQL(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
tests := map[string]string{
|
||||||
|
"String": "text", "Boolean": "boolean", "Int": "integer", "BigInt": "bigint",
|
||||||
|
"Float": "double precision", "Decimal": "decimal", "DateTime": "timestamp",
|
||||||
|
"Json": "jsonb", "Bytes": "bytea", "Custom": "Custom",
|
||||||
|
}
|
||||||
|
for in, want := range tests {
|
||||||
|
if got := r.prismaTypeToSQL(in); got != want {
|
||||||
|
t.Errorf("%s = %s, want %s", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,375 @@
|
|||||||
|
package typeorm
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const exampleTS = "../../../tests/assets/typeorm/example.ts"
|
||||||
|
|
||||||
|
func readFixture(t *testing.T) *models.Schema {
|
||||||
|
t.Helper()
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{FilePath: exampleTS}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(db.Schemas) != 1 {
|
||||||
|
t.Fatalf("schemas: %d", len(db.Schemas))
|
||||||
|
}
|
||||||
|
return db.Schemas[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func tableByName(s *models.Schema, name string) *models.Table {
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
if t.Name == name {
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseSource(t *testing.T, src string) *models.Schema {
|
||||||
|
t.Helper()
|
||||||
|
db, err := NewReader(&readers.ReaderOptions{}).parseTypeORM(src)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db.Schemas[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadFixture_Tables(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
for _, name := range []string{"User", "Project", "Task", "Comment", "Tag", "user_project", "tag_task"} {
|
||||||
|
if tableByName(s, name) == nil {
|
||||||
|
t.Errorf("table %q missing", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(s.Tables) != 7 {
|
||||||
|
t.Errorf("tables: %d, want 7 (5 entities + 2 join tables)", len(s.Tables))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadFixture_ColumnsAndKeys(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
|
||||||
|
user := tableByName(s, "User")
|
||||||
|
id := user.Columns["id"]
|
||||||
|
if id == nil || id.Type != "uuid" || !id.IsPrimaryKey || id.Default != "gen_random_uuid()" {
|
||||||
|
t.Errorf("User.id: %+v", id)
|
||||||
|
}
|
||||||
|
if c := user.Columns["createdAt"]; c == nil || c.Type != "timestamp" || c.Default != "now()" {
|
||||||
|
t.Errorf("User.createdAt: %+v", c)
|
||||||
|
}
|
||||||
|
if c := user.Columns["updatedAt"]; c == nil || c.Type != "timestamp" || !strings.Contains(c.Comment, "auto-update") {
|
||||||
|
t.Errorf("User.updatedAt: %+v", c)
|
||||||
|
}
|
||||||
|
if uq := user.Constraints["uq_email"]; uq == nil || uq.Type != models.UniqueConstraint || uq.Columns[0] != "email" {
|
||||||
|
t.Errorf("unique email: %+v", user.Constraints)
|
||||||
|
}
|
||||||
|
if _, ok := user.Columns["ownedProjects"]; ok {
|
||||||
|
t.Error("relation fields must not become columns")
|
||||||
|
}
|
||||||
|
|
||||||
|
project := tableByName(s, "Project")
|
||||||
|
if c := project.Columns["description"]; c == nil || c.NotNull {
|
||||||
|
t.Errorf("nullable description: %+v", c)
|
||||||
|
}
|
||||||
|
if c := project.Columns["status"]; c == nil || c.Default != "active" {
|
||||||
|
t.Errorf("status default: %+v", c)
|
||||||
|
}
|
||||||
|
if c := tableByName(s, "Task").Columns["description"]; c == nil || c.Type != "text" || c.NotNull {
|
||||||
|
t.Errorf("Task.description: %+v", c)
|
||||||
|
}
|
||||||
|
if c := tableByName(s, "Comment").Columns["content"]; c == nil || c.Type != "text" {
|
||||||
|
t.Errorf("shorthand type: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadFixture_Relationships(t *testing.T) {
|
||||||
|
s := readFixture(t)
|
||||||
|
|
||||||
|
fk := tableByName(s, "Project").Constraints["fk_Project_owner"]
|
||||||
|
if fk == nil || fk.Type != models.ForeignKeyConstraint || fk.Columns[0] != "ownerId" || fk.ReferencedTable != "User" {
|
||||||
|
t.Errorf("Project.owner fk: %+v", fk)
|
||||||
|
}
|
||||||
|
if c := tableByName(s, "Project").Columns["ownerId"]; c == nil || c.Type != "uuid" || !c.NotNull {
|
||||||
|
t.Errorf("ownerId column: %+v", c)
|
||||||
|
}
|
||||||
|
// ManyToOne with { nullable: true } produces a nullable FK column.
|
||||||
|
if c := tableByName(s, "Task").Columns["assigneeId"]; c == nil || c.NotNull {
|
||||||
|
t.Errorf("assigneeId must be nullable: %+v", c)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, jt := range []string{"user_project", "tag_task"} {
|
||||||
|
tbl := tableByName(s, jt)
|
||||||
|
if len(tbl.Columns) != 2 {
|
||||||
|
t.Errorf("%s columns: %d", jt, len(tbl.Columns))
|
||||||
|
}
|
||||||
|
pk := 0
|
||||||
|
fks := 0
|
||||||
|
for _, c := range tbl.Constraints {
|
||||||
|
switch c.Type {
|
||||||
|
case models.PrimaryKeyConstraint:
|
||||||
|
pk++
|
||||||
|
if len(c.Columns) != 2 {
|
||||||
|
t.Errorf("%s composite pk: %v", jt, c.Columns)
|
||||||
|
}
|
||||||
|
case models.ForeignKeyConstraint:
|
||||||
|
fks++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if pk != 1 || fks != 2 {
|
||||||
|
t.Errorf("%s: pk=%d fks=%d", jt, pk, fks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadSchemaAndTable(t *testing.T) {
|
||||||
|
r := NewReader(&readers.ReaderOptions{FilePath: exampleTS})
|
||||||
|
s, err := r.ReadSchema()
|
||||||
|
if err != nil || s.Name != "public" {
|
||||||
|
t.Fatalf("schema: %v %+v", err, s)
|
||||||
|
}
|
||||||
|
tbl, err := r.ReadTable()
|
||||||
|
if err != nil || tbl.Name != "User" {
|
||||||
|
t.Fatalf("table: %v %+v", err, tbl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReader_Errors(t *testing.T) {
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||||
|
t.Errorf("empty path: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{FilePath: filepath.Join(t.TempDir(), "x.ts")}).ReadDatabase(); err == nil || !strings.Contains(err.Error(), "failed to read file") {
|
||||||
|
t.Errorf("missing file: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadSchema(); err == nil {
|
||||||
|
t.Error("ReadSchema without path must fail")
|
||||||
|
}
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadTable(); err == nil {
|
||||||
|
t.Error("ReadTable without path must fail")
|
||||||
|
}
|
||||||
|
|
||||||
|
empty := filepath.Join(t.TempDir(), "empty.ts")
|
||||||
|
if err := os.WriteFile(empty, []byte("// nothing here\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
r := NewReader(&readers.ReaderOptions{FilePath: empty})
|
||||||
|
if db, err := r.ReadDatabase(); err != nil || len(db.Schemas[0].Tables) != 0 {
|
||||||
|
t.Errorf("empty file: %v %+v", err, db)
|
||||||
|
}
|
||||||
|
if _, err := r.ReadTable(); err == nil || !strings.Contains(err.Error(), "no tables found") {
|
||||||
|
t.Errorf("ReadTable on empty: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEntityOptions(t *testing.T) {
|
||||||
|
s := parseSource(t, `
|
||||||
|
@Entity({ name: "app_users", schema: "auth", database: "main", engine: "InnoDB" })
|
||||||
|
export class User {
|
||||||
|
@PrimaryGeneratedColumn()
|
||||||
|
id: number;
|
||||||
|
|
||||||
|
@Column({ type: 'varchar', length: 100, nullable: true })
|
||||||
|
login: string;
|
||||||
|
|
||||||
|
@Column({ type: 'numeric', precision: 12, scale: 4 })
|
||||||
|
balance: number;
|
||||||
|
|
||||||
|
@Column({ type: 'boolean' })
|
||||||
|
active: boolean;
|
||||||
|
}
|
||||||
|
|
||||||
|
@Entity('legacy')
|
||||||
|
export class Legacy {
|
||||||
|
@PrimaryGeneratedColumn('increment')
|
||||||
|
id: number;
|
||||||
|
|
||||||
|
@Column('jsonb')
|
||||||
|
payload: any;
|
||||||
|
}
|
||||||
|
`)
|
||||||
|
user := tableByName(s, "app_users")
|
||||||
|
if user == nil || user.Schema != "auth" {
|
||||||
|
t.Fatalf("tables: %+v", s.Tables)
|
||||||
|
}
|
||||||
|
if c := user.Columns["id"]; c == nil || !c.AutoIncrement || c.Type != "integer" {
|
||||||
|
t.Errorf("id: %+v", c)
|
||||||
|
}
|
||||||
|
if c := user.Columns["login"]; c == nil || c.Type != "varchar(100)" || c.Length != 100 || c.NotNull {
|
||||||
|
t.Errorf("login: %+v", c)
|
||||||
|
}
|
||||||
|
if c := user.Columns["balance"]; c == nil || c.Type != "numeric(12,4)" {
|
||||||
|
t.Errorf("balance: %+v", c)
|
||||||
|
}
|
||||||
|
if c := user.Columns["active"]; c == nil || c.Type != "boolean" {
|
||||||
|
t.Errorf("active: %+v", c)
|
||||||
|
}
|
||||||
|
if c := tableByName(s, "Legacy").Columns["payload"]; c == nil || c.Type != "jsonb" {
|
||||||
|
t.Errorf("payload: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestViewEntity(t *testing.T) {
|
||||||
|
s := parseSource(t, `
|
||||||
|
@ViewEntity({
|
||||||
|
name: "active_users",
|
||||||
|
schema: "reporting",
|
||||||
|
expression: `+"`"+`SELECT id, email FROM users WHERE active`+"`"+`
|
||||||
|
})
|
||||||
|
export class ActiveUsers {
|
||||||
|
id: number;
|
||||||
|
email: string;
|
||||||
|
}
|
||||||
|
|
||||||
|
@ViewEntity({ expression: "SELECT 1" })
|
||||||
|
export class OneView {
|
||||||
|
n: number;
|
||||||
|
}
|
||||||
|
`)
|
||||||
|
if len(s.Views) != 2 || len(s.Tables) != 0 {
|
||||||
|
t.Fatalf("views=%d tables=%d", len(s.Views), len(s.Tables))
|
||||||
|
}
|
||||||
|
v := s.Views[0]
|
||||||
|
if v.Name != "active_users" || v.Schema != "reporting" || !strings.Contains(v.Definition, "SELECT id, email FROM users") {
|
||||||
|
t.Errorf("view: %+v", v)
|
||||||
|
}
|
||||||
|
if c := v.Columns["email"]; c == nil || c.Type != "text" {
|
||||||
|
t.Errorf("view column: %+v", v.Columns)
|
||||||
|
}
|
||||||
|
if s.Views[1].Name != "OneView" || s.Views[1].Definition != "SELECT 1" {
|
||||||
|
t.Errorf("second view: %+v", s.Views[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseColumnDecorator_IdentityAndGenerated(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
tbl := models.InitTable("t", "public")
|
||||||
|
|
||||||
|
col := models.InitColumn("id", "t", "public")
|
||||||
|
r.parseColumnDecorator(`@PrimaryGeneratedColumn('identity', { generatedIdentity: 'ALWAYS' })`, col, tbl)
|
||||||
|
if !col.IsPrimaryKey || !col.Identity || !col.AutoIncrement {
|
||||||
|
t.Errorf("identity pk: %+v", col)
|
||||||
|
}
|
||||||
|
|
||||||
|
other := models.InitColumn("seq", "t", "public")
|
||||||
|
r.parseColumnDecorator(`@Generated('identity')`, other, tbl)
|
||||||
|
if !other.Identity || other.IdentityGeneration != "BY DEFAULT" {
|
||||||
|
t.Errorf("@Generated: %+v", other)
|
||||||
|
}
|
||||||
|
r.parseColumnDecorator(`@Generated('uuid')`, models.InitColumn("u", "t", "public"), tbl) // no-op, no panic
|
||||||
|
|
||||||
|
gen := models.InitColumn("full", "t", "public")
|
||||||
|
r.parseColumnOptions(`@Column({ type: 'text', generatedType: 'STORED', asExpression: 'a || \'x\'' })`, gen, tbl)
|
||||||
|
if !gen.Generated || !strings.Contains(gen.GenerationExpression, "a ||") {
|
||||||
|
t.Errorf("generated column: %+v", gen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseGeneratedIdentity(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{`{ generatedIdentity: 'ALWAYS' }`, "ALWAYS"},
|
||||||
|
{`{ generatedIdentity: 'BY DEFAULT' }`, "BY DEFAULT"},
|
||||||
|
{`no option`, "BY DEFAULT"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := parseGeneratedIdentity(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnescapeSingleQuoted(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"plain", "plain"}, {`it\'s`, "it's"}, {`a\\b`, `a\b`}, {`trailing\`, `trailing\`}, {"", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := unescapeSingleQuoted(tt.in); got != tt.want {
|
||||||
|
t.Errorf("unescape(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchDecorator(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
line string
|
||||||
|
want string
|
||||||
|
wantOK bool
|
||||||
|
}{
|
||||||
|
{"@Entity()", "@Entity()", true},
|
||||||
|
{"@Column() name: string;", "@Column()", true},
|
||||||
|
{"@Column({ type: 'text' })", "@Column({ type: 'text' })", true},
|
||||||
|
{`@Column({ asExpression: 'f(a)' }) x: string;`, `@Column({ asExpression: 'f(a)' })`, true},
|
||||||
|
{"@Generated", "@Generated", true},
|
||||||
|
{"@Column({ unterminated", "@Column({ unterminated", true},
|
||||||
|
{"name: string;", "", false},
|
||||||
|
{"", "", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
got, ok := matchDecorator(tt.line)
|
||||||
|
if got != tt.want || ok != tt.wantOK {
|
||||||
|
t.Errorf("matchDecorator(%q) = (%q,%v), want (%q,%v)", tt.line, got, ok, tt.want, tt.wantOK)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTypeScriptTypeToSQL(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"string", "text"}, {"number", "integer"}, {"boolean", "boolean"}, {"Date", "timestamp"},
|
||||||
|
{"any", "jsonb"}, {"string[]", "text"}, {"string | null", "text"}, {"Unknown", "text"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := r.typeScriptTypeToSQL(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsRelationField(t *testing.T) {
|
||||||
|
r := &Reader{}
|
||||||
|
for _, d := range []string{"@ManyToOne(() => A)", "@OneToMany(() => A, a => a.b)", "@ManyToMany(() => A)", "@OneToOne(() => A)"} {
|
||||||
|
if !r.isRelationField(fieldInfo{decorators: []string{d}}) {
|
||||||
|
t.Errorf("%s should be a relation", d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if r.isRelationField(fieldInfo{decorators: []string{"@Column()"}}) || r.isRelationField(fieldInfo{}) {
|
||||||
|
t.Error("non-relation misdetected")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOneToOne_And_MultiLineDecorators(t *testing.T) {
|
||||||
|
s := parseSource(t, `
|
||||||
|
@Entity()
|
||||||
|
export class Profile {
|
||||||
|
@PrimaryGeneratedColumn()
|
||||||
|
id: number;
|
||||||
|
|
||||||
|
@Column({
|
||||||
|
type: 'varchar',
|
||||||
|
length: 50,
|
||||||
|
nullable: true,
|
||||||
|
})
|
||||||
|
bio: string;
|
||||||
|
|
||||||
|
@OneToOne(() => Account)
|
||||||
|
@JoinColumn()
|
||||||
|
account: Account;
|
||||||
|
}
|
||||||
|
|
||||||
|
@Entity()
|
||||||
|
export class Account {
|
||||||
|
@PrimaryGeneratedColumn()
|
||||||
|
id: number;
|
||||||
|
}
|
||||||
|
`)
|
||||||
|
p := tableByName(s, "Profile")
|
||||||
|
if c := p.Columns["bio"]; c == nil || c.Type != "varchar(50)" || c.NotNull {
|
||||||
|
t.Errorf("multi-line @Column not parsed: %+v", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,207 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
"encoding/json"
|
||||||
|
"encoding/xml"
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"gopkg.in/yaml.v3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// arrayPtr is the pointer-receiver surface shared by every nullable array type.
|
||||||
|
type arrayPtr[T any] interface {
|
||||||
|
*T
|
||||||
|
Scan(any) error
|
||||||
|
UnmarshalJSON([]byte) error
|
||||||
|
UnmarshalYAML(*yaml.Node) error
|
||||||
|
UnmarshalXML(*xml.Decoder, xml.StartElement) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// arrayValue is the value-receiver surface shared by every nullable array type.
|
||||||
|
type arrayValue interface {
|
||||||
|
Value() (driver.Value, error)
|
||||||
|
MarshalJSON() ([]byte, error)
|
||||||
|
MarshalYAML() (any, error)
|
||||||
|
MarshalXML(*xml.Encoder, xml.StartElement) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type wrapped[T any] struct {
|
||||||
|
XMLName xml.Name `yaml:"-" xml:"w"`
|
||||||
|
V T `yaml:"v" xml:"v"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// arrayRoundTrip runs the full Scan/Value/JSON/YAML/XML contract for one array type.
|
||||||
|
// badScan is a literal the type's Scan must reject ("" skips the check).
|
||||||
|
func arrayRoundTrip[T any, P arrayPtr[T]](t *testing.T, sample T, null T, badScan string) {
|
||||||
|
t.Helper()
|
||||||
|
sv, ok := any(sample).(arrayValue)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("%T does not implement the array value surface", sample)
|
||||||
|
}
|
||||||
|
nv := any(null).(arrayValue)
|
||||||
|
|
||||||
|
t.Run("scan-value", func(t *testing.T) {
|
||||||
|
val, err := sv.Value()
|
||||||
|
if err != nil || val == nil {
|
||||||
|
t.Fatalf("Value: %v %v", val, err)
|
||||||
|
}
|
||||||
|
for _, in := range []any{val, []byte(val.(string))} {
|
||||||
|
var got T
|
||||||
|
if err := P(&got).Scan(in); err != nil {
|
||||||
|
t.Fatalf("Scan(%T): %v", in, err)
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(got, sample) {
|
||||||
|
t.Errorf("Scan(%T) = %+v, want %+v", in, got, sample)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if v, err := nv.Value(); v != nil || err != nil {
|
||||||
|
t.Errorf("null Value = %v, %v", v, err)
|
||||||
|
}
|
||||||
|
got := sample
|
||||||
|
if err := P(&got).Scan(nil); err != nil || !reflect.DeepEqual(got, null) {
|
||||||
|
t.Errorf("Scan(nil) = %+v, %v", got, err)
|
||||||
|
}
|
||||||
|
if err := P(&got).Scan(12345); err == nil {
|
||||||
|
t.Error("Scan(int) must fail")
|
||||||
|
}
|
||||||
|
if badScan != "" {
|
||||||
|
var bad T
|
||||||
|
if err := P(&bad).Scan(badScan); err == nil {
|
||||||
|
t.Errorf("Scan(%q) must fail", badScan)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("json", func(t *testing.T) {
|
||||||
|
b, err := sv.MarshalJSON()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var got T
|
||||||
|
if err := P(&got).UnmarshalJSON(b); err != nil || !reflect.DeepEqual(got, sample) {
|
||||||
|
t.Errorf("round trip = %+v, %v", got, err)
|
||||||
|
}
|
||||||
|
nb, _ := nv.MarshalJSON()
|
||||||
|
if string(nb) != "null" {
|
||||||
|
t.Errorf("null marshals to %s", nb)
|
||||||
|
}
|
||||||
|
got = sample
|
||||||
|
if err := P(&got).UnmarshalJSON([]byte(" null ")); err != nil || !reflect.DeepEqual(got, null) {
|
||||||
|
t.Errorf("null unmarshal = %+v, %v", got, err)
|
||||||
|
}
|
||||||
|
if err := P(&got).UnmarshalJSON([]byte(`{}`)); err == nil {
|
||||||
|
t.Error("object must be rejected")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("yaml", func(t *testing.T) {
|
||||||
|
b, err := yaml.Marshal(wrapped[T]{V: sample})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var got wrapped[T]
|
||||||
|
if err := yaml.Unmarshal(b, &got); err != nil || !reflect.DeepEqual(got.V, sample) {
|
||||||
|
t.Errorf("round trip = %+v, %v\n%s", got.V, err, b)
|
||||||
|
}
|
||||||
|
nb, err := yaml.Marshal(wrapped[T]{V: null})
|
||||||
|
if err != nil || !strings.Contains(string(nb), "null") {
|
||||||
|
t.Errorf("null marshal = %q, %v", nb, err)
|
||||||
|
}
|
||||||
|
// yaml.v3 skips UnmarshalYAML for null, so decode into a fresh value.
|
||||||
|
got = wrapped[T]{}
|
||||||
|
if err := yaml.Unmarshal(nb, &got); err != nil || !reflect.DeepEqual(got.V, null) {
|
||||||
|
t.Errorf("null unmarshal = %+v, %v", got.V, err)
|
||||||
|
}
|
||||||
|
var bad wrapped[T]
|
||||||
|
if err := yaml.Unmarshal([]byte("v: {a: b}\n"), &bad); err == nil {
|
||||||
|
t.Error("mapping must be rejected")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("xml", func(t *testing.T) {
|
||||||
|
b, err := xml.Marshal(wrapped[T]{V: sample})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var got wrapped[T]
|
||||||
|
if err := xml.Unmarshal(b, &got); err != nil || !reflect.DeepEqual(got.V, sample) {
|
||||||
|
t.Errorf("round trip = %+v, %v\n%s", got.V, err, b)
|
||||||
|
}
|
||||||
|
if _, err := xml.Marshal(wrapped[T]{V: null}); err != nil {
|
||||||
|
t.Errorf("null marshal: %v", err)
|
||||||
|
}
|
||||||
|
var bad wrapped[T]
|
||||||
|
if err := xml.Unmarshal([]byte("<w><v><item>1</item>"), &bad); err == nil {
|
||||||
|
t.Error("truncated xml must fail")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestArrayTypes_FullContract(t *testing.T) {
|
||||||
|
u1, u2 := uuid.New(), uuid.New()
|
||||||
|
t.Run("string", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlStringArray([]string{"a", "b c", `q"uote`, "x,y"}), SqlStringArray{}, "")
|
||||||
|
})
|
||||||
|
t.Run("int16", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlInt16Array([]int16{1, -2, 300}), SqlInt16Array{}, "{99999}")
|
||||||
|
})
|
||||||
|
t.Run("int32", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlInt32Array([]int32{1, -2, 300000}), SqlInt32Array{}, "{x}")
|
||||||
|
})
|
||||||
|
t.Run("int64", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlInt64Array([]int64{1, -2, 1 << 40}), SqlInt64Array{}, "{x}")
|
||||||
|
})
|
||||||
|
t.Run("float32", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlFloat32Array([]float32{1.5, -2.25}), SqlFloat32Array{}, "{x}")
|
||||||
|
})
|
||||||
|
t.Run("float64", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlFloat64Array([]float64{1.5, -2.25, 1e10}), SqlFloat64Array{}, "{x}")
|
||||||
|
})
|
||||||
|
t.Run("bool", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlBoolArray([]bool{true, false, true}), SqlBoolArray{}, "not an array")
|
||||||
|
})
|
||||||
|
t.Run("uuid", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlUUIDArray([]uuid.UUID{u1, u2}), SqlUUIDArray{}, "{not-a-uuid}")
|
||||||
|
})
|
||||||
|
t.Run("vector", func(t *testing.T) {
|
||||||
|
arrayRoundTrip(t, NewSqlVector([]float32{1, 2.5, -3}), SqlVector{}, "1,2,3")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestArrayTypes_EmptyAndMalformedScan(t *testing.T) {
|
||||||
|
var s SqlStringArray
|
||||||
|
if err := s.Scan("{}"); err != nil || !s.Valid || len(s.Val) != 0 {
|
||||||
|
t.Errorf("empty array: %+v %v", s, err)
|
||||||
|
}
|
||||||
|
var i SqlInt32Array
|
||||||
|
if err := i.Scan("{}"); err != nil || !i.Valid || len(i.Val) != 0 {
|
||||||
|
t.Errorf("empty int array: %+v %v", i, err)
|
||||||
|
}
|
||||||
|
var v SqlVector
|
||||||
|
if err := v.Scan("[]"); err != nil || !v.Valid || len(v.Val) != 0 {
|
||||||
|
t.Errorf("empty vector: %+v %v", v, err)
|
||||||
|
}
|
||||||
|
if err := v.Scan("[1,x]"); err == nil {
|
||||||
|
t.Error("bad vector element must fail")
|
||||||
|
}
|
||||||
|
if err := v.Scan(42); err == nil {
|
||||||
|
t.Error("vector Scan(int) must fail")
|
||||||
|
}
|
||||||
|
for _, bad := range []string{"not an array", "{unterminated"} {
|
||||||
|
var a SqlInt32Array
|
||||||
|
if err := a.Scan(bad); err == nil {
|
||||||
|
t.Errorf("Scan(%q) must fail", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestArrayJSONIsPlainSlice(t *testing.T) {
|
||||||
|
b, err := json.Marshal(NewSqlInt32Array([]int32{1, 2}))
|
||||||
|
if err != nil || string(b) != "[1,2]" {
|
||||||
|
t.Errorf("got %s, %v", b, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,157 @@
|
|||||||
|
package sqltypes
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql/driver"
|
||||||
|
"encoding/json"
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSqlNull_ValueScalarCases(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlNull[any]
|
||||||
|
want driver.Value
|
||||||
|
}{
|
||||||
|
{name: "invalid", input: SqlNull[any]{}, want: nil},
|
||||||
|
{name: "integer", input: Null[any](int64(42), true), want: int64(42)},
|
||||||
|
{name: "string", input: Null[any]("hello", true), want: "hello"},
|
||||||
|
{name: "boolean", input: Null[any](true, true), want: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := tt.input.Value()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Value returned error: %v", err)
|
||||||
|
}
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("Value() = %v (%T), want %v (%T)", got, got, tt.want, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_Int64Conversions(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlNull[any]
|
||||||
|
want int64
|
||||||
|
}{
|
||||||
|
{name: "invalid", input: SqlNull[any]{}, want: 0},
|
||||||
|
{name: "signed integer", input: Null[any](int32(-12), true), want: -12},
|
||||||
|
{name: "unsigned integer", input: Null[any](uint16(12), true), want: 12},
|
||||||
|
{name: "float truncates", input: Null[any](float64(12.9), true), want: 12},
|
||||||
|
{name: "numeric string", input: Null[any]("123", true), want: 123},
|
||||||
|
{name: "invalid string", input: Null[any]("not a number", true), want: 0},
|
||||||
|
{name: "true", input: Null[any](true, true), want: 1},
|
||||||
|
{name: "false", input: Null[any](false, true), want: 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := tt.input.Int64(); got != tt.want {
|
||||||
|
t.Errorf("Int64() = %d, want %d", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_Float64Conversions(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input SqlNull[any]
|
||||||
|
want float64
|
||||||
|
}{
|
||||||
|
{name: "invalid", input: SqlNull[any]{}, want: 0},
|
||||||
|
{name: "float", input: Null[any](float32(1.25), true), want: 1.25},
|
||||||
|
{name: "signed integer", input: Null[any](int64(-12), true), want: -12},
|
||||||
|
{name: "unsigned integer", input: Null[any](uint16(12), true), want: 12},
|
||||||
|
{name: "numeric string", input: Null[any]("12.5", true), want: 12.5},
|
||||||
|
{name: "invalid string", input: Null[any]("not a number", true), want: 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := tt.input.Float64(); got != tt.want {
|
||||||
|
t.Errorf("Float64() = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlDate_JSONNullAndInvalid(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
json string
|
||||||
|
valid bool
|
||||||
|
}{
|
||||||
|
{name: "null", json: "null", valid: false},
|
||||||
|
{name: "invalid date", json: `"not-a-date"`, valid: false},
|
||||||
|
{name: "valid date", json: `"2024-01-15"`, valid: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
var got SqlDate
|
||||||
|
if err := json.Unmarshal([]byte(tt.json), &got); err != nil {
|
||||||
|
t.Fatalf("UnmarshalJSON returned error: %v", err)
|
||||||
|
}
|
||||||
|
if got.Valid != tt.valid {
|
||||||
|
t.Errorf("Valid = %v, want %v", got.Valid, tt.valid)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if data, err := json.Marshal(SqlDate{}); err != nil {
|
||||||
|
t.Fatalf("MarshalJSON returned error: %v", err)
|
||||||
|
} else if string(data) != "null" {
|
||||||
|
t.Errorf("MarshalJSON() = %s, want null", data)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlTypeNowConstructors(t *testing.T) {
|
||||||
|
before := time.Now()
|
||||||
|
timestamp := SqlTimeStampNow()
|
||||||
|
date := SqlDateNow()
|
||||||
|
tm := SqlTimeNow()
|
||||||
|
after := time.Now()
|
||||||
|
|
||||||
|
for name, got := range map[string]time.Time{
|
||||||
|
"timestamp": timestamp.Time(),
|
||||||
|
"date": date.Time(),
|
||||||
|
"time": tm.Time(),
|
||||||
|
} {
|
||||||
|
if !got.After(before) && !got.Equal(before) || got.After(after) {
|
||||||
|
t.Errorf("%s constructor returned %v outside [%v, %v]", name, got, before, after)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !timestamp.Valid || !date.Valid || !tm.Valid {
|
||||||
|
t.Fatal("Now constructors must return valid values")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewSqlAndToJSONDT(t *testing.T) {
|
||||||
|
if got := NewSql[int64]("42"); !got.Valid || got.Val != 42 {
|
||||||
|
t.Errorf("NewSql[int64](\"42\") = %#v, want valid 42", got)
|
||||||
|
}
|
||||||
|
if got := NewSql[int64](nil); got.Valid {
|
||||||
|
t.Errorf("NewSql[int64](nil) = %#v, want invalid", got)
|
||||||
|
}
|
||||||
|
if got := NewSqlFloat32(1.5); !got.Valid || got.Val != 1.5 {
|
||||||
|
t.Errorf("NewSqlFloat32(1.5) = %#v, want valid 1.5", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
when := time.Date(2024, 1, 15, 10, 30, 45, 0, time.UTC)
|
||||||
|
if got := ToJSONDT(when); got != "2024-01-15T10:30:45Z" {
|
||||||
|
t.Errorf("ToJSONDT() = %q, want RFC3339 timestamp", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSqlNull_Float64PreservesInfinity(t *testing.T) {
|
||||||
|
got := Null[float64](math.Inf(1), true).Float64()
|
||||||
|
if !math.IsInf(got, 1) {
|
||||||
|
t.Errorf("Float64() = %v, want +Inf", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package transform
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Validation and normalization are currently pass-through stubs; these tests
|
||||||
|
// pin that contract (no error, input returned unchanged).
|
||||||
|
func TestTransformerStubs(t *testing.T) {
|
||||||
|
tr := NewTransformer()
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("nil transformer")
|
||||||
|
}
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
table := models.InitTable("t", "public")
|
||||||
|
|
||||||
|
if err := tr.ValidateDatabase(db); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
if err := tr.ValidateSchema(schema); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
if err := tr.ValidateTable(table); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
if got, err := tr.NormalizeDatabase(db); err != nil || got != db {
|
||||||
|
t.Errorf("NormalizeDatabase = %v, %v", got, err)
|
||||||
|
}
|
||||||
|
if got, err := tr.NormalizeSchema(schema); err != nil || got != schema {
|
||||||
|
t.Errorf("NormalizeSchema = %v, %v", got, err)
|
||||||
|
}
|
||||||
|
if got, err := tr.NormalizeTable(table); err != nil || got != table {
|
||||||
|
t.Errorf("NormalizeTable = %v, %v", got, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,186 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ConnKind identifies the database type a connection string targets.
|
||||||
|
type ConnKind string
|
||||||
|
|
||||||
|
const (
|
||||||
|
ConnPostgres ConnKind = "postgres"
|
||||||
|
ConnMSSQL ConnKind = "mssql"
|
||||||
|
ConnSQLite ConnKind = "sqlite"
|
||||||
|
)
|
||||||
|
|
||||||
|
// connKinds lists the kinds offered by the builder dialog, in display order.
|
||||||
|
var connKinds = []ConnKind{ConnPostgres, ConnMSSQL, ConnSQLite}
|
||||||
|
|
||||||
|
// maskedPassword is substituted for the password in previews.
|
||||||
|
const maskedPassword = "****"
|
||||||
|
|
||||||
|
// ConnFields holds the editable parts of a connection string.
|
||||||
|
type ConnFields struct {
|
||||||
|
Kind ConnKind
|
||||||
|
Host string
|
||||||
|
Port string
|
||||||
|
Database string
|
||||||
|
User string
|
||||||
|
Password string
|
||||||
|
SSLMode string
|
||||||
|
FilePath string // SQLite only
|
||||||
|
|
||||||
|
// Extra keeps query parameters the builder has no field for, so that
|
||||||
|
// parsing and rebuilding an existing string does not drop them.
|
||||||
|
Extra url.Values
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultConnFields returns sensible defaults for the given kind.
|
||||||
|
func DefaultConnFields(kind ConnKind) ConnFields {
|
||||||
|
f := ConnFields{Kind: kind}
|
||||||
|
switch kind {
|
||||||
|
case ConnPostgres:
|
||||||
|
f.Host, f.Port, f.User, f.SSLMode = "localhost", "5432", "postgres", "disable"
|
||||||
|
case ConnMSSQL:
|
||||||
|
f.Host, f.Port, f.User, f.SSLMode = "localhost", "1433", "sa", "disable"
|
||||||
|
}
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSLModes returns the valid SSL/encryption options for a kind.
|
||||||
|
func SSLModes(kind ConnKind) []string {
|
||||||
|
switch kind {
|
||||||
|
case ConnPostgres:
|
||||||
|
return []string{"disable", "allow", "prefer", "require", "verify-ca", "verify-full"}
|
||||||
|
case ConnMSSQL:
|
||||||
|
return []string{"disable", "false", "true"}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f ConnFields) sslParam() string {
|
||||||
|
if f.Kind == ConnMSSQL {
|
||||||
|
return "encrypt"
|
||||||
|
}
|
||||||
|
return "sslmode"
|
||||||
|
}
|
||||||
|
|
||||||
|
// BuildConnString renders the fields as a connection string. With mask set,
|
||||||
|
// a non-empty password is replaced by asterisks (for previews).
|
||||||
|
func BuildConnString(f ConnFields, mask bool) string {
|
||||||
|
if f.Kind == ConnSQLite {
|
||||||
|
return f.FilePath
|
||||||
|
}
|
||||||
|
|
||||||
|
u := &url.URL{Scheme: "postgres"}
|
||||||
|
if f.Kind == ConnMSSQL {
|
||||||
|
u.Scheme = "sqlserver"
|
||||||
|
}
|
||||||
|
|
||||||
|
if f.Port != "" {
|
||||||
|
u.Host = net.JoinHostPort(f.Host, f.Port)
|
||||||
|
} else {
|
||||||
|
u.Host = f.Host
|
||||||
|
}
|
||||||
|
|
||||||
|
if f.User != "" {
|
||||||
|
if f.Password != "" {
|
||||||
|
pw := f.Password
|
||||||
|
if mask {
|
||||||
|
pw = maskedPassword
|
||||||
|
}
|
||||||
|
u.User = url.UserPassword(f.User, pw)
|
||||||
|
} else {
|
||||||
|
u.User = url.User(f.User)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
query := url.Values{}
|
||||||
|
for k, v := range f.Extra {
|
||||||
|
query[k] = v
|
||||||
|
}
|
||||||
|
if f.Kind == ConnMSSQL {
|
||||||
|
if f.Database != "" {
|
||||||
|
query.Set("database", f.Database)
|
||||||
|
}
|
||||||
|
} else if f.Database != "" {
|
||||||
|
u.Path = "/" + f.Database
|
||||||
|
}
|
||||||
|
if f.SSLMode != "" {
|
||||||
|
query.Set(f.sslParam(), f.SSLMode)
|
||||||
|
}
|
||||||
|
u.RawQuery = query.Encode()
|
||||||
|
|
||||||
|
out := u.String()
|
||||||
|
if mask {
|
||||||
|
// url escapes '*' in the userinfo; keep the preview readable.
|
||||||
|
out = strings.Replace(out, url.QueryEscape(maskedPassword), maskedPassword, 1)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// DetectConnKind guesses the kind from a connection string's scheme. Anything
|
||||||
|
// that is not a recognised URL is treated as a SQLite file path.
|
||||||
|
func DetectConnKind(s string) ConnKind {
|
||||||
|
lower := strings.ToLower(strings.TrimSpace(s))
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(lower, "postgres://"), strings.HasPrefix(lower, "postgresql://"):
|
||||||
|
return ConnPostgres
|
||||||
|
case strings.HasPrefix(lower, "sqlserver://"), strings.HasPrefix(lower, "mssql://"):
|
||||||
|
return ConnMSSQL
|
||||||
|
}
|
||||||
|
return ConnSQLite
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseConnString splits a connection string into fields. An empty string
|
||||||
|
// yields the defaults for hint. Missing ports fall back to the kind default.
|
||||||
|
func ParseConnString(s string, hint ConnKind) (ConnFields, error) {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if s == "" {
|
||||||
|
return DefaultConnFields(hint), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
kind := DetectConnKind(s)
|
||||||
|
if kind == ConnSQLite {
|
||||||
|
path := s
|
||||||
|
for _, prefix := range []string{"sqlite://", "sqlite3://"} {
|
||||||
|
path = strings.TrimPrefix(path, prefix)
|
||||||
|
}
|
||||||
|
return ConnFields{Kind: ConnSQLite, FilePath: path}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
u, err := url.Parse(s)
|
||||||
|
if err != nil {
|
||||||
|
return DefaultConnFields(kind), fmt.Errorf("invalid connection string: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
f := ConnFields{
|
||||||
|
Kind: kind,
|
||||||
|
Host: u.Hostname(),
|
||||||
|
Port: u.Port(),
|
||||||
|
}
|
||||||
|
if f.Port == "" {
|
||||||
|
f.Port = DefaultConnFields(kind).Port
|
||||||
|
}
|
||||||
|
if u.User != nil {
|
||||||
|
f.User = u.User.Username()
|
||||||
|
f.Password, _ = u.User.Password()
|
||||||
|
}
|
||||||
|
|
||||||
|
query := u.Query()
|
||||||
|
if kind == ConnMSSQL {
|
||||||
|
f.Database = query.Get("database")
|
||||||
|
query.Del("database")
|
||||||
|
} else {
|
||||||
|
f.Database = strings.TrimPrefix(u.Path, "/")
|
||||||
|
}
|
||||||
|
f.SSLMode = query.Get(f.sslParam())
|
||||||
|
query.Del(f.sslParam())
|
||||||
|
if len(query) > 0 {
|
||||||
|
f.Extra = query
|
||||||
|
}
|
||||||
|
return f, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
_ "github.com/microsoft/go-mssqldb"
|
||||||
|
_ "modernc.org/sqlite"
|
||||||
|
)
|
||||||
|
|
||||||
|
// connTestTimeout bounds how long "Test connection" may block.
|
||||||
|
const connTestTimeout = 5 * time.Second
|
||||||
|
|
||||||
|
// TestConnection opens and pings the database described by f. Any occurrence
|
||||||
|
// of the password in the returned error is masked.
|
||||||
|
func TestConnection(f ConnFields) error {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), connTestTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
err := testConnection(ctx, f)
|
||||||
|
if err != nil && f.Password != "" {
|
||||||
|
err = fmt.Errorf("%s", strings.ReplaceAll(err.Error(), f.Password, maskedPassword))
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func testConnection(ctx context.Context, f ConnFields) error {
|
||||||
|
switch f.Kind {
|
||||||
|
case ConnPostgres:
|
||||||
|
conn, err := pgx.Connect(ctx, BuildConnString(f, false))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return conn.Close(ctx)
|
||||||
|
case ConnMSSQL:
|
||||||
|
return pingSQL(ctx, "sqlserver", BuildConnString(f, false))
|
||||||
|
case ConnSQLite:
|
||||||
|
if f.FilePath == "" {
|
||||||
|
return fmt.Errorf("file path is required")
|
||||||
|
}
|
||||||
|
// Opening a missing SQLite file would silently create it.
|
||||||
|
if _, err := os.Stat(f.FilePath); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return pingSQL(ctx, "sqlite", f.FilePath)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("unsupported connection type %q", f.Kind)
|
||||||
|
}
|
||||||
|
|
||||||
|
func pingSQL(ctx context.Context, driver, dsn string) error {
|
||||||
|
db, err := sql.Open(driver, dsn)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
return db.PingContext(ctx)
|
||||||
|
}
|
||||||
@@ -0,0 +1,210 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gdamore/tcell/v2"
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
)
|
||||||
|
|
||||||
|
// connBuilderPage is the page name of the connection string builder dialog.
|
||||||
|
const connBuilderPage = "conn-builder"
|
||||||
|
|
||||||
|
// showConnStringBuilder opens the connection string builder, pre-filled by
|
||||||
|
// parsing current. Save calls onDone with the built string; Esc/Back leaves
|
||||||
|
// the caller's input untouched.
|
||||||
|
func (se *SchemaEditor) showConnStringBuilder(current string, hint ConnKind, returnPage string, onDone func(connString string)) {
|
||||||
|
fields, err := ParseConnString(current, hint)
|
||||||
|
if err != nil {
|
||||||
|
se.showErrorDialog("Error", err.Error()+"\nStarting from defaults.")
|
||||||
|
}
|
||||||
|
|
||||||
|
title := tview.NewTextView().
|
||||||
|
SetText("[::b]Connection String Builder").
|
||||||
|
SetTextAlign(tview.AlignCenter).
|
||||||
|
SetDynamicColors(true)
|
||||||
|
|
||||||
|
preview := tview.NewTextView()
|
||||||
|
preview.SetBorder(true).SetTitle(" Preview (password masked) ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
|
||||||
|
form := tview.NewForm()
|
||||||
|
form.SetBorder(true).SetTitle(" Connection ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
|
||||||
|
updatePreview := func() {
|
||||||
|
preview.SetText(tview.Escape(BuildConnString(fields, true)))
|
||||||
|
}
|
||||||
|
|
||||||
|
closeBuilder := func() {
|
||||||
|
se.pages.RemovePage(connBuilderPage)
|
||||||
|
se.pages.SwitchToPage(returnPage)
|
||||||
|
}
|
||||||
|
|
||||||
|
var render func(focus int)
|
||||||
|
render = func(focus int) {
|
||||||
|
form.Clear(false)
|
||||||
|
|
||||||
|
kindIndex := 0
|
||||||
|
kindLabels := make([]string, len(connKinds))
|
||||||
|
for i, k := range connKinds {
|
||||||
|
kindLabels[i] = string(k)
|
||||||
|
if k == fields.Kind {
|
||||||
|
kindIndex = i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
form.AddDropDown("Type", kindLabels, kindIndex, func(_ string, index int) {
|
||||||
|
if connKinds[index] == fields.Kind {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fields = DefaultConnFields(connKinds[index])
|
||||||
|
render(0)
|
||||||
|
})
|
||||||
|
|
||||||
|
if fields.Kind == ConnSQLite {
|
||||||
|
form.AddInputField("File Path", fields.FilePath, 50, nil, func(v string) {
|
||||||
|
fields.FilePath = v
|
||||||
|
updatePreview()
|
||||||
|
})
|
||||||
|
if item, ok := form.GetFormItemByLabel("File Path").(*tview.InputField); ok {
|
||||||
|
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() != tcell.KeyEnter {
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
se.showFileBrowser(FileBrowserConfig{
|
||||||
|
Mode: FileBrowserLoad,
|
||||||
|
StartPath: fields.FilePath,
|
||||||
|
Extensions: FormatExtensions("sqlite"),
|
||||||
|
ReturnPage: connBuilderPage,
|
||||||
|
OnSelect: func(path string) { item.SetText(path) },
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
form.AddInputField("Host", fields.Host, 50, nil, func(v string) { fields.Host = v; updatePreview() })
|
||||||
|
form.AddInputField("Port", fields.Port, 10, tview.InputFieldInteger, func(v string) { fields.Port = v; updatePreview() })
|
||||||
|
form.AddInputField("Database", fields.Database, 50, nil, func(v string) { fields.Database = v; updatePreview() })
|
||||||
|
form.AddInputField("User", fields.User, 50, nil, func(v string) { fields.User = v; updatePreview() })
|
||||||
|
form.AddPasswordField("Password", fields.Password, 50, '*', func(v string) { fields.Password = v; updatePreview() })
|
||||||
|
|
||||||
|
label := "SSL Mode"
|
||||||
|
if fields.Kind == ConnMSSQL {
|
||||||
|
label = "Encrypt"
|
||||||
|
}
|
||||||
|
modes := SSLModes(fields.Kind)
|
||||||
|
modeIndex := -1
|
||||||
|
for i, m := range modes {
|
||||||
|
if m == fields.SSLMode {
|
||||||
|
modeIndex = i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if modeIndex < 0 {
|
||||||
|
// Keep a value parsed from an existing string even if it is not a listed option.
|
||||||
|
modes = append([]string{fields.SSLMode}, modes...)
|
||||||
|
modeIndex = 0
|
||||||
|
}
|
||||||
|
form.AddDropDown(label, modes, modeIndex, func(option string, _ int) {
|
||||||
|
fields.SSLMode = option
|
||||||
|
updatePreview()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
form.AddButton("Save [F2]", connBuilderSave(se, &fields, closeBuilder, onDone))
|
||||||
|
form.AddButton("Test [F3]", func() { se.testConnectionDialog(fields) })
|
||||||
|
form.AddButton("Back [Esc]", closeBuilder)
|
||||||
|
|
||||||
|
updatePreview()
|
||||||
|
form.SetFocus(focus)
|
||||||
|
se.app.SetFocus(form)
|
||||||
|
}
|
||||||
|
|
||||||
|
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyEscape:
|
||||||
|
closeBuilder()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyF2:
|
||||||
|
connBuilderSave(se, &fields, closeBuilder, onDone)()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyF3:
|
||||||
|
se.testConnectionDialog(fields)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
render(0)
|
||||||
|
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow).
|
||||||
|
AddItem(title, 1, 0, false).
|
||||||
|
AddItem(form, 0, 1, true).
|
||||||
|
AddItem(preview, 4, 0, false)
|
||||||
|
|
||||||
|
se.pages.AddAndSwitchToPage(connBuilderPage, flex, true)
|
||||||
|
se.app.SetFocus(form)
|
||||||
|
}
|
||||||
|
|
||||||
|
// connBuilderSave returns the Save action: validate, write back, close.
|
||||||
|
func connBuilderSave(se *SchemaEditor, fields *ConnFields, closeBuilder func(), onDone func(string)) func() {
|
||||||
|
return func() {
|
||||||
|
if msg := validateConnFields(*fields); msg != "" {
|
||||||
|
se.showErrorDialog("Error", msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
result := BuildConnString(*fields, false)
|
||||||
|
closeBuilder()
|
||||||
|
onDone(result)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateConnFields returns a message describing the first missing required field, or "".
|
||||||
|
func validateConnFields(f ConnFields) string {
|
||||||
|
if f.Kind == ConnSQLite {
|
||||||
|
if strings.TrimSpace(f.FilePath) == "" {
|
||||||
|
return "File path is required"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(f.Host) == "" {
|
||||||
|
return "Host is required"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// testConnectionDialog runs TestConnection in the background and reports the result.
|
||||||
|
func (se *SchemaEditor) testConnectionDialog(fields ConnFields) {
|
||||||
|
if msg := validateConnFields(fields); msg != "" {
|
||||||
|
se.showErrorDialog("Error", msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
err := TestConnection(fields)
|
||||||
|
se.app.QueueUpdateDraw(func() {
|
||||||
|
if err != nil {
|
||||||
|
se.showErrorDialog("Connection Failed", fmt.Sprintf("Connection failed:\n%v", err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
se.showSuccessDialog("Connection OK", "Connection successful", nil)
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// attachConnStringBuilder makes Enter on the named input open the builder.
|
||||||
|
func (se *SchemaEditor) attachConnStringBuilder(form *tview.Form, label, returnPage string, format func() string) {
|
||||||
|
item, ok := form.GetFormItemByLabel(label).(*tview.InputField)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() != tcell.KeyEnter {
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
hint := ConnPostgres
|
||||||
|
if format != nil && format() == "sqlite" {
|
||||||
|
hint = ConnSQLite
|
||||||
|
}
|
||||||
|
se.showConnStringBuilder(item.GetText(), hint, returnPage, func(s string) { item.SetText(s) })
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBuildConnString(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
fields ConnFields
|
||||||
|
mask bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "postgres defaults with db",
|
||||||
|
fields: func() ConnFields { f := DefaultConnFields(ConnPostgres); f.Database = "app"; return f }(),
|
||||||
|
want: "postgres://postgres@localhost:5432/app?sslmode=disable",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "postgres password unmasked",
|
||||||
|
fields: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5433", Database: "x", User: "u", Password: "p@ss/w", SSLMode: "require"},
|
||||||
|
want: "postgres://u:p%40ss%2Fw@db:5433/x?sslmode=require",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "postgres password masked",
|
||||||
|
fields: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5432", Database: "x", User: "u", Password: "secret"},
|
||||||
|
mask: true,
|
||||||
|
want: "postgres://u:****@db:5432/x",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mssql",
|
||||||
|
fields: ConnFields{Kind: ConnMSSQL, Host: "sql", Port: "1433", Database: "shop", User: "sa", Password: "pw", SSLMode: "disable"},
|
||||||
|
want: "sqlserver://sa:pw@sql:1433?database=shop&encrypt=disable",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sqlite is the plain path",
|
||||||
|
fields: ConnFields{Kind: ConnSQLite, FilePath: "/tmp/a b.db"},
|
||||||
|
want: "/tmp/a b.db",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := BuildConnString(tt.fields, tt.mask); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMaskedBuildHidesPassword(t *testing.T) {
|
||||||
|
f := ConnFields{Kind: ConnMSSQL, Host: "h", User: "u", Password: "hunter2"}
|
||||||
|
if got := BuildConnString(f, true); strings.Contains(got, "hunter2") {
|
||||||
|
t.Errorf("masked string leaks password: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseConnString(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in string
|
||||||
|
want ConnFields
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "postgres full",
|
||||||
|
in: "postgres://u:p%40ss@db:5433/app?sslmode=require&application_name=x",
|
||||||
|
want: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5433", Database: "app", User: "u", Password: "p@ss", SSLMode: "require"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "postgresql scheme, default port",
|
||||||
|
in: "postgresql://u@db/app",
|
||||||
|
want: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5432", Database: "app", User: "u"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mssql",
|
||||||
|
in: "sqlserver://sa:pw@sql:1444?database=shop&encrypt=true",
|
||||||
|
want: ConnFields{Kind: ConnMSSQL, Host: "sql", Port: "1444", Database: "shop", User: "sa", Password: "pw", SSLMode: "true"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sqlite path",
|
||||||
|
in: "/data/app.db",
|
||||||
|
want: ConnFields{Kind: ConnSQLite, FilePath: "/data/app.db"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "sqlite scheme",
|
||||||
|
in: "sqlite:///data/app.db",
|
||||||
|
want: ConnFields{Kind: ConnSQLite, FilePath: "/data/app.db"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := ParseConnString(tt.in, ConnPostgres)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got.Extra = nil
|
||||||
|
if !reflect.DeepEqual(got, tt.want) {
|
||||||
|
t.Errorf("got %+v, want %+v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseConnStringEmptyUsesHintDefaults(t *testing.T) {
|
||||||
|
got, err := ParseConnString(" ", ConnMSSQL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got.Kind != ConnMSSQL || got.Port != "1433" || got.Host != "localhost" {
|
||||||
|
t.Errorf("unexpected defaults: %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseConnStringInvalid(t *testing.T) {
|
||||||
|
if _, err := ParseConnString("postgres://u:p@host:badport/db", ConnPostgres); err == nil {
|
||||||
|
t.Error("expected error for invalid port")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConnStringRoundTrip(t *testing.T) {
|
||||||
|
for _, in := range []string{
|
||||||
|
"postgres://u:pw@db:5433/app?application_name=x&sslmode=require",
|
||||||
|
"sqlserver://sa:pw@sql:1433?application+name=x&database=shop&encrypt=false",
|
||||||
|
} {
|
||||||
|
f, err := ParseConnString(in, ConnPostgres)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := BuildConnString(f, false); got != in {
|
||||||
|
t.Errorf("round trip: got %q, want %q", got, in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTestConnectionSQLite(t *testing.T) {
|
||||||
|
if err := TestConnection(ConnFields{Kind: ConnSQLite}); err == nil {
|
||||||
|
t.Error("expected error for empty path")
|
||||||
|
}
|
||||||
|
if err := TestConnection(ConnFields{Kind: ConnSQLite, FilePath: t.TempDir() + "/missing.db"}); err == nil {
|
||||||
|
t.Error("expected error for missing file")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,198 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestColumnDataOps(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
|
||||||
|
if se.CreateColumn(5, 0, "x", "int", false, false) != nil || se.CreateColumn(0, 5, "x", "int", false, false) != nil {
|
||||||
|
t.Error("create with bad index must return nil")
|
||||||
|
}
|
||||||
|
col := se.CreateColumn(0, 0, "age", "integer", true, true)
|
||||||
|
if col == nil || col.Type != "integer" || !col.IsPrimaryKey || !col.NotNull {
|
||||||
|
t.Fatalf("create: %+v", col)
|
||||||
|
}
|
||||||
|
if se.GetColumn(0, 0, "age") != col || se.GetColumn(0, 0, "nope") != nil || se.GetColumn(9, 0, "age") != nil {
|
||||||
|
t.Error("get mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
|
||||||
|
t.Error("create second column")
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
si, ti int
|
||||||
|
old, new string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"bad table", 0, 9, "age", "age", false},
|
||||||
|
{"missing column", 0, 0, "zzz", "zzz", false},
|
||||||
|
{"in place", 0, 0, "age", "age", true},
|
||||||
|
{"rename", 0, 0, "age", "years", true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := se.UpdateColumn(tt.si, tt.ti, tt.old, tt.new, "bigint", false, true, "0", "desc"); got != tt.want {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
got := se.GetColumn(0, 0, "years")
|
||||||
|
if got == nil || got.Name != "years" || got.Type != "bigint" || got.IsPrimaryKey || got.Default != "0" || got.Description != "desc" {
|
||||||
|
t.Errorf("after update: %+v", got)
|
||||||
|
}
|
||||||
|
if se.GetColumn(0, 0, "age") != nil {
|
||||||
|
t.Error("old name must be gone")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(se.GetAllColumns(0, 0)) != 4 || se.GetAllColumns(0, 9) != nil {
|
||||||
|
t.Error("GetAllColumns")
|
||||||
|
}
|
||||||
|
if se.DeleteColumn(0, 9, "a") || se.DeleteColumn(0, 0, "zzz") {
|
||||||
|
t.Error("delete bad target must fail")
|
||||||
|
}
|
||||||
|
if !se.DeleteColumn(0, 0, "a") || se.DeleteColumn(0, 0, "a") {
|
||||||
|
t.Error("delete should succeed once")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateColumn_NilMap(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
se.db.Schemas[0].Tables[0].Columns = nil
|
||||||
|
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
|
||||||
|
t.Error("create with nil map")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRelationshipDataOps(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
rel := &models.Relationship{Name: "fk_a", FromTable: "users", ToTable: "orders"}
|
||||||
|
|
||||||
|
if se.CreateRelationship(9, 0, rel) != nil || se.CreateRelationship(0, 9, rel) != nil || se.CreateRelationship(0, -1, rel) != nil {
|
||||||
|
t.Error("create bad index")
|
||||||
|
}
|
||||||
|
// Before any relationship exists, update/delete/get/names report nothing.
|
||||||
|
se.db.Schemas[0].Tables[0].Relationships = nil
|
||||||
|
if se.UpdateRelationship(0, 0, "fk_a", rel) || se.DeleteRelationship(0, 0, "fk_a") ||
|
||||||
|
se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationshipNames(0, 0) != nil {
|
||||||
|
t.Error("nil map handling")
|
||||||
|
}
|
||||||
|
|
||||||
|
if se.CreateRelationship(0, 0, rel) != rel {
|
||||||
|
t.Fatal("create")
|
||||||
|
}
|
||||||
|
se.CreateRelationship(0, 0, &models.Relationship{Name: "fk_0"})
|
||||||
|
if got := se.GetRelationshipNames(0, 0); !reflect.DeepEqual(got, []string{"fk_0", "fk_a"}) {
|
||||||
|
t.Errorf("names must be sorted: %v", got)
|
||||||
|
}
|
||||||
|
if se.GetRelationship(0, 0, "fk_a") != rel || se.GetRelationship(0, 0, "none") != nil {
|
||||||
|
t.Error("get")
|
||||||
|
}
|
||||||
|
|
||||||
|
renamed := &models.Relationship{Name: "fk_b"}
|
||||||
|
if !se.UpdateRelationship(0, 0, "fk_a", renamed) {
|
||||||
|
t.Fatal("update")
|
||||||
|
}
|
||||||
|
if se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationship(0, 0, "fk_b") != renamed {
|
||||||
|
t.Error("rename")
|
||||||
|
}
|
||||||
|
if se.UpdateRelationship(9, 0, "x", renamed) || se.UpdateRelationship(0, 9, "x", renamed) {
|
||||||
|
t.Error("update bad index")
|
||||||
|
}
|
||||||
|
if se.DeleteRelationship(9, 0, "x") || se.DeleteRelationship(0, 9, "x") {
|
||||||
|
t.Error("delete bad index")
|
||||||
|
}
|
||||||
|
if !se.DeleteRelationship(0, 0, "fk_b") || se.GetRelationship(0, 0, "fk_b") != nil {
|
||||||
|
t.Error("delete")
|
||||||
|
}
|
||||||
|
if se.GetRelationship(9, 0, "x") != nil || se.GetRelationship(0, 9, "x") != nil ||
|
||||||
|
se.GetRelationshipNames(9, 0) != nil || se.GetRelationshipNames(0, 9) != nil {
|
||||||
|
t.Error("bad index reads")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSchemaDataOps(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
s := se.CreateSchema("sales", "desc")
|
||||||
|
if s == nil || s.Name != "sales" || s.Description != "desc" || s.Tables == nil || s.Sequences == nil || s.Enums == nil {
|
||||||
|
t.Fatalf("create: %+v", s)
|
||||||
|
}
|
||||||
|
if len(se.GetAllSchemas()) != 2 || se.GetSchema(1) != s || se.GetSchema(2) != nil || se.GetSchema(-1) != nil {
|
||||||
|
t.Error("get")
|
||||||
|
}
|
||||||
|
se.UpdateSchema(1, "billing", "owner", "d2")
|
||||||
|
if s.Name != "billing" || s.Owner != "owner" || s.Description != "d2" {
|
||||||
|
t.Errorf("update: %+v", s)
|
||||||
|
}
|
||||||
|
se.UpdateSchema(9, "x", "x", "x") // no panic
|
||||||
|
if se.DeleteSchema(9) || se.DeleteSchema(-1) {
|
||||||
|
t.Error("delete bad index")
|
||||||
|
}
|
||||||
|
if !se.DeleteSchema(1) || len(se.db.Schemas) != 1 {
|
||||||
|
t.Error("delete")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTableDataOps(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
if se.CreateTable(9, "x", "") != nil {
|
||||||
|
t.Error("create bad schema")
|
||||||
|
}
|
||||||
|
tbl := se.CreateTable(0, "orders", "d")
|
||||||
|
if tbl == nil || tbl.Schema != "public" || tbl.Columns == nil || tbl.Constraints == nil || tbl.Indexes == nil {
|
||||||
|
t.Fatalf("create: %+v", tbl)
|
||||||
|
}
|
||||||
|
if se.GetTable(0, 1) != tbl || se.GetTable(0, 2) != nil || se.GetTable(9, 0) != nil || se.GetTable(0, -1) != nil {
|
||||||
|
t.Error("get")
|
||||||
|
}
|
||||||
|
if len(se.GetAllTables()) != 2 || len(se.GetTablesInSchema(0)) != 2 || se.GetTablesInSchema(9) != nil {
|
||||||
|
t.Error("get all")
|
||||||
|
}
|
||||||
|
se.UpdateTable(0, 1, "orders2", "d2")
|
||||||
|
if tbl.Name != "orders2" || tbl.Description != "d2" {
|
||||||
|
t.Errorf("update: %+v", tbl)
|
||||||
|
}
|
||||||
|
se.UpdateTable(9, 0, "x", "x")
|
||||||
|
se.UpdateTable(0, 9, "x", "x")
|
||||||
|
if se.DeleteTable(9, 0) || se.DeleteTable(0, 9) {
|
||||||
|
t.Error("delete bad index")
|
||||||
|
}
|
||||||
|
if !se.DeleteTable(0, 1) || len(se.db.Schemas[0].Tables) != 1 {
|
||||||
|
t.Error("delete")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateDatabase(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
se.updateDatabase("n", "d", "c", "pgsql", "16")
|
||||||
|
db := se.db
|
||||||
|
if db.Name != "n" || db.Description != "d" || db.Comment != "c" || db.DatabaseType != models.PostgresqlDatabaseType || db.DatabaseVersion != "16" {
|
||||||
|
t.Errorf("%+v", db)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDomainDataOps(t *testing.T) {
|
||||||
|
se := NewSchemaEditor(models.InitDatabase("d"))
|
||||||
|
se.createDomain("a", "da")
|
||||||
|
se.createDomain("b", "db")
|
||||||
|
if len(se.db.Domains) != 2 || se.db.Domains[1].Sequence != 1 {
|
||||||
|
t.Fatalf("create: %+v", se.db.Domains)
|
||||||
|
}
|
||||||
|
se.updateDomain(0, "a2", "da2")
|
||||||
|
se.updateDomain(9, "x", "x")
|
||||||
|
if se.db.Domains[0].Name != "a2" || se.db.Domains[0].Description != "da2" {
|
||||||
|
t.Error("update")
|
||||||
|
}
|
||||||
|
se.deleteDomain(9)
|
||||||
|
se.deleteDomain(-1)
|
||||||
|
se.deleteDomain(0)
|
||||||
|
if len(se.db.Domains) != 1 || se.db.Domains[0].Name != "b" {
|
||||||
|
t.Errorf("delete: %+v", se.db.Domains)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -207,6 +207,10 @@ func (se *SchemaEditor) showDomainEditor(index int, domain *models.Domain) {
|
|||||||
se.showDomainList()
|
se.showDomainList()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
form.AddButton("Tables", func() {
|
||||||
|
se.showDomainTables(index)
|
||||||
|
})
|
||||||
|
|
||||||
form.AddButton("Delete", func() {
|
form.AddButton("Delete", func() {
|
||||||
se.showDeleteDomainConfirm(index)
|
se.showDeleteDomainConfirm(index)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -0,0 +1,134 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileEntry is a single row in the file browser.
|
||||||
|
type FileEntry struct {
|
||||||
|
Name string
|
||||||
|
IsDir bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// formatExtensions maps a UI format name to the file extensions it reads or writes.
|
||||||
|
var formatExtensions = map[string][]string{
|
||||||
|
"dbml": {".dbml"},
|
||||||
|
"dctx": {".dctx"},
|
||||||
|
"drawdb": {".json"},
|
||||||
|
"graphql": {".graphql", ".gql"},
|
||||||
|
"json": {".json"},
|
||||||
|
"yaml": {".yaml", ".yml"},
|
||||||
|
"gorm": {".go"},
|
||||||
|
"bun": {".go"},
|
||||||
|
"drizzle": {".ts"},
|
||||||
|
"prisma": {".prisma"},
|
||||||
|
"typeorm": {".ts"},
|
||||||
|
"pgsql": {".sql"},
|
||||||
|
"sqlite": {".db", ".sqlite", ".sqlite3"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// directoryFormats are formats whose reader/writer accepts a directory.
|
||||||
|
var directoryFormats = map[string]bool{
|
||||||
|
"gorm": true, "bun": true, "drizzle": true, "typeorm": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatExtensions returns the extensions for a format, or nil (no filter) if unknown.
|
||||||
|
func FormatExtensions(format string) []string {
|
||||||
|
return formatExtensions[format]
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsDirectoryFormat reports whether a format can be loaded from or saved to a directory.
|
||||||
|
func IsDirectoryFormat(format string) bool {
|
||||||
|
return directoryFormats[format]
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExpandHome replaces a leading ~ with the user's home directory.
|
||||||
|
func ExpandHome(p string) string {
|
||||||
|
if strings.HasPrefix(p, "~") {
|
||||||
|
if home, err := os.UserHomeDir(); err == nil {
|
||||||
|
return filepath.Join(home, p[1:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchesExtension reports whether name has one of exts (case-insensitive).
|
||||||
|
// An empty extension list matches everything.
|
||||||
|
func MatchesExtension(name string, exts []string) bool {
|
||||||
|
if len(exts) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
ext := strings.ToLower(filepath.Ext(name))
|
||||||
|
for _, e := range exts {
|
||||||
|
if strings.EqualFold(e, ext) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListDir returns the entries of dir: directories first, then files that match
|
||||||
|
// exts, each group sorted case-insensitively. Hidden (dot) entries are skipped
|
||||||
|
// unless showHidden is set.
|
||||||
|
func ListDir(dir string, exts []string, showHidden bool) ([]FileEntry, error) {
|
||||||
|
items, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var dirs, files []FileEntry
|
||||||
|
for _, item := range items {
|
||||||
|
name := item.Name()
|
||||||
|
if !showHidden && strings.HasPrefix(name, ".") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
isDir := item.IsDir()
|
||||||
|
if !isDir && item.Type()&os.ModeSymlink != 0 {
|
||||||
|
// Follow symlinks so links to directories are navigable.
|
||||||
|
if info, err := os.Stat(filepath.Join(dir, name)); err == nil {
|
||||||
|
isDir = info.IsDir()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if isDir {
|
||||||
|
dirs = append(dirs, FileEntry{Name: name, IsDir: true})
|
||||||
|
} else if MatchesExtension(name, exts) {
|
||||||
|
files = append(files, FileEntry{Name: name})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
byName := func(s []FileEntry) {
|
||||||
|
sort.Slice(s, func(i, j int) bool {
|
||||||
|
return strings.ToLower(s[i].Name) < strings.ToLower(s[j].Name)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
byName(dirs)
|
||||||
|
byName(files)
|
||||||
|
return append(dirs, files...), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveStart works out where the browser should open for the current input
|
||||||
|
// value. It returns the directory to show and, if the input named a file, its
|
||||||
|
// base name. Falls back to the working directory.
|
||||||
|
func ResolveStart(input string) (dir, name string) {
|
||||||
|
input = strings.TrimSpace(input)
|
||||||
|
if input != "" {
|
||||||
|
p := ExpandHome(input)
|
||||||
|
if abs, err := filepath.Abs(p); err == nil {
|
||||||
|
p = abs
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(p); err == nil && info.IsDir() {
|
||||||
|
return p, ""
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(filepath.Dir(p)); err == nil && info.IsDir() {
|
||||||
|
return filepath.Dir(p), filepath.Base(p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wd, err := os.Getwd()
|
||||||
|
if err != nil {
|
||||||
|
wd = "."
|
||||||
|
}
|
||||||
|
return wd, ""
|
||||||
|
}
|
||||||
@@ -0,0 +1,363 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"github.com/gdamore/tcell/v2"
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileBrowserMode selects between picking an existing path and choosing a save target.
|
||||||
|
type FileBrowserMode int
|
||||||
|
|
||||||
|
const (
|
||||||
|
FileBrowserLoad FileBrowserMode = iota
|
||||||
|
FileBrowserSave
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileBrowserConfig configures the file browser dialog.
|
||||||
|
type FileBrowserConfig struct {
|
||||||
|
Mode FileBrowserMode
|
||||||
|
StartPath string // current value of the input; may be empty
|
||||||
|
Extensions []string // empty = show all files
|
||||||
|
AllowDir bool // a directory is a valid result (directory-based formats)
|
||||||
|
ReturnPage string // page to switch back to when the dialog closes
|
||||||
|
OnSelect func(path string)
|
||||||
|
}
|
||||||
|
|
||||||
|
// attachFileBrowser makes Enter on the named input open the file browser,
|
||||||
|
// filtered for the currently selected format.
|
||||||
|
func (se *SchemaEditor) attachFileBrowser(form *tview.Form, label, returnPage string, mode FileBrowserMode, format func() string) {
|
||||||
|
item, ok := form.GetFormItemByLabel(label).(*tview.InputField)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() != tcell.KeyEnter {
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
f := format()
|
||||||
|
se.showFileBrowser(FileBrowserConfig{
|
||||||
|
Mode: mode,
|
||||||
|
StartPath: item.GetText(),
|
||||||
|
Extensions: FormatExtensions(f),
|
||||||
|
AllowDir: IsDirectoryFormat(f),
|
||||||
|
ReturnPage: returnPage,
|
||||||
|
OnSelect: func(path string) { item.SetText(path) },
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// showFileBrowser displays the file browser page. Esc closes it without
|
||||||
|
// calling OnSelect, leaving the originating input unchanged.
|
||||||
|
func (se *SchemaEditor) showFileBrowser(cfg FileBrowserConfig) {
|
||||||
|
const pageName = "file-browser"
|
||||||
|
|
||||||
|
dir, startName := ResolveStart(cfg.StartPath)
|
||||||
|
showHidden := false
|
||||||
|
useFilter := len(cfg.Extensions) > 0
|
||||||
|
var entries []FileEntry // rows shown below the ".." row
|
||||||
|
|
||||||
|
title := tview.NewTextView().
|
||||||
|
SetText("[::b]Select File").
|
||||||
|
SetTextAlign(tview.AlignCenter).
|
||||||
|
SetDynamicColors(true)
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
title.SetText("[::b]Save As")
|
||||||
|
}
|
||||||
|
|
||||||
|
info := tview.NewTextView().SetDynamicColors(true)
|
||||||
|
|
||||||
|
fileTable := tview.NewTable().SetSelectable(true, false).SetFixed(0, 0)
|
||||||
|
fileTable.SetBorder(true)
|
||||||
|
|
||||||
|
nameInput := tview.NewInputField().SetLabel("File name: ").SetFieldWidth(0)
|
||||||
|
nameInput.SetText(startName)
|
||||||
|
|
||||||
|
closeBrowser := func() {
|
||||||
|
se.pages.RemovePage(pageName)
|
||||||
|
se.pages.SwitchToPage(cfg.ReturnPage)
|
||||||
|
}
|
||||||
|
|
||||||
|
finish := func(path string) {
|
||||||
|
closeBrowser()
|
||||||
|
cfg.OnSelect(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
refresh := func() {
|
||||||
|
exts := cfg.Extensions
|
||||||
|
if !useFilter {
|
||||||
|
exts = nil
|
||||||
|
}
|
||||||
|
list, err := ListDir(dir, exts, showHidden)
|
||||||
|
if err != nil {
|
||||||
|
se.showErrorDialog("Error", fmt.Sprintf("Cannot read %s: %v", dir, err))
|
||||||
|
list = nil
|
||||||
|
}
|
||||||
|
entries = list
|
||||||
|
|
||||||
|
fileTable.Clear()
|
||||||
|
fileTable.SetCell(0, 0, tview.NewTableCell("[..]").SetTextColor(tcell.ColorAqua))
|
||||||
|
for i, e := range entries {
|
||||||
|
cell := tview.NewTableCell(e.Name)
|
||||||
|
if e.IsDir {
|
||||||
|
cell.SetText(e.Name + "/").SetTextColor(tcell.ColorAqua)
|
||||||
|
}
|
||||||
|
fileTable.SetCell(i+1, 0, cell)
|
||||||
|
}
|
||||||
|
fileTable.Select(0, 0)
|
||||||
|
if len(entries) > 0 {
|
||||||
|
fileTable.Select(1, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
filterText := "all files"
|
||||||
|
if useFilter {
|
||||||
|
filterText = fmt.Sprintf("%v", cfg.Extensions)
|
||||||
|
}
|
||||||
|
hiddenText := "hidden: off"
|
||||||
|
if showHidden {
|
||||||
|
hiddenText = "hidden: on"
|
||||||
|
}
|
||||||
|
info.SetText(fmt.Sprintf("%s [yellow](%s, filter: %s)[-]", tview.Escape(dir), hiddenText, tview.Escape(filterText)))
|
||||||
|
fileTable.SetTitle(" Files ")
|
||||||
|
}
|
||||||
|
|
||||||
|
goUp := func() {
|
||||||
|
parent := filepath.Dir(dir)
|
||||||
|
if parent == dir {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
prev := filepath.Base(dir)
|
||||||
|
dir = parent
|
||||||
|
refresh()
|
||||||
|
for i, e := range entries {
|
||||||
|
if e.Name == prev {
|
||||||
|
fileTable.Select(i+1, 0)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
selected := func() (FileEntry, bool) {
|
||||||
|
row, _ := fileTable.GetSelection()
|
||||||
|
if row < 1 || row > len(entries) {
|
||||||
|
return FileEntry{}, false
|
||||||
|
}
|
||||||
|
return entries[row-1], true
|
||||||
|
}
|
||||||
|
|
||||||
|
// confirmOverwrite asks before replacing an existing file (not directories).
|
||||||
|
confirmOverwrite := func(path string) {
|
||||||
|
modal := tview.NewModal().
|
||||||
|
SetText(fmt.Sprintf("File already exists:\n%s\n\nOverwrite it?", path)).
|
||||||
|
AddButtons([]string{"Cancel", "Overwrite"}).
|
||||||
|
SetDoneFunc(func(_ int, label string) {
|
||||||
|
se.pages.RemovePage("overwrite-confirm")
|
||||||
|
se.pages.SwitchToPage(pageName)
|
||||||
|
if label == "Overwrite" {
|
||||||
|
finish(path)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
modal.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage("overwrite-confirm")
|
||||||
|
se.pages.SwitchToPage(pageName)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddAndSwitchToPage("overwrite-confirm", modal, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
chooseSave := func() {
|
||||||
|
name := nameInput.GetText()
|
||||||
|
if name == "" {
|
||||||
|
if cfg.AllowDir {
|
||||||
|
finish(dir)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
se.showErrorDialog("Error", "Enter a file name")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
path := filepath.Join(dir, name)
|
||||||
|
if st, err := os.Stat(path); err == nil {
|
||||||
|
if st.IsDir() {
|
||||||
|
se.showErrorDialog("Error", name+" is a directory")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
confirmOverwrite(path)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// chooseHighlighted handles Select: the highlighted entry in load mode, or
|
||||||
|
// the typed name in save mode.
|
||||||
|
chooseHighlighted := func() {
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
chooseSave()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e, ok := selected()
|
||||||
|
switch {
|
||||||
|
case ok && !e.IsDir:
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
case ok && cfg.AllowDir:
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
case cfg.AllowDir:
|
||||||
|
finish(dir)
|
||||||
|
default:
|
||||||
|
se.showErrorDialog("Error", "Select a file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
activate := func() {
|
||||||
|
row, _ := fileTable.GetSelection()
|
||||||
|
if row == 0 {
|
||||||
|
goUp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e, ok := selected()
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if e.IsDir {
|
||||||
|
dir = filepath.Join(dir, e.Name)
|
||||||
|
refresh()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
nameInput.SetText(e.Name)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
}
|
||||||
|
|
||||||
|
toggleHidden := func() { showHidden = !showHidden; refresh() }
|
||||||
|
toggleFilter := func() {
|
||||||
|
if len(cfg.Extensions) > 0 {
|
||||||
|
useFilter = !useFilter
|
||||||
|
refresh()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
btnSelect := tview.NewButton("Select [s]").SetSelectedFunc(chooseHighlighted)
|
||||||
|
btnHidden := tview.NewButton("Hidden [h]").SetSelectedFunc(toggleHidden)
|
||||||
|
btnFilter := tview.NewButton("Filter [f]").SetSelectedFunc(toggleFilter)
|
||||||
|
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(closeBrowser)
|
||||||
|
|
||||||
|
btnFlex := tview.NewFlex().
|
||||||
|
AddItem(btnSelect, 0, 1, false).
|
||||||
|
AddItem(btnHidden, 0, 1, false).
|
||||||
|
AddItem(btnFilter, 0, 1, false).
|
||||||
|
AddItem(btnBack, 0, 1, false)
|
||||||
|
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow).
|
||||||
|
AddItem(title, 1, 0, false).
|
||||||
|
AddItem(info, 1, 0, false).
|
||||||
|
AddItem(fileTable, 0, 1, true)
|
||||||
|
|
||||||
|
focusOrder := []tview.Primitive{fileTable}
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
flex.AddItem(nameInput, 1, 0, false)
|
||||||
|
focusOrder = append(focusOrder, nameInput)
|
||||||
|
}
|
||||||
|
flex.AddItem(btnFlex, 1, 0, false)
|
||||||
|
focusOrder = append(focusOrder, btnSelect, btnHidden, btnFilter, btnBack)
|
||||||
|
|
||||||
|
// Circular Tab / Shift+Tab across every focusable widget.
|
||||||
|
cycle := func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
step := 0
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyTab:
|
||||||
|
step = 1
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
step = -1
|
||||||
|
default:
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
for i, p := range focusOrder {
|
||||||
|
if p.HasFocus() {
|
||||||
|
se.app.SetFocus(focusOrder[(i+step+len(focusOrder))%len(focusOrder)])
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fileTable.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyEscape:
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyEnter:
|
||||||
|
activate()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyBackspace, tcell.KeyBackspace2, tcell.KeyLeft:
|
||||||
|
goUp()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Rune() {
|
||||||
|
case 's':
|
||||||
|
chooseHighlighted()
|
||||||
|
return nil
|
||||||
|
case 'h':
|
||||||
|
toggleHidden()
|
||||||
|
return nil
|
||||||
|
case 'f':
|
||||||
|
toggleFilter()
|
||||||
|
return nil
|
||||||
|
case 'b':
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
nameInput.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyEscape:
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyEnter:
|
||||||
|
chooseSave()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, b := range []*tview.Button{btnSelect, btnHidden, btnFilter, btnBack} {
|
||||||
|
b.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
refresh()
|
||||||
|
if startName != "" {
|
||||||
|
for i, e := range entries {
|
||||||
|
if e.Name == startName {
|
||||||
|
fileTable.Select(i+1, 0)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
se.pages.AddAndSwitchToPage(pageName, flex, true)
|
||||||
|
se.app.SetFocus(fileTable)
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func touch(t *testing.T, path string) {
|
||||||
|
t.Helper()
|
||||||
|
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(path, nil, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func names(entries []FileEntry) []string {
|
||||||
|
var out []string
|
||||||
|
for _, e := range entries {
|
||||||
|
if e.IsDir {
|
||||||
|
out = append(out, e.Name+"/")
|
||||||
|
} else {
|
||||||
|
out = append(out, e.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchesExtension(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
exts []string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"a.dbml", []string{".dbml"}, true},
|
||||||
|
{"A.DBML", []string{".dbml"}, true},
|
||||||
|
{"a.json", []string{".dbml"}, false},
|
||||||
|
{"a.yml", []string{".yaml", ".yml"}, true},
|
||||||
|
{"noext", []string{".sql"}, false},
|
||||||
|
{"anything", nil, true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := MatchesExtension(tt.name, tt.exts); got != tt.want {
|
||||||
|
t.Errorf("MatchesExtension(%q, %v) = %v, want %v", tt.name, tt.exts, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListDirFilterAndHidden(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
touch(t, filepath.Join(dir, "b.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "A.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "c.json"))
|
||||||
|
touch(t, filepath.Join(dir, ".hidden.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "sub", "x.txt"))
|
||||||
|
touch(t, filepath.Join(dir, ".git", "x"))
|
||||||
|
|
||||||
|
got, err := ListDir(dir, FormatExtensions("dbml"), false)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if want := []string{"sub/", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("filtered: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, _ = ListDir(dir, FormatExtensions("dbml"), true)
|
||||||
|
if want := []string{".git/", "sub/", ".hidden.dbml", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("hidden: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, _ = ListDir(dir, nil, false)
|
||||||
|
if want := []string{"sub/", "A.dbml", "b.dbml", "c.json"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("no filter: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListDirMissing(t *testing.T) {
|
||||||
|
if _, err := ListDir(filepath.Join(t.TempDir(), "nope"), nil, false); err == nil {
|
||||||
|
t.Error("expected error for missing directory")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveStart(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
file := filepath.Join(dir, "schema.dbml")
|
||||||
|
touch(t, file)
|
||||||
|
wd, _ := os.Getwd()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in string
|
||||||
|
wantDir string
|
||||||
|
wantFileName string
|
||||||
|
}{
|
||||||
|
{"existing file", file, dir, "schema.dbml"},
|
||||||
|
{"directory", dir, dir, ""},
|
||||||
|
{"new file in existing dir", filepath.Join(dir, "new.dbml"), dir, "new.dbml"},
|
||||||
|
{"empty", "", wd, ""},
|
||||||
|
{"nonexistent parent", filepath.Join(dir, "no", "such", "f.dbml"), wd, ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
d, n := ResolveStart(tt.in)
|
||||||
|
if d != tt.wantDir || n != tt.wantFileName {
|
||||||
|
t.Errorf("got (%q, %q), want (%q, %q)", d, n, tt.wantDir, tt.wantFileName)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatExtensions(t *testing.T) {
|
||||||
|
if got := FormatExtensions("yaml"); !reflect.DeepEqual(got, []string{".yaml", ".yml"}) {
|
||||||
|
t.Errorf("yaml: %v", got)
|
||||||
|
}
|
||||||
|
if FormatExtensions("unknown") != nil {
|
||||||
|
t.Error("unknown format should not filter")
|
||||||
|
}
|
||||||
|
if !IsDirectoryFormat("gorm") || IsDirectoryFormat("json") {
|
||||||
|
t.Error("directory format detection wrong")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,294 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
const uiFixtures = "../../tests/assets"
|
||||||
|
|
||||||
|
func newUIEditor() *SchemaEditor {
|
||||||
|
se := NewSchemaEditor(models.InitDatabase("start"))
|
||||||
|
se.db = newTestEditor().db
|
||||||
|
return se
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasPage(se *SchemaEditor, name string) bool {
|
||||||
|
return se.pages.HasPage(name)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortedKeysAndColumnNames(t *testing.T) {
|
||||||
|
if got := sortedKeys(map[string]int{"b": 1, "a": 2, "c": 3}); strings.Join(got, ",") != "a,b,c" {
|
||||||
|
t.Errorf("sortedKeys: %v", got)
|
||||||
|
}
|
||||||
|
if got := sortedKeys[int](nil); len(got) != 0 {
|
||||||
|
t.Errorf("nil map: %v", got)
|
||||||
|
}
|
||||||
|
tbl := models.InitTable("t", "s")
|
||||||
|
tbl.Columns["z"] = models.InitColumn("z", "t", "s")
|
||||||
|
tbl.Columns["a"] = models.InitColumn("a", "t", "s")
|
||||||
|
if got := getColumnNames(tbl); strings.Join(got, ",") != "a,z" {
|
||||||
|
t.Errorf("getColumnNames: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocations(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
se.db.Schemas = append(se.db.Schemas, models.InitSchema("empty"))
|
||||||
|
sl := se.schemaLocations()
|
||||||
|
if len(sl) != 2 || sl[0].label != "public" || sl[1].schemaIndex != 1 || sl[0].tableIndex != -1 {
|
||||||
|
t.Errorf("schemaLocations: %+v", sl)
|
||||||
|
}
|
||||||
|
tl := se.tableLocations()
|
||||||
|
if len(tl) != 1 || tl[0].label != "public.users" || tl[0].schemaIndex != 0 || tl[0].tableIndex != 0 {
|
||||||
|
t.Errorf("tableLocations: %+v", tl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseSkipTablesUI(t *testing.T) {
|
||||||
|
if got := parseSkipTablesUI(""); len(got) != 0 {
|
||||||
|
t.Errorf("empty: %v", got)
|
||||||
|
}
|
||||||
|
got := parseSkipTablesUI(" Users , ORDERS ,, ")
|
||||||
|
if len(got) != 2 || !got["users"] || !got["orders"] {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHelpTexts(t *testing.T) {
|
||||||
|
for name, fn := range map[string]func() string{"load": getLoadHelpText, "save": getSaveHelpText, "import": getImportHelpText} {
|
||||||
|
if txt := fn(); !strings.Contains(txt, "dbml") && name != "save" || txt == "" {
|
||||||
|
t.Errorf("%s help text: %q", name, txt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObjectKinds(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
if err := se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}, Unique: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "v1", Definition: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s1", IncrementBy: 1, StartValue: 1}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveScript(0, -1, &models.Script{Name: "sc1", SQL: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
kinds := map[string]objectKind{
|
||||||
|
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
|
||||||
|
}
|
||||||
|
for page, k := range kinds {
|
||||||
|
t.Run(page, func(t *testing.T) {
|
||||||
|
if k.page != page || k.title == "" || k.singular == "" || len(k.headers) == 0 {
|
||||||
|
t.Fatalf("metadata: %+v", k)
|
||||||
|
}
|
||||||
|
rows := k.rows()
|
||||||
|
if len(rows) != 1 {
|
||||||
|
t.Fatalf("rows: %+v", rows)
|
||||||
|
}
|
||||||
|
for _, r := range rows {
|
||||||
|
if len(r.cells) != len(k.headers) {
|
||||||
|
t.Errorf("cells %v do not match headers %v", r.cells, k.headers)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(k.locations()) == 0 {
|
||||||
|
t.Error("no locations")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Editing an existing row without changes keeps it valid.
|
||||||
|
form := tview.NewForm()
|
||||||
|
save := k.buildForm(form, &rows[0])
|
||||||
|
if form.GetFormItemCount() == 0 {
|
||||||
|
t.Error("no form fields")
|
||||||
|
}
|
||||||
|
loc := k.locations()[0]
|
||||||
|
loc.schemaIndex, loc.tableIndex = rows[0].schemaIndex, rows[0].tableIndex
|
||||||
|
if err := save(loc); err != nil {
|
||||||
|
t.Errorf("save unchanged: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A blank new form is rejected by validation.
|
||||||
|
blank := tview.NewForm()
|
||||||
|
saveBlank := k.buildForm(blank, nil)
|
||||||
|
if err := saveBlank(k.locations()[0]); err == nil {
|
||||||
|
t.Error("blank form accepted")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !k.remove(rows[0]) || len(k.rows()) != 0 {
|
||||||
|
t.Error("remove failed")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObjectKind_CreateIndexFromForm(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
k := se.indexKind()
|
||||||
|
form := tview.NewForm()
|
||||||
|
save := k.buildForm(form, nil)
|
||||||
|
form.GetFormItemByLabel("Name").(*tview.InputField).SetText("idx_new")
|
||||||
|
form.GetFormItemByLabel("Columns (comma separated)").(*tview.InputField).SetText("id, email")
|
||||||
|
if err := save(k.locations()[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
idx := se.db.Schemas[0].Tables[0].Indexes["idx_new"]
|
||||||
|
if idx == nil || len(idx.Columns) != 2 || idx.Type != "btree" {
|
||||||
|
t.Errorf("index: %+v", idx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadDatabase(t *testing.T) {
|
||||||
|
for _, tt := range []struct{ format, path string }{
|
||||||
|
{"dbml", "dbml/simple.dbml"}, {"json", "json/database.json"}, {"yaml", "yaml/database.yaml"},
|
||||||
|
{"drawdb", "drawdb/simple.json"}, {"dctx", "dctx/p1.dctx"}, {"graphql", "graphql/simple.graphql"},
|
||||||
|
{"prisma", "prisma/example.prisma"}, {"typeorm", "typeorm/example.ts"},
|
||||||
|
{"drizzle", "drizzle/schema.ts"}, {"gorm", "gorm/simple.go"}, {"bun", "bun/simple.go"},
|
||||||
|
} {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
se.loadDatabase(tt.format, filepath.Join(uiFixtures, tt.path), "")
|
||||||
|
if hasPage(se, "error-dialog") || !hasPage(se, "success-dialog") {
|
||||||
|
t.Fatalf("expected success dialog (pages: error=%v)", hasPage(se, "error-dialog"))
|
||||||
|
}
|
||||||
|
if se.loadConfig == nil || se.loadConfig.SourceType != tt.format || len(se.db.Schemas) == 0 {
|
||||||
|
t.Errorf("state: %+v db=%+v", se.loadConfig, se.db)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
errCases := []struct {
|
||||||
|
name, format, path, conn string
|
||||||
|
}{
|
||||||
|
{"pgsql no conn", "pgsql", "", ""},
|
||||||
|
{"file required", "json", "", ""},
|
||||||
|
{"unsupported", "nope", "x", ""},
|
||||||
|
{"missing file", "json", filepath.Join(t.TempDir(), "missing.json"), ""},
|
||||||
|
}
|
||||||
|
for _, tt := range errCases {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
before := se.db
|
||||||
|
se.loadDatabase(tt.format, tt.path, tt.conn)
|
||||||
|
if !hasPage(se, "error-dialog") {
|
||||||
|
t.Error("expected error dialog")
|
||||||
|
}
|
||||||
|
if se.db != before || se.loadConfig != nil {
|
||||||
|
t.Error("state must be unchanged on error")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateNewDatabase(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
se.loadConfig = &LoadConfig{SourceType: "json"}
|
||||||
|
se.createNewDatabase()
|
||||||
|
if se.db.Name != "New Database" || len(se.db.Schemas) != 0 || se.loadConfig != nil || !hasPage(se, "success-dialog") {
|
||||||
|
t.Errorf("state: %+v", se.db)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSaveDatabase(t *testing.T) {
|
||||||
|
for _, tt := range []struct{ format, file string }{
|
||||||
|
{"json", "o.json"}, {"yaml", "o.yaml"}, {"dbml", "o.dbml"}, {"drawdb", "o.drawdb.json"},
|
||||||
|
{"graphql", "o.graphql"}, {"prisma", "o.prisma"}, {"typeorm", "o.ts"}, {"drizzle", "d.ts"},
|
||||||
|
{"gorm", "g.go"}, {"bun", "b.go"},
|
||||||
|
} {
|
||||||
|
t.Run(tt.format, func(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
out := filepath.Join(t.TempDir(), tt.file)
|
||||||
|
se.saveDatabase(tt.format, out)
|
||||||
|
if hasPage(se, "error-dialog") {
|
||||||
|
t.Fatal("unexpected error dialog")
|
||||||
|
}
|
||||||
|
if se.saveConfig == nil || se.saveConfig.FilePath != out || se.saveConfig.TargetType != tt.format {
|
||||||
|
t.Errorf("saveConfig: %+v", se.saveConfig)
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(out); err != nil || info.Size() == 0 {
|
||||||
|
t.Errorf("output: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
for name, args := range map[string][2]string{
|
||||||
|
"pgsql unsupported": {"pgsql", "x.sql"},
|
||||||
|
"path required": {"json", ""},
|
||||||
|
"unknown format": {"nope", "x"},
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
se.saveDatabase(args[0], args[1])
|
||||||
|
if !hasPage(se, "error-dialog") || se.saveConfig != nil {
|
||||||
|
t.Error("expected error dialog and no saveConfig")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImportAndMerge(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
se.importAndMergeDatabase("json", filepath.Join(uiFixtures, "json/database.json"), "", false, false, false, false, false, "")
|
||||||
|
if hasPage(se, "error-dialog") {
|
||||||
|
t.Fatal("unexpected error dialog")
|
||||||
|
}
|
||||||
|
|
||||||
|
for name, args := range map[string][3]string{
|
||||||
|
"pgsql no conn": {"pgsql", "", ""},
|
||||||
|
"file required": {"json", "", ""},
|
||||||
|
"unsupported": {"nope", "x", ""},
|
||||||
|
"missing file": {"json", filepath.Join(t.TempDir(), "missing.json"), ""},
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
se.importAndMergeDatabase(args[0], args[1], args[2], false, false, false, false, false, "")
|
||||||
|
if !hasPage(se, "error-dialog") {
|
||||||
|
t.Error("expected error dialog")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPerformMerge(t *testing.T) {
|
||||||
|
se := newUIEditor()
|
||||||
|
src := models.InitDatabase("src")
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
tbl := models.InitTable("orders", "public")
|
||||||
|
tbl.Columns["id"] = models.InitColumn("id", "orders", "public")
|
||||||
|
skip := models.InitTable("skipme", "public")
|
||||||
|
s.Tables = append(s.Tables, tbl, skip)
|
||||||
|
src.Schemas = append(src.Schemas, s)
|
||||||
|
|
||||||
|
se.performMerge(src, false, false, false, false, false, "SkipMe")
|
||||||
|
if !hasPage(se, "success-dialog") {
|
||||||
|
t.Error("expected success dialog")
|
||||||
|
}
|
||||||
|
names := map[string]bool{}
|
||||||
|
for _, tb := range se.db.Schemas[0].Tables {
|
||||||
|
names[tb.Name] = true
|
||||||
|
}
|
||||||
|
if !names["users"] || !names["orders"] || names["skipme"] || len(names) != 2 {
|
||||||
|
t.Errorf("tables after merge: %v", names)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEditorAccessors(t *testing.T) {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
lc, sc := &LoadConfig{SourceType: "json"}, &SaveConfig{TargetType: "yaml"}
|
||||||
|
se := NewSchemaEditorWithConfigs(db, lc, sc)
|
||||||
|
if se.GetDatabase() != db || se.loadConfig != lc || se.saveConfig != sc || se.app == nil || se.pages == nil {
|
||||||
|
t.Errorf("%+v", se)
|
||||||
|
}
|
||||||
|
if se.createMainMenu() == nil {
|
||||||
|
t.Error("main menu")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newDialogTestEditor() *SchemaEditor {
|
||||||
|
se := &SchemaEditor{app: tview.NewApplication(), pages: tview.NewPages()}
|
||||||
|
se.pages.AddPage("origin", tview.NewBox(), true, true)
|
||||||
|
return se
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFileBrowserOpensOnEachMode(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
touch(t, filepath.Join(dir, "a.dbml"))
|
||||||
|
|
||||||
|
for _, mode := range []FileBrowserMode{FileBrowserLoad, FileBrowserSave} {
|
||||||
|
se := newDialogTestEditor()
|
||||||
|
se.showFileBrowser(FileBrowserConfig{
|
||||||
|
Mode: mode,
|
||||||
|
StartPath: filepath.Join(dir, "a.dbml"),
|
||||||
|
Extensions: FormatExtensions("dbml"),
|
||||||
|
ReturnPage: "origin",
|
||||||
|
OnSelect: func(string) { t.Error("OnSelect must not fire without a selection") },
|
||||||
|
})
|
||||||
|
if !se.pages.HasPage("file-browser") {
|
||||||
|
t.Errorf("mode %d: file-browser page missing", mode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConnStringBuilderOpensForEachKind(t *testing.T) {
|
||||||
|
for _, in := range []string{
|
||||||
|
"",
|
||||||
|
"postgres://u:pw@db:5432/app?sslmode=disable",
|
||||||
|
"sqlserver://sa:pw@sql:1433?database=shop&encrypt=disable",
|
||||||
|
"/tmp/app.db",
|
||||||
|
"postgres://u:p@host:badport/db", // parse error falls back to defaults
|
||||||
|
} {
|
||||||
|
se := newDialogTestEditor()
|
||||||
|
se.showConnStringBuilder(in, ConnPostgres, "origin", func(string) {
|
||||||
|
t.Error("onDone must not fire without Save")
|
||||||
|
})
|
||||||
|
if !se.pages.HasPage(connBuilderPage) {
|
||||||
|
t.Errorf("%q: builder page missing", in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateConnFields(t *testing.T) {
|
||||||
|
if validateConnFields(ConnFields{Kind: ConnSQLite}) == "" {
|
||||||
|
t.Error("sqlite without path should be invalid")
|
||||||
|
}
|
||||||
|
if validateConnFields(ConnFields{Kind: ConnPostgres}) == "" {
|
||||||
|
t.Error("postgres without host should be invalid")
|
||||||
|
}
|
||||||
|
if msg := validateConnFields(DefaultConnFields(ConnMSSQL)); msg != "" {
|
||||||
|
t.Errorf("defaults should be valid, got %q", msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -92,6 +92,9 @@ func (se *SchemaEditor) showLoadScreen() {
|
|||||||
connString = value
|
connString = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "load-database", FileBrowserLoad, func() string { return currentFormat })
|
||||||
|
se.attachConnStringBuilder(form, "Connection String", "load-database", func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddTextView("Help", getLoadHelpText(), 0, 5, true, false)
|
form.AddTextView("Help", getLoadHelpText(), 0, 5, true, false)
|
||||||
|
|
||||||
// Buttons
|
// Buttons
|
||||||
@@ -190,6 +193,8 @@ func (se *SchemaEditor) showSaveScreen() {
|
|||||||
filePath = value
|
filePath = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "save-database", FileBrowserSave, func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddTextView("Help", getSaveHelpText(), 0, 5, true, false)
|
form.AddTextView("Help", getSaveHelpText(), 0, 5, true, false)
|
||||||
|
|
||||||
// Buttons
|
// Buttons
|
||||||
@@ -469,6 +474,8 @@ func getLoadHelpText() string {
|
|||||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm
|
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm
|
||||||
Database formats: pgsql (requires connection string)
|
Database formats: pgsql (requires connection string)
|
||||||
|
|
||||||
|
Press Enter in File Path to browse files, or in Connection String to open the builder.
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
- File path: ~/schemas/mydb.dbml or /path/to/schema.json
|
- File path: ~/schemas/mydb.dbml or /path/to/schema.json
|
||||||
- Connection: postgres://user:pass@localhost/dbname`
|
- Connection: postgres://user:pass@localhost/dbname`
|
||||||
@@ -520,6 +527,8 @@ func (se *SchemaEditor) showUpdateExistingDatabaseConfirm() {
|
|||||||
func getSaveHelpText() string {
|
func getSaveHelpText() string {
|
||||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql (SQL export)
|
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql (SQL export)
|
||||||
|
|
||||||
|
Press Enter in File Path to browse for a target.
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
- File: ~/schemas/mydb.dbml
|
- File: ~/schemas/mydb.dbml
|
||||||
- Directory (for code formats): ./models/`
|
- Directory (for code formats): ./models/`
|
||||||
@@ -570,6 +579,9 @@ func (se *SchemaEditor) showImportScreen() {
|
|||||||
connString = value
|
connString = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "import-database", FileBrowserLoad, func() string { return currentFormat })
|
||||||
|
se.attachConnStringBuilder(form, "Connection String", "import-database", func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddInputField("Skip Tables (comma-separated)", "", 50, nil, func(value string) {
|
form.AddInputField("Skip Tables (comma-separated)", "", 50, nil, func(value string) {
|
||||||
skipTables = value
|
skipTables = value
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -39,6 +39,18 @@ func (se *SchemaEditor) createMainMenu() tview.Primitive {
|
|||||||
AddItem("Manage Domains", "View, create, edit, and delete domains", 'd', func() {
|
AddItem("Manage Domains", "View, create, edit, and delete domains", 'd', func() {
|
||||||
se.showDomainList()
|
se.showDomainList()
|
||||||
}).
|
}).
|
||||||
|
AddItem("Manage Indexes", "View, create, edit, and delete table indexes", 'x', func() {
|
||||||
|
se.showObjectList(se.indexKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Views", "View, create, edit, and delete views", 'v', func() {
|
||||||
|
se.showObjectList(se.viewKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Sequences", "View, create, edit, and delete sequences", 'u', func() {
|
||||||
|
se.showObjectList(se.sequenceKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Scripts", "View, create, edit, and delete SQL scripts", 'c', func() {
|
||||||
|
se.showObjectList(se.scriptKind())
|
||||||
|
}).
|
||||||
AddItem("Import & Merge", "Import and merge schema from another database", 'i', func() {
|
AddItem("Import & Merge", "Import and merge schema from another database", 'i', func() {
|
||||||
se.showImportScreen()
|
se.showImportScreen()
|
||||||
}).
|
}).
|
||||||
|
|||||||
@@ -0,0 +1,263 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Data operations for indexes, views, sequences, scripts and domain/table assignment.
|
||||||
|
|
||||||
|
func (se *SchemaEditor) schemaAt(schemaIndex int) (*models.Schema, error) {
|
||||||
|
if schemaIndex < 0 || schemaIndex >= len(se.db.Schemas) {
|
||||||
|
return nil, errors.New("schema not found")
|
||||||
|
}
|
||||||
|
return se.db.Schemas[schemaIndex], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) tableAt(schemaIndex, tableIndex int) (*models.Schema, *models.Table, error) {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
if tableIndex < 0 || tableIndex >= len(schema.Tables) {
|
||||||
|
return nil, nil, errors.New("table not found")
|
||||||
|
}
|
||||||
|
return schema, schema.Tables[tableIndex], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// splitList splits a comma separated list, trimming blanks and dropping empty entries.
|
||||||
|
func splitList(s string) []string {
|
||||||
|
parts := make([]string, 0)
|
||||||
|
for _, p := range strings.Split(s, ",") {
|
||||||
|
if p = strings.TrimSpace(p); p != "" {
|
||||||
|
parts = append(parts, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return parts
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveIndex adds an index to a table. When oldName is non-empty the index of that
|
||||||
|
// name is replaced (and renamed if needed).
|
||||||
|
func (se *SchemaEditor) SaveIndex(schemaIndex, tableIndex int, oldName string, idx *models.Index) error {
|
||||||
|
schema, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
idx.Name = strings.TrimSpace(idx.Name)
|
||||||
|
if idx.Name == "" {
|
||||||
|
return errors.New("index name is required")
|
||||||
|
}
|
||||||
|
if len(idx.Columns) == 0 {
|
||||||
|
return errors.New("index needs at least one column")
|
||||||
|
}
|
||||||
|
for _, c := range idx.Columns {
|
||||||
|
if _, ok := table.Columns[c]; !ok {
|
||||||
|
return fmt.Errorf("column %q not found in table %s", c, table.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, exists := table.Indexes[idx.Name]; exists && idx.Name != oldName {
|
||||||
|
return fmt.Errorf("index %q already exists", idx.Name)
|
||||||
|
}
|
||||||
|
if table.Indexes == nil {
|
||||||
|
table.Indexes = make(map[string]*models.Index)
|
||||||
|
}
|
||||||
|
if oldName != "" {
|
||||||
|
delete(table.Indexes, oldName)
|
||||||
|
}
|
||||||
|
idx.Table = table.Name
|
||||||
|
idx.Schema = schema.Name
|
||||||
|
table.Indexes[idx.Name] = idx
|
||||||
|
table.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteIndex removes an index from a table.
|
||||||
|
func (se *SchemaEditor) DeleteIndex(schemaIndex, tableIndex int, name string) bool {
|
||||||
|
_, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if _, ok := table.Indexes[name]; !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
delete(table.Indexes, name)
|
||||||
|
table.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveView adds a view to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveView(schemaIndex, at int, v *models.View) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
v.Name = strings.TrimSpace(v.Name)
|
||||||
|
if v.Name == "" {
|
||||||
|
return errors.New("view name is required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(v.Definition) == "" {
|
||||||
|
return errors.New("view definition is required")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Views {
|
||||||
|
if i != at && o.Name == v.Name {
|
||||||
|
return fmt.Errorf("view %q already exists", v.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
v.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Views) {
|
||||||
|
schema.Views[at] = v
|
||||||
|
} else {
|
||||||
|
schema.Views = append(schema.Views, v)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteView removes the view at position at.
|
||||||
|
func (se *SchemaEditor) DeleteView(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Views) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Views = append(schema.Views[:at], schema.Views[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveSequence adds a sequence to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveSequence(schemaIndex, at int, s *models.Sequence) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.Name = strings.TrimSpace(s.Name)
|
||||||
|
if s.Name == "" {
|
||||||
|
return errors.New("sequence name is required")
|
||||||
|
}
|
||||||
|
if s.IncrementBy == 0 {
|
||||||
|
return errors.New("increment must not be zero")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Sequences {
|
||||||
|
if i != at && o.Name == s.Name {
|
||||||
|
return fmt.Errorf("sequence %q already exists", s.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Sequences) {
|
||||||
|
schema.Sequences[at] = s
|
||||||
|
} else {
|
||||||
|
schema.Sequences = append(schema.Sequences, s)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteSequence removes the sequence at position at.
|
||||||
|
func (se *SchemaEditor) DeleteSequence(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Sequences) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Sequences = append(schema.Sequences[:at], schema.Sequences[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveScript adds a script to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveScript(schemaIndex, at int, s *models.Script) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.Name = strings.TrimSpace(s.Name)
|
||||||
|
if s.Name == "" {
|
||||||
|
return errors.New("script name is required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(s.SQL) == "" {
|
||||||
|
return errors.New("script SQL is required")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Scripts {
|
||||||
|
if i != at && o.Name == s.Name {
|
||||||
|
return fmt.Errorf("script %q already exists", s.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Scripts) {
|
||||||
|
schema.Scripts[at] = s
|
||||||
|
} else {
|
||||||
|
schema.Scripts = append(schema.Scripts, s)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteScript removes the script at position at.
|
||||||
|
func (se *SchemaEditor) DeleteScript(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Scripts) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Scripts = append(schema.Scripts[:at], schema.Scripts[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// AssignTableToDomain adds a reference to schemaName.tableName to the domain at domainIndex.
|
||||||
|
func (se *SchemaEditor) AssignTableToDomain(domainIndex int, schemaName, tableName string) error {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return errors.New("domain not found")
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
var table *models.Table
|
||||||
|
for _, s := range se.db.Schemas {
|
||||||
|
if s.Name != schemaName {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
if t.Name == tableName {
|
||||||
|
table = t
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if table == nil {
|
||||||
|
return fmt.Errorf("table %s.%s not found", schemaName, tableName)
|
||||||
|
}
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||||
|
return fmt.Errorf("table %s.%s is already in domain %s", schemaName, tableName, domain.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
dt := models.InitDomainTable(tableName, schemaName)
|
||||||
|
dt.RefTable = table
|
||||||
|
dt.Sequence = uint(len(domain.Tables))
|
||||||
|
domain.Tables = append(domain.Tables, dt)
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnassignTableFromDomain removes the reference to schemaName.tableName from the domain.
|
||||||
|
func (se *SchemaEditor) UnassignTableFromDomain(domainIndex int, schemaName, tableName string) bool {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
for i, dt := range domain.Tables {
|
||||||
|
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||||
|
domain.Tables = append(domain.Tables[:i], domain.Tables[i+1:]...)
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestEditor() *SchemaEditor {
|
||||||
|
db := models.InitDatabase("test")
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
table := models.InitTable("users", "public")
|
||||||
|
table.Columns["id"] = models.InitColumn("id", "users", "public")
|
||||||
|
table.Columns["email"] = models.InitColumn("email", "users", "public")
|
||||||
|
schema.Tables = append(schema.Tables, table)
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
return &SchemaEditor{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSaveIndex(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
table := se.db.Schemas[0].Tables[0]
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
old string
|
||||||
|
idx *models.Index
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"valid", "", &models.Index{Name: "idx_email", Columns: []string{"email"}, Unique: true}, false},
|
||||||
|
{"duplicate", "", &models.Index{Name: "idx_email", Columns: []string{"email"}}, true},
|
||||||
|
{"missing name", "", &models.Index{Columns: []string{"email"}}, true},
|
||||||
|
{"no columns", "", &models.Index{Name: "idx_none"}, true},
|
||||||
|
{"unknown column", "", &models.Index{Name: "idx_bad", Columns: []string{"nope"}}, true},
|
||||||
|
{"rename", "idx_email", &models.Index{Name: "idx_email2", Columns: []string{"email", "id"}}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if err := se.SaveIndex(0, 0, tt.old, tt.idx); (err != nil) != tt.wantErr {
|
||||||
|
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if _, ok := table.Indexes["idx_email"]; ok {
|
||||||
|
t.Error("renamed index should be gone under old name")
|
||||||
|
}
|
||||||
|
if idx := table.Indexes["idx_email2"]; idx == nil || idx.Table != "users" || idx.Schema != "public" {
|
||||||
|
t.Errorf("unexpected renamed index: %+v", idx)
|
||||||
|
}
|
||||||
|
if !se.DeleteIndex(0, 0, "idx_email2") || se.DeleteIndex(0, 0, "idx_email2") {
|
||||||
|
t.Error("delete should succeed once")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSaveViewSequenceScript(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
schema := se.db.Schemas[0]
|
||||||
|
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 2"}); err == nil {
|
||||||
|
t.Error("duplicate view accepted")
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, 0, &models.View{Name: "v", Definition: "select 3"}); err != nil {
|
||||||
|
t.Errorf("editing in place should not conflict: %v", err)
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "w"}); err == nil {
|
||||||
|
t.Error("view without definition accepted")
|
||||||
|
}
|
||||||
|
if len(schema.Views) != 1 || schema.Views[0].Definition != "select 3" || schema.Views[0].Schema != "public" {
|
||||||
|
t.Errorf("unexpected views: %+v", schema.Views)
|
||||||
|
}
|
||||||
|
if !se.DeleteView(0, 0) || se.DeleteView(0, 0) {
|
||||||
|
t.Error("view delete mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "z"}); err == nil {
|
||||||
|
t.Error("zero increment accepted")
|
||||||
|
}
|
||||||
|
if !se.DeleteSequence(0, 0) || len(schema.Sequences) != 0 {
|
||||||
|
t.Error("sequence delete failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := se.SaveScript(0, -1, &models.Script{Name: "init", SQL: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveScript(0, -1, &models.Script{Name: "empty"}); err == nil {
|
||||||
|
t.Error("script without SQL accepted")
|
||||||
|
}
|
||||||
|
if err := se.SaveScript(5, -1, &models.Script{Name: "x", SQL: "y"}); err == nil {
|
||||||
|
t.Error("bad schema index accepted")
|
||||||
|
}
|
||||||
|
if !se.DeleteScript(0, 0) || len(schema.Scripts) != 0 {
|
||||||
|
t.Error("script delete failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDomainTableAssignment(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
se.createDomainNoUI("core")
|
||||||
|
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "users"); err == nil {
|
||||||
|
t.Error("duplicate assignment accepted")
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "missing"); err == nil {
|
||||||
|
t.Error("unknown table accepted")
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(3, "public", "users"); err == nil {
|
||||||
|
t.Error("bad domain index accepted")
|
||||||
|
}
|
||||||
|
dt := se.db.Domains[0].Tables[0]
|
||||||
|
if dt.RefTable != se.db.Schemas[0].Tables[0] {
|
||||||
|
t.Error("RefTable not linked")
|
||||||
|
}
|
||||||
|
if !se.UnassignTableFromDomain(0, "public", "users") || se.UnassignTableFromDomain(0, "public", "users") {
|
||||||
|
t.Error("unassign mismatch")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) createDomainNoUI(name string) {
|
||||||
|
se.db.Domains = append(se.db.Domains, models.InitDomain(name))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSplitList(t *testing.T) {
|
||||||
|
got := splitList(" a, b,, c ,")
|
||||||
|
if len(got) != 3 || got[0] != "a" || got[2] != "c" {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,476 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gdamore/tcell/v2"
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// objectLocation identifies where a new object is created: a schema, and for indexes also a table.
|
||||||
|
type objectLocation struct {
|
||||||
|
label string
|
||||||
|
schemaIndex int
|
||||||
|
tableIndex int
|
||||||
|
}
|
||||||
|
|
||||||
|
// objectRow is one existing object shown in an object list.
|
||||||
|
type objectRow struct {
|
||||||
|
cells []string
|
||||||
|
schemaIndex int
|
||||||
|
tableIndex int
|
||||||
|
at int // position within the schema slice (views, sequences, scripts)
|
||||||
|
name string // map key (indexes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// objectKind describes how a kind of schema object is listed and edited.
|
||||||
|
type objectKind struct {
|
||||||
|
page string
|
||||||
|
title string
|
||||||
|
singular string
|
||||||
|
headers []string
|
||||||
|
rows func() []objectRow
|
||||||
|
locations func() []objectLocation
|
||||||
|
// buildForm adds the editable fields to the form for row (nil when creating) and
|
||||||
|
// returns a function that validates and saves the values at the given location.
|
||||||
|
buildForm func(form *tview.Form, row *objectRow) func(loc objectLocation) error
|
||||||
|
remove func(row objectRow) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) schemaLocations() []objectLocation {
|
||||||
|
locs := make([]objectLocation, 0, len(se.db.Schemas))
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
locs = append(locs, objectLocation{label: s.Name, schemaIndex: si, tableIndex: -1})
|
||||||
|
}
|
||||||
|
return locs
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) tableLocations() []objectLocation {
|
||||||
|
locs := make([]objectLocation, 0)
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for ti, t := range s.Tables {
|
||||||
|
locs = append(locs, objectLocation{label: s.Name + "." + t.Name, schemaIndex: si, tableIndex: ti})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return locs
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) indexKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "indexes",
|
||||||
|
title: "Manage Indexes",
|
||||||
|
singular: "Index",
|
||||||
|
headers: []string{"Name", "Schema", "Table", "Type", "Unique", "Columns"},
|
||||||
|
locations: se.tableLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for ti, t := range s.Tables {
|
||||||
|
for _, name := range sortedKeys(t.Indexes) {
|
||||||
|
idx := t.Indexes[name]
|
||||||
|
rows = append(rows, objectRow{
|
||||||
|
cells: []string{idx.Name, s.Name, t.Name, idx.Type, strconv.FormatBool(idx.Unique), strings.Join(idx.Columns, ",")},
|
||||||
|
schemaIndex: si, tableIndex: ti, name: name,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
idx := models.InitIndex("", "", "")
|
||||||
|
idx.Type = "btree"
|
||||||
|
if row != nil {
|
||||||
|
idx = se.db.Schemas[row.schemaIndex].Tables[row.tableIndex].Indexes[row.name]
|
||||||
|
}
|
||||||
|
name, columns, typ, where := idx.Name, strings.Join(idx.Columns, ", "), idx.Type, idx.Where
|
||||||
|
unique := idx.Unique
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Columns (comma separated)", columns, 50, nil, func(v string) { columns = v })
|
||||||
|
form.AddInputField("Type", typ, 20, nil, func(v string) { typ = v })
|
||||||
|
form.AddCheckbox("Unique", unique, func(v bool) { unique = v })
|
||||||
|
form.AddInputField("Where", where, 50, nil, func(v string) { where = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
oldName := ""
|
||||||
|
if row != nil {
|
||||||
|
oldName = row.name
|
||||||
|
}
|
||||||
|
next := *idx
|
||||||
|
next.Name, next.Columns, next.Type, next.Unique, next.Where = name, splitList(columns), typ, unique, where
|
||||||
|
return se.SaveIndex(loc.schemaIndex, loc.tableIndex, oldName, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteIndex(r.schemaIndex, r.tableIndex, r.name) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) viewKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "views",
|
||||||
|
title: "Manage Views",
|
||||||
|
singular: "View",
|
||||||
|
headers: []string{"Name", "Schema", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, v := range s.Views {
|
||||||
|
rows = append(rows, objectRow{cells: []string{v.Name, s.Name, v.Description}, schemaIndex: si, at: i})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
view := models.InitView("", "")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
view, at = se.db.Schemas[row.schemaIndex].Views[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc, def := view.Name, view.Description, view.Definition
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddTextArea("Definition (SQL)", def, 60, 8, 0, func(v string) { def = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
next := *view
|
||||||
|
next.Name, next.Description, next.Definition = name, desc, def
|
||||||
|
return se.SaveView(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteView(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) sequenceKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "sequences",
|
||||||
|
title: "Manage Sequences",
|
||||||
|
singular: "Sequence",
|
||||||
|
headers: []string{"Name", "Schema", "Start", "Increment", "Cycle", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, q := range s.Sequences {
|
||||||
|
rows = append(rows, objectRow{
|
||||||
|
cells: []string{q.Name, s.Name, strconv.FormatInt(q.StartValue, 10), strconv.FormatInt(q.IncrementBy, 10), strconv.FormatBool(q.Cycle), q.Description}, schemaIndex: si, at: i,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
seq := models.InitSequence("", "")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
seq, at = se.db.Schemas[row.schemaIndex].Sequences[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc := seq.Name, seq.Description
|
||||||
|
start, incr := strconv.FormatInt(seq.StartValue, 10), strconv.FormatInt(seq.IncrementBy, 10)
|
||||||
|
minV, maxV := strconv.FormatInt(seq.MinValue, 10), strconv.FormatInt(seq.MaxValue, 10)
|
||||||
|
cycle := seq.Cycle
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddInputField("Start", start, 20, nil, func(v string) { start = v })
|
||||||
|
form.AddInputField("Increment", incr, 20, nil, func(v string) { incr = v })
|
||||||
|
form.AddInputField("Min (0 = none)", minV, 20, nil, func(v string) { minV = v })
|
||||||
|
form.AddInputField("Max (0 = none)", maxV, 20, nil, func(v string) { maxV = v })
|
||||||
|
form.AddCheckbox("Cycle", cycle, func(v bool) { cycle = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
next := *seq
|
||||||
|
next.Name, next.Description, next.Cycle = name, desc, cycle
|
||||||
|
for _, f := range []struct {
|
||||||
|
label string
|
||||||
|
text string
|
||||||
|
dst *int64
|
||||||
|
}{{"start", start, &next.StartValue}, {"increment", incr, &next.IncrementBy}, {"min", minV, &next.MinValue}, {"max", maxV, &next.MaxValue}} {
|
||||||
|
n, err := strconv.ParseInt(strings.TrimSpace(f.text), 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s must be an integer", f.label)
|
||||||
|
}
|
||||||
|
*f.dst = n
|
||||||
|
}
|
||||||
|
return se.SaveSequence(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteSequence(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) scriptKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "scripts",
|
||||||
|
title: "Manage Scripts",
|
||||||
|
singular: "Script",
|
||||||
|
headers: []string{"Name", "Schema", "Version", "Priority", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, sc := range s.Scripts {
|
||||||
|
rows = append(rows, objectRow{cells: []string{sc.Name, s.Name, sc.Version, strconv.Itoa(sc.Priority), sc.Description}, schemaIndex: si, at: i})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
script := models.InitScript("")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
script, at = se.db.Schemas[row.schemaIndex].Scripts[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc, version, sql, rollback := script.Name, script.Description, script.Version, script.SQL, script.Rollback
|
||||||
|
priority, runAfter := strconv.Itoa(script.Priority), strings.Join(script.RunAfter, ", ")
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddInputField("Version", version, 20, nil, func(v string) { version = v })
|
||||||
|
form.AddInputField("Priority", priority, 10, nil, func(v string) { priority = v })
|
||||||
|
form.AddInputField("Run after (comma separated)", runAfter, 50, nil, func(v string) { runAfter = v })
|
||||||
|
form.AddTextArea("SQL", sql, 60, 8, 0, func(v string) { sql = v })
|
||||||
|
form.AddTextArea("Rollback SQL", rollback, 60, 4, 0, func(v string) { rollback = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
prio, err := strconv.Atoi(strings.TrimSpace(priority))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("priority must be an integer")
|
||||||
|
}
|
||||||
|
next := *script
|
||||||
|
next.Name, next.Description, next.Version, next.Priority = name, desc, version, prio
|
||||||
|
next.RunAfter, next.SQL, next.Rollback = splitList(runAfter), sql, rollback
|
||||||
|
return se.SaveScript(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteScript(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedKeys[V any](m map[string]V) []string {
|
||||||
|
keys := make([]string, 0, len(m))
|
||||||
|
for k := range m {
|
||||||
|
keys = append(keys, k)
|
||||||
|
}
|
||||||
|
sort.Strings(keys)
|
||||||
|
return keys
|
||||||
|
}
|
||||||
|
|
||||||
|
// showObjectList displays all objects of a kind across schemas.
|
||||||
|
func (se *SchemaEditor) showObjectList(k objectKind) {
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow)
|
||||||
|
title := tview.NewTextView().SetText("[::b]" + k.title).SetDynamicColors(true).SetTextAlign(tview.AlignCenter)
|
||||||
|
|
||||||
|
table := tview.NewTable().SetBorders(true).SetSelectable(true, false).SetFixed(1, 0)
|
||||||
|
for i, h := range k.headers {
|
||||||
|
table.SetCell(0, i, tview.NewTableCell(h).SetTextColor(tcell.ColorYellow).SetSelectable(false).SetAlign(tview.AlignLeft))
|
||||||
|
}
|
||||||
|
rows := k.rows()
|
||||||
|
for r, row := range rows {
|
||||||
|
for c, text := range row.cells {
|
||||||
|
table.SetCell(r+1, c, tview.NewTableCell(text).SetSelectable(true))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
table.SetTitle(" " + k.title[len("Manage "):] + " ").SetBorder(true).SetTitleAlign(tview.AlignLeft)
|
||||||
|
|
||||||
|
back := func() {
|
||||||
|
se.pages.SwitchToPage("main")
|
||||||
|
se.pages.RemovePage(k.page)
|
||||||
|
}
|
||||||
|
btnNew := tview.NewButton("New " + k.singular + " [n]").SetSelectedFunc(func() { se.showObjectForm(k, nil) })
|
||||||
|
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(back)
|
||||||
|
btnNew.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
se.app.SetFocus(table)
|
||||||
|
return nil
|
||||||
|
case tcell.KeyTab:
|
||||||
|
se.app.SetFocus(btnBack)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
btnBack.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
se.app.SetFocus(btnNew)
|
||||||
|
return nil
|
||||||
|
case tcell.KeyTab:
|
||||||
|
se.app.SetFocus(table)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
btnFlex := tview.NewFlex().AddItem(btnNew, 0, 1, true).AddItem(btnBack, 0, 1, false)
|
||||||
|
|
||||||
|
table.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch {
|
||||||
|
case event.Key() == tcell.KeyEscape, event.Rune() == 'b':
|
||||||
|
back()
|
||||||
|
return nil
|
||||||
|
case event.Key() == tcell.KeyTab:
|
||||||
|
se.app.SetFocus(btnNew)
|
||||||
|
return nil
|
||||||
|
case event.Key() == tcell.KeyEnter:
|
||||||
|
if row, _ := table.GetSelection(); row > 0 && row <= len(rows) {
|
||||||
|
se.showObjectForm(k, &rows[row-1])
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
case event.Rune() == 'n':
|
||||||
|
se.showObjectForm(k, nil)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
flex.AddItem(title, 1, 0, false).AddItem(table, 0, 1, true).AddItem(btnFlex, 1, 0, false)
|
||||||
|
se.pages.AddPage(k.page, flex, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showObjectForm shows the create (row == nil) or edit form for an object.
|
||||||
|
func (se *SchemaEditor) showObjectForm(k objectKind, row *objectRow) {
|
||||||
|
formPage := k.page + "-form"
|
||||||
|
form := tview.NewForm()
|
||||||
|
errView := tview.NewTextView().SetDynamicColors(true)
|
||||||
|
|
||||||
|
locs := k.locations()
|
||||||
|
loc := objectLocation{schemaIndex: -1, tableIndex: -1}
|
||||||
|
switch {
|
||||||
|
case row != nil:
|
||||||
|
loc = objectLocation{schemaIndex: row.schemaIndex, tableIndex: row.tableIndex}
|
||||||
|
case len(locs) > 0:
|
||||||
|
loc = locs[0]
|
||||||
|
labels := make([]string, len(locs))
|
||||||
|
for i, l := range locs {
|
||||||
|
labels[i] = l.label
|
||||||
|
}
|
||||||
|
form.AddDropDown("Location", labels, 0, func(_ string, i int) { loc = locs[i] })
|
||||||
|
}
|
||||||
|
|
||||||
|
save := k.buildForm(form, row)
|
||||||
|
|
||||||
|
closeForm := func() {
|
||||||
|
se.pages.RemovePage(formPage)
|
||||||
|
se.pages.RemovePage(k.page)
|
||||||
|
se.showObjectList(k)
|
||||||
|
}
|
||||||
|
form.AddButton("Save", func() {
|
||||||
|
if err := save(loc); err != nil {
|
||||||
|
errView.SetText("[red]" + tview.Escape(err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
closeForm()
|
||||||
|
})
|
||||||
|
if row != nil {
|
||||||
|
form.AddButton("Delete", func() {
|
||||||
|
modal := tview.NewModal().
|
||||||
|
SetText(fmt.Sprintf("Delete %s '%s'? This action cannot be undone.", strings.ToLower(k.singular), row.cells[0])).
|
||||||
|
AddButtons([]string{"Cancel", "Delete"}).
|
||||||
|
SetDoneFunc(func(_ int, label string) {
|
||||||
|
se.pages.RemovePage(formPage + "-delete")
|
||||||
|
if label == "Delete" {
|
||||||
|
k.remove(*row)
|
||||||
|
closeForm()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
se.pages.AddAndSwitchToPage(formPage+"-delete", modal, true)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
form.AddButton("Back", closeForm)
|
||||||
|
|
||||||
|
verb := "New"
|
||||||
|
if row != nil {
|
||||||
|
verb = "Edit"
|
||||||
|
}
|
||||||
|
form.SetBorder(true).SetTitle(" " + verb + " " + k.singular + " ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.showExitConfirmation(formPage, k.page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(locs) == 0 && row == nil {
|
||||||
|
errView.SetText("[red]No schema/table available. Create one first.")
|
||||||
|
}
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow).AddItem(form, 0, 1, true).AddItem(errView, 1, 0, false)
|
||||||
|
se.pages.AddPage(formPage, flex, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showDomainTables lists the tables assigned to a domain and allows assigning/unassigning.
|
||||||
|
func (se *SchemaEditor) showDomainTables(domainIndex int) {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
page := "domain-tables"
|
||||||
|
list := tview.NewList().ShowSecondaryText(true)
|
||||||
|
refresh := func() {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
se.showDomainTables(domainIndex)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
dt := dt
|
||||||
|
list.AddItem(dt.SchemaName+"."+dt.TableName, "Enter to remove from domain", 0, func() {
|
||||||
|
se.UnassignTableFromDomain(domainIndex, dt.SchemaName, dt.TableName)
|
||||||
|
refresh()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
list.AddItem("[Assign Table]", "Add a table to this domain", 'a', func() {
|
||||||
|
se.showAssignDomainTable(domainIndex, refresh)
|
||||||
|
})
|
||||||
|
list.AddItem("[Back]", "Return to domain", 'b', func() {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
})
|
||||||
|
list.SetBorder(true).SetTitle(" Domain " + domain.Name + " - Tables ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
list.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddPage(page, list, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showAssignDomainTable shows a form to pick a table not yet in the domain.
|
||||||
|
func (se *SchemaEditor) showAssignDomainTable(domainIndex int, done func()) {
|
||||||
|
page := "assign-domain-table"
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
var options []string
|
||||||
|
var refs []models.DomainTable
|
||||||
|
for _, s := range se.db.Schemas {
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
taken := false
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
taken = taken || (dt.SchemaName == s.Name && dt.TableName == t.Name)
|
||||||
|
}
|
||||||
|
if !taken {
|
||||||
|
options = append(options, s.Name+"."+t.Name)
|
||||||
|
refs = append(refs, models.DomainTable{SchemaName: s.Name, TableName: t.Name})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
form := tview.NewForm()
|
||||||
|
selected := 0
|
||||||
|
form.AddDropDown("Table", options, 0, func(_ string, i int) { selected = i })
|
||||||
|
form.AddButton("Assign", func() {
|
||||||
|
if len(refs) > 0 {
|
||||||
|
_ = se.AssignTableToDomain(domainIndex, refs[selected].SchemaName, refs[selected].TableName)
|
||||||
|
}
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
done()
|
||||||
|
})
|
||||||
|
form.AddButton("Back", func() { se.pages.RemovePage(page) })
|
||||||
|
form.SetBorder(true).SetTitle(" Assign Table ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddPage(page, form, true, true)
|
||||||
|
}
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// richEditor returns an editor whose database has one of every object the screens render.
|
||||||
|
func richEditor(t *testing.T) *SchemaEditor {
|
||||||
|
t.Helper()
|
||||||
|
se := NewSchemaEditor(newTestEditor().db)
|
||||||
|
db := se.db
|
||||||
|
tbl := db.Schemas[0].Tables[0]
|
||||||
|
col := tbl.Columns["id"]
|
||||||
|
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||||
|
tbl.Relationships["fk_self"] = &models.Relationship{Name: "fk_self", FromTable: "users", ToTable: "users", FromColumns: []string{"id"}, ToColumns: []string{"id"}}
|
||||||
|
se.createDomainNoUI("core")
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}})
|
||||||
|
_ = se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"})
|
||||||
|
_ = se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1})
|
||||||
|
_ = se.SaveScript(0, -1, &models.Script{Name: "sc", SQL: "select 1"})
|
||||||
|
return se
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestScreensRender builds every screen and dialog against a populated database
|
||||||
|
// and checks that none panics and that each registers a page.
|
||||||
|
func TestScreensRender(t *testing.T) {
|
||||||
|
col := func(se *SchemaEditor) *models.Column { return se.db.Schemas[0].Tables[0].Columns["id"] }
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
page string
|
||||||
|
run func(se *SchemaEditor)
|
||||||
|
}{
|
||||||
|
{"schema list", "schemas", func(se *SchemaEditor) { se.showSchemaList() }},
|
||||||
|
{"schema editor", "schema-editor", func(se *SchemaEditor) { se.showSchemaEditor(0, se.db.Schemas[0]) }},
|
||||||
|
{"new schema", "new-schema", func(se *SchemaEditor) { se.showNewSchemaDialog() }},
|
||||||
|
{"edit schema", "edit-schema", func(se *SchemaEditor) { se.showEditSchemaDialog(0) }},
|
||||||
|
{"table list", "tables", func(se *SchemaEditor) { se.showTableList() }},
|
||||||
|
{"table editor", "table-editor", func(se *SchemaEditor) { se.showTableEditor(0, 0, se.db.Schemas[0].Tables[0]) }},
|
||||||
|
{"new table", "new-table", func(se *SchemaEditor) { se.showNewTableDialog(0) }},
|
||||||
|
{"new table from list", "new-table-from-list", func(se *SchemaEditor) { se.showNewTableDialogFromList() }},
|
||||||
|
{"edit table", "edit-table", func(se *SchemaEditor) { se.showEditTableDialog(0, 0) }},
|
||||||
|
{"column editor", "column-editor", func(se *SchemaEditor) { se.showColumnEditor(0, 0, 0, col(se)) }},
|
||||||
|
{"new column", "new-column", func(se *SchemaEditor) { se.showNewColumnDialog(0, 0) }},
|
||||||
|
{"relationship list", "relationships", func(se *SchemaEditor) { se.showRelationshipList(0, 0) }},
|
||||||
|
{"new relationship", "new-relationship", func(se *SchemaEditor) { se.showNewRelationshipDialog(0, 0) }},
|
||||||
|
{"edit relationship", "edit-relationship", func(se *SchemaEditor) { se.showEditRelationshipDialog(0, 0, "fk_self") }},
|
||||||
|
{"delete relationship", "delete-relationship-confirm", func(se *SchemaEditor) { se.showDeleteRelationshipConfirm(0, 0, "fk_self") }},
|
||||||
|
{"domain list", "domains", func(se *SchemaEditor) { se.showDomainList() }},
|
||||||
|
{"new domain", "new-domain", func(se *SchemaEditor) { se.showNewDomainDialog() }},
|
||||||
|
{"domain editor", "edit-domain", func(se *SchemaEditor) { se.showDomainEditor(0, se.db.Domains[0]) }},
|
||||||
|
{"delete domain", "delete-domain-confirm", func(se *SchemaEditor) { se.showDeleteDomainConfirm(0) }},
|
||||||
|
{"domain tables", "domain-tables", func(se *SchemaEditor) { se.showDomainTables(0) }},
|
||||||
|
{"assign domain table", "assign-domain-table", func(se *SchemaEditor) { se.showAssignDomainTable(0, func() {}) }},
|
||||||
|
{"edit database", "edit-database", func(se *SchemaEditor) { se.showEditDatabaseForm() }},
|
||||||
|
{"exit confirm", "exit-confirm", func(se *SchemaEditor) { se.showExitConfirmation("a", "main") }},
|
||||||
|
{"exit editor confirm", "exit-editor-confirm", func(se *SchemaEditor) { se.showExitEditorConfirm() }},
|
||||||
|
{"delete schema confirm", "confirm-delete-schema", func(se *SchemaEditor) { se.showDeleteSchemaConfirm(0) }},
|
||||||
|
{"delete table confirm", "confirm-delete-table", func(se *SchemaEditor) { se.showDeleteTableConfirm(0, 0) }},
|
||||||
|
{"delete column confirm", "confirm-delete-column", func(se *SchemaEditor) { se.showDeleteColumnConfirm(0, 0, "id") }},
|
||||||
|
{"load screen", "load-database", func(se *SchemaEditor) { se.showLoadScreen() }},
|
||||||
|
{"save screen", "save-database", func(se *SchemaEditor) { se.showSaveScreen() }},
|
||||||
|
{"import screen", "import-database", func(se *SchemaEditor) { se.showImportScreen() }},
|
||||||
|
{"update existing confirm", "update-confirm", func(se *SchemaEditor) {
|
||||||
|
se.loadConfig = &LoadConfig{SourceType: "json", FilePath: "x.json"}
|
||||||
|
se.showUpdateExistingDatabaseConfirm()
|
||||||
|
}},
|
||||||
|
{"import confirm", "import-confirm", func(se *SchemaEditor) {
|
||||||
|
se.showImportConfirmation(models.InitDatabase("src"), false, false, false, false, false, "")
|
||||||
|
}},
|
||||||
|
{"conn builder", "", func(se *SchemaEditor) { se.showConnStringBuilder("", "", "main", func(string) {}) }},
|
||||||
|
}
|
||||||
|
for _, tt := range cases {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
se := richEditor(t)
|
||||||
|
before := len(se.pages.GetPageNames(false))
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
t.Fatalf("panic: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
tt.run(se)
|
||||||
|
if tt.page != "" && !se.pages.HasPage(tt.page) {
|
||||||
|
t.Errorf("page %q not registered; pages: %v", tt.page, se.pages.GetPageNames(false))
|
||||||
|
}
|
||||||
|
if tt.page == "" && len(se.pages.GetPageNames(false)) <= before {
|
||||||
|
t.Error("no page added")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestObjectScreensRender(t *testing.T) {
|
||||||
|
se := richEditor(t)
|
||||||
|
for name, k := range map[string]objectKind{
|
||||||
|
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
se.showObjectList(k)
|
||||||
|
if !se.pages.HasPage(k.page) {
|
||||||
|
t.Errorf("list page %q missing; pages: %v", k.page, se.pages.GetPageNames(false))
|
||||||
|
}
|
||||||
|
rows := k.rows()
|
||||||
|
se.showObjectForm(k, nil)
|
||||||
|
se.showObjectForm(k, &rows[0])
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -205,12 +205,24 @@ Organize UI code into these files:
|
|||||||
- **column_screens.go** - Column editor, new column dialog
|
- **column_screens.go** - Column editor, new column dialog
|
||||||
- **domain_screens.go** - Domain list, domain editor, new/edit domain dialogs
|
- **domain_screens.go** - Domain list, domain editor, new/edit domain dialogs
|
||||||
- **dialogs.go** - Confirmation dialogs (exit, delete)
|
- **dialogs.go** - Confirmation dialogs (exit, delete)
|
||||||
|
- **filebrowser_screens.go** - File browser dialog (`file-browser`), opened with Enter on File Path inputs
|
||||||
|
- **connstring_screens.go** - Connection string builder dialog (`conn-builder`), opened with Enter on Connection String inputs
|
||||||
|
|
||||||
### Data Operations Files (Business Logic)
|
### Data Operations Files (Business Logic)
|
||||||
|
|
||||||
- **schema_dataops.go** - Schema CRUD operations (Create, Read, Update, Delete)
|
- **schema_dataops.go** - Schema CRUD operations (Create, Read, Update, Delete)
|
||||||
- **table_dataops.go** - Table CRUD operations
|
- **table_dataops.go** - Table CRUD operations
|
||||||
- **column_dataops.go** - Column CRUD operations
|
- **column_dataops.go** - Column CRUD operations
|
||||||
|
- **filebrowser.go** - Directory listing, extension filtering and start-path resolution (no tview)
|
||||||
|
- **connstring.go**, **connstring_check.go** - Connection string build/parse/mask and connection test (no tview)
|
||||||
|
|
||||||
|
### Input Dialogs
|
||||||
|
|
||||||
|
- **File browser** - Enter on a File Path input opens it. Up/Down move, Enter opens a directory or selects a file,
|
||||||
|
Backspace/`[..]` goes to the parent, `h` toggles hidden files, `f` toggles the format extension filter,
|
||||||
|
`s` selects, Esc/`b` cancels (input unchanged). Save mode adds a file name field and asks before overwriting.
|
||||||
|
- **Connection string builder** - Enter on a Connection String input opens it. Fields per type (PostgreSQL, MSSQL,
|
||||||
|
SQLite), masked password and preview, F2 save, F3 test connection, Esc cancels (input unchanged).
|
||||||
|
|
||||||
## Code Separation Rules
|
## Code Separation Rules
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
package bun
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestSnakeCaseToCamelCase(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"user", "user"},
|
||||||
|
{"User_Name", "userName"},
|
||||||
|
{"user_id", "userID"},
|
||||||
|
{"http_request", "httpRequest"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPascalCaseToSnakeCase(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"User", "user"},
|
||||||
|
{"UserName", "user_name"},
|
||||||
|
{"UserID", "user_id"},
|
||||||
|
{"HTTPRequest", "http_request"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSingularize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"people", "person"},
|
||||||
|
{"People", "person"},
|
||||||
|
{"categories", "category"},
|
||||||
|
{"wolves", "wolf"},
|
||||||
|
{"boxes", "box"},
|
||||||
|
{"churches", "church"},
|
||||||
|
{"users", "user"},
|
||||||
|
{"class", "class"},
|
||||||
|
{"user", "user"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Singularize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPluralize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"person", "people"},
|
||||||
|
{"category", "categories"},
|
||||||
|
{"box", "boxes"},
|
||||||
|
{"church", "churches"},
|
||||||
|
{"user", "users"},
|
||||||
|
{"day", "days"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Pluralize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsVowel(t *testing.T) {
|
||||||
|
for _, c := range []byte("aeiouAEIOU") {
|
||||||
|
if !isVowel(c) {
|
||||||
|
t.Errorf("%c should be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, c := range []byte("bcxyzBZ1_") {
|
||||||
|
if isVowel(c) {
|
||||||
|
t.Errorf("%c should not be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types for Bun
|
// TypeMapper handles type conversions between SQL and Go types for Bun
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
|
typeMappings map[string]string
|
||||||
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
||||||
}
|
}
|
||||||
@@ -37,6 +38,10 @@ func NewTypeMapper(typeStyle, arrayNullable string) *TypeMapper {
|
|||||||
|
|
||||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
||||||
|
if goType, ok := tm.overrideGoType(sqlType, notNull); ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
|
||||||
// Array columns always use a native Go slice, regardless of typeStyle.
|
// Array columns always use a native Go slice, regardless of typeStyle.
|
||||||
if pgsql.IsArrayType(sqlType) {
|
if pgsql.IsArrayType(sqlType) {
|
||||||
goType := tm.arrayGoType(tm.extractBaseType(sqlType))
|
goType := tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||||
@@ -68,6 +73,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
|||||||
return tm.bunGoType(baseType)
|
return tm.bunGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetTypeMappings installs user-configured SQL-to-Go type overrides.
|
||||||
|
func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) {
|
||||||
|
tm.typeMappings = mappings
|
||||||
|
}
|
||||||
|
|
||||||
|
// overrideGoType applies a configured override, if any, for the column type.
|
||||||
|
func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) {
|
||||||
|
if len(tm.typeMappings) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType))
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if pgsql.IsArrayType(sqlType) {
|
||||||
|
return "[]" + goType, true
|
||||||
|
}
|
||||||
|
return writers.ApplyTypeMapping(goType, notNull), true
|
||||||
|
}
|
||||||
|
|
||||||
// extractBaseType extracts the base type from a SQL type string
|
// extractBaseType extracts the base type from a SQL type string
|
||||||
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
||||||
return pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
return pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package bun
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestTypeMapper_CustomTypeMappings(t *testing.T) {
|
||||||
|
mapper := NewTypeMapper("", "")
|
||||||
|
mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"uuid", true, "uuid.UUID"},
|
||||||
|
{"uuid", false, "*uuid.UUID"},
|
||||||
|
{"numeric(10,2)", true, "decimal.Decimal"},
|
||||||
|
{"UUID[]", false, "[]uuid.UUID"},
|
||||||
|
{"bigint", true, "int64"}, // unmapped types keep defaults
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want {
|
||||||
|
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
package bun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSQLTypeToGoType_Styles(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
style string
|
||||||
|
arrays string
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{writers.NullableTypeSqlTypes, "", "integer", true, "int32"},
|
||||||
|
{writers.NullableTypeSqlTypes, "", "bigint", true, "int64"},
|
||||||
|
{writers.NullableTypeSqlTypes, "", "text", true, "sql_types.SqlString"},
|
||||||
|
{writers.NullableTypeSqlTypes, "", "boolean", true, "bool"},
|
||||||
|
{writers.NullableTypeSqlTypes, "", "bigint", false, "sql_types.SqlInt64"},
|
||||||
|
{writers.NullableTypeSqlTypes, "", "text", false, "sql_types.SqlString"},
|
||||||
|
{writers.NullableTypeStdlib, "", "integer", true, "int32"},
|
||||||
|
{writers.NullableTypeStdlib, "", "integer", false, "sql.NullInt32"},
|
||||||
|
{writers.NullableTypeStdlib, "", "bigint", false, "sql.NullInt64"},
|
||||||
|
{writers.NullableTypeStdlib, "", "boolean", false, "sql.NullBool"},
|
||||||
|
{writers.NullableTypeStdlib, "", "text", false, "sql.NullString"},
|
||||||
|
{writers.NullableTypeStdlib, "", "timestamptz", false, "sql.NullTime"},
|
||||||
|
{writers.NullableTypeStdlib, "", "mystery", false, "sql.NullString"},
|
||||||
|
{writers.NullableTypeBaselib, "", "integer", false, "*int32"},
|
||||||
|
{"", "", "text", false, "*string"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
|
||||||
|
got := NewTypeMapper(tt.style, tt.arrays).SQLTypeToGoType(tt.sqlType, tt.notNull)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("got %q want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSQLTypeToGoType_Arrays(t *testing.T) {
|
||||||
|
slice := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysSlice)
|
||||||
|
ptr := NewTypeMapper(writers.NullableTypeSqlTypes, writers.NullableArraysPointerSlice)
|
||||||
|
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]"} {
|
||||||
|
s := slice.SQLTypeToGoType(sqlType, false)
|
||||||
|
p := ptr.SQLTypeToGoType(sqlType, false)
|
||||||
|
if s == "" || strings.HasPrefix(s, "*") {
|
||||||
|
t.Errorf("%s slice mode: %q", sqlType, s)
|
||||||
|
}
|
||||||
|
if p != "*"+s {
|
||||||
|
t.Errorf("%s pointer mode: %q want %q", sqlType, p, "*"+s)
|
||||||
|
}
|
||||||
|
if got := ptr.SQLTypeToGoType(sqlType, true); got != s {
|
||||||
|
t.Errorf("%s not-null should stay a slice: %q", sqlType, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImportHelpers(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
style string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{writers.NullableTypeStdlib, `"database/sql"`},
|
||||||
|
{writers.NullableTypeBaselib, ""},
|
||||||
|
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := NewTypeMapper(tt.style, "").GetNullableTypeImportLine(); got != tt.want {
|
||||||
|
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tm := NewTypeMapper("", "")
|
||||||
|
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
|
||||||
|
t.Error("NeedsFmtImport should echo its argument")
|
||||||
|
}
|
||||||
|
if tm.GetSQLTypesImport() == "" || tm.GetBunImport() != "github.com/uptrace/bun" {
|
||||||
|
t.Error("unexpected imports")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer {
|
|||||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
w.typeMapper.SetTypeMappings(options.TypeMappings)
|
||||||
|
|
||||||
// Initialize templates
|
// Initialize templates
|
||||||
tmpl, err := NewTemplates()
|
tmpl, err := NewTemplates()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -100,8 +100,8 @@ func (tm *TypeMapper) BuildColumnChain(col *models.Column, table *models.Table,
|
|||||||
// Determine Drizzle column type
|
// Determine Drizzle column type
|
||||||
var drizzleType string
|
var drizzleType string
|
||||||
if isEnum {
|
if isEnum {
|
||||||
// For enum types, use the type name directly
|
// Enum columns call the enum constant declared via pgEnum(...)
|
||||||
drizzleType = fmt.Sprintf("pgEnum('%s')", col.Type)
|
drizzleType = tm.ToCamelCase(col.Type)
|
||||||
} else {
|
} else {
|
||||||
drizzleType = tm.SQLTypeToDrizzle(col.Type)
|
drizzleType = tm.SQLTypeToDrizzle(col.Type)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,300 @@
|
|||||||
|
package drizzle
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const drizzleFixture = "../../../tests/assets/drizzle/schema.ts"
|
||||||
|
|
||||||
|
func fixtureDB(t *testing.T) *models.Database {
|
||||||
|
t.Helper()
|
||||||
|
db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
// shopDB builds a database with an enum, FK, unique, index and varied defaults.
|
||||||
|
func shopDB() *models.Database {
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}})
|
||||||
|
|
||||||
|
users := models.InitTable("users", "public")
|
||||||
|
id := models.InitColumn("id", "users", "public")
|
||||||
|
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
|
||||||
|
email := models.InitColumn("email", "users", "public")
|
||||||
|
email.Type, email.NotNull = "varchar(255)", true
|
||||||
|
role := models.InitColumn("role", "users", "public")
|
||||||
|
role.Type, role.NotNull, role.Default = "role", true, "user"
|
||||||
|
active := models.InitColumn("active", "users", "public")
|
||||||
|
active.Type, active.Default = "boolean", true
|
||||||
|
created := models.InitColumn("created_at", "users", "public")
|
||||||
|
created.Type, created.Default = "timestamp", "now()"
|
||||||
|
score := models.InitColumn("score", "users", "public")
|
||||||
|
score.Type, score.Default = "integer", "10"
|
||||||
|
for _, c := range []*models.Column{id, email, role, active, created, score} {
|
||||||
|
users.Columns[c.Name] = c
|
||||||
|
}
|
||||||
|
uq := models.InitConstraint("uq_email", models.UniqueConstraint)
|
||||||
|
uq.Columns = []string{"email"}
|
||||||
|
users.Constraints["uq_email"] = uq
|
||||||
|
ix := models.InitIndex("idx_role", "users", "public")
|
||||||
|
ix.Columns = []string{"role"}
|
||||||
|
users.Indexes["idx_role"] = ix
|
||||||
|
|
||||||
|
posts := models.InitTable("blog_posts", "public")
|
||||||
|
pid := models.InitColumn("id", "blog_posts", "public")
|
||||||
|
pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true
|
||||||
|
pid.Default = "gen_random_uuid()"
|
||||||
|
author := models.InitColumn("author_id", "blog_posts", "public")
|
||||||
|
author.Type, author.NotNull = "integer", true
|
||||||
|
posts.Columns["id"], posts.Columns["author_id"] = pid, author
|
||||||
|
fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint)
|
||||||
|
fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"}
|
||||||
|
posts.Constraints["fk_author"] = fk
|
||||||
|
|
||||||
|
s.Tables = append(s.Tables, users, posts)
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeSingle(t *testing.T, db *models.Database) string {
|
||||||
|
t.Helper()
|
||||||
|
out := filepath.Join(t.TempDir(), "schema.ts")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSingleFile_Content(t *testing.T) {
|
||||||
|
got := writeSingle(t, shopDB())
|
||||||
|
for _, want := range []string{
|
||||||
|
"drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'",
|
||||||
|
"pgTable('users'", "pgTable('blog_posts'", "blogPosts",
|
||||||
|
".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)",
|
||||||
|
"default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`",
|
||||||
|
"varchar", "uuid(",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("output missing %q\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSingleFile_Deterministic(t *testing.T) {
|
||||||
|
first := writeSingle(t, shopDB())
|
||||||
|
for i := 0; i < 15; i++ {
|
||||||
|
if got := writeSingle(t, shopDB()); got != first {
|
||||||
|
t.Fatalf("output differs on run %d", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFixtureRoundTrip(t *testing.T) {
|
||||||
|
db := fixtureDB(t)
|
||||||
|
got := writeSingle(t, db)
|
||||||
|
if got == "" {
|
||||||
|
t.Fatal("empty output")
|
||||||
|
}
|
||||||
|
out := filepath.Join(t.TempDir(), "again.ts")
|
||||||
|
if err := os.WriteFile(out, []byte(got), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("re-read: %v", err)
|
||||||
|
}
|
||||||
|
count := func(d *models.Database) (n int) {
|
||||||
|
for _, s := range d.Schemas {
|
||||||
|
n += len(s.Tables)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if count(db) != count(again) {
|
||||||
|
t.Errorf("tables: %d -> %d", count(db), count(again))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMultiFile(t *testing.T) {
|
||||||
|
dir := filepath.Join(t.TempDir(), "schema")
|
||||||
|
w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}})
|
||||||
|
if err := w.WriteDatabase(shopDB()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} {
|
||||||
|
if _, err := os.Stat(filepath.Join(dir, f)); err != nil {
|
||||||
|
t.Errorf("%s not written: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
users, _ := os.ReadFile(filepath.Join(dir, "users.ts"))
|
||||||
|
if !strings.Contains(string(users), "from './enums'") {
|
||||||
|
t.Errorf("users.ts must import its enum:\n%s", users)
|
||||||
|
}
|
||||||
|
posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts"))
|
||||||
|
if strings.Contains(string(posts), "from './enums'") {
|
||||||
|
t.Errorf("blog_posts.ts uses no enum:\n%s", posts)
|
||||||
|
}
|
||||||
|
enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts"))
|
||||||
|
if !strings.Contains(string(enums), "pgEnum('role'") {
|
||||||
|
t.Errorf("enums.ts:\n%s", enums)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMultiFile_RequiresOutputPath(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}})
|
||||||
|
if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShouldUseMultiFile(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
opts writers.WriterOptions
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"stdout", writers.WriterOptions{}, false},
|
||||||
|
{"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true},
|
||||||
|
{"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false},
|
||||||
|
{"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false},
|
||||||
|
{"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true},
|
||||||
|
{"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true},
|
||||||
|
{"existing dir", writers.WriterOptions{OutputPath: dir}, true},
|
||||||
|
{"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
opts := tt.opts
|
||||||
|
if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want {
|
||||||
|
t.Errorf("got %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteSchemaAndTable(t *testing.T) {
|
||||||
|
db := shopDB()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
sOut := filepath.Join(dir, "s.ts")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tOut := filepath.Join(dir, "t.ts")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, _ := os.ReadFile(tOut)
|
||||||
|
if !strings.Contains(string(b), "pgTable('users'") {
|
||||||
|
t.Errorf("table output:\n%s", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "missing", "x.ts")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil {
|
||||||
|
t.Error("expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatDefaultValue(t *testing.T) {
|
||||||
|
tm := NewTypeMapper()
|
||||||
|
tests := []struct {
|
||||||
|
in any
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"now()", "sql`now()`"}, {"CURRENT_TIMESTAMP", "sql`now()`"},
|
||||||
|
{"gen_random_uuid()", "sql`gen_random_uuid()`"}, {"uuid_generate_v4()", "sql`gen_random_uuid()`"},
|
||||||
|
{"42", "42"}, {"-1.5", "-1.5"}, {"it's", `'it\'s'`}, {"plain", "'plain'"},
|
||||||
|
{true, "true"}, {false, "false"},
|
||||||
|
{7, "7"}, {int64(8), "8"}, {2.5, "2.5"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := tm.formatDefaultValue(tt.in); got != tt.want {
|
||||||
|
t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsNumericString(t *testing.T) {
|
||||||
|
for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} {
|
||||||
|
if got := isNumericString(in); got != want {
|
||||||
|
t.Errorf("isNumericString(%q) = %v", in, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildReferencesChain(t *testing.T) {
|
||||||
|
tm := NewTypeMapper()
|
||||||
|
fk := &models.Constraint{ReferencedColumns: []string{"id"}}
|
||||||
|
if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" {
|
||||||
|
t.Errorf("no columns: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortHelpers(t *testing.T) {
|
||||||
|
idxs := map[string]*models.Index{
|
||||||
|
"b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1},
|
||||||
|
}
|
||||||
|
got := sortIndexes(idxs)
|
||||||
|
if len(got) != 4 {
|
||||||
|
t.Fatalf("len %d", len(got))
|
||||||
|
}
|
||||||
|
// Items with a sequence are ordered by it relative to each other.
|
||||||
|
pos := map[string]int{}
|
||||||
|
for i, ix := range got {
|
||||||
|
pos[ix.Name] = i
|
||||||
|
}
|
||||||
|
if pos["y"] > pos["z"] {
|
||||||
|
t.Errorf("sequence order violated: %v", pos)
|
||||||
|
}
|
||||||
|
|
||||||
|
cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}})
|
||||||
|
if len(cons) != 2 || cons[0].Name != "a" {
|
||||||
|
t.Errorf("sortConstraints: %+v", cons)
|
||||||
|
}
|
||||||
|
|
||||||
|
strs := []string{"c", "a", "b"}
|
||||||
|
sortStrings(strs)
|
||||||
|
if strings.Join(strs, "") != "abc" {
|
||||||
|
t.Errorf("sortStrings: %v", strs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnumColumnCallsConstant(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "schema.ts")
|
||||||
|
w := NewWriter(&writers.WriterOptions{OutputPath: out})
|
||||||
|
if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := string(b)
|
||||||
|
if !strings.Contains(got, "role('role')") {
|
||||||
|
t.Errorf("enum column should call constant:\n%s", got)
|
||||||
|
}
|
||||||
|
if strings.Contains(got, "pgEnum('role')(") {
|
||||||
|
t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
package gorm
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestSnakeCaseToCamelCase(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"user", "user"},
|
||||||
|
{"User_Name", "userName"},
|
||||||
|
{"user_id", "userID"},
|
||||||
|
{"http_request", "httpRequest"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := SnakeCaseToCamelCase(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPascalCaseToSnakeCase(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"User", "user"},
|
||||||
|
{"UserName", "user_name"},
|
||||||
|
{"UserID", "user_id"},
|
||||||
|
{"HTTPRequest", "http_request"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := PascalCaseToSnakeCase(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSingularize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"people", "person"},
|
||||||
|
{"People", "person"},
|
||||||
|
{"categories", "category"},
|
||||||
|
{"wolves", "wolf"},
|
||||||
|
{"boxes", "box"},
|
||||||
|
{"churches", "church"},
|
||||||
|
{"users", "user"},
|
||||||
|
{"class", "class"},
|
||||||
|
{"user", "user"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Singularize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPluralize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"person", "people"},
|
||||||
|
{"category", "categories"},
|
||||||
|
{"box", "boxes"},
|
||||||
|
{"church", "churches"},
|
||||||
|
{"user", "users"},
|
||||||
|
{"day", "days"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Pluralize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("%q: got %q want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsVowel(t *testing.T) {
|
||||||
|
for _, c := range []byte("aeiouAEIOU") {
|
||||||
|
if !isVowel(c) {
|
||||||
|
t.Errorf("%c should be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, c := range []byte("bcxyzBZ1_") {
|
||||||
|
if isVowel(c) {
|
||||||
|
t.Errorf("%c should not be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types
|
// TypeMapper handles type conversions between SQL and Go types
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
|
typeMappings map[string]string
|
||||||
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,6 +31,10 @@ func NewTypeMapper(typeStyle string) *TypeMapper {
|
|||||||
|
|
||||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
||||||
|
if goType, ok := tm.overrideGoType(sqlType, notNull); ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
|
||||||
// Array types are handled separately for both styles.
|
// Array types are handled separately for both styles.
|
||||||
if pgsql.IsArrayType(sqlType) {
|
if pgsql.IsArrayType(sqlType) {
|
||||||
return tm.arrayGoType(tm.extractBaseType(sqlType))
|
return tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||||
@@ -57,6 +62,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
|||||||
return tm.nullableGoType(baseType)
|
return tm.nullableGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetTypeMappings installs user-configured SQL-to-Go type overrides.
|
||||||
|
func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) {
|
||||||
|
tm.typeMappings = mappings
|
||||||
|
}
|
||||||
|
|
||||||
|
// overrideGoType applies a configured override, if any, for the column type.
|
||||||
|
func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) {
|
||||||
|
if len(tm.typeMappings) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType))
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if pgsql.IsArrayType(sqlType) {
|
||||||
|
return "[]" + goType, true
|
||||||
|
}
|
||||||
|
return writers.ApplyTypeMapping(goType, notNull), true
|
||||||
|
}
|
||||||
|
|
||||||
// extractBaseType extracts the base type from a SQL type string
|
// extractBaseType extracts the base type from a SQL type string
|
||||||
// Examples: varchar(100) → varchar, numeric(10,2) → numeric
|
// Examples: varchar(100) → varchar, numeric(10,2) → numeric
|
||||||
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package gorm
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestTypeMapper_CustomTypeMappings(t *testing.T) {
|
||||||
|
mapper := NewTypeMapper("")
|
||||||
|
mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"uuid", true, "uuid.UUID"},
|
||||||
|
{"uuid", false, "*uuid.UUID"},
|
||||||
|
{"numeric(10,2)", true, "decimal.Decimal"},
|
||||||
|
{"UUID[]", false, "[]uuid.UUID"},
|
||||||
|
{"bigint", true, "int64"}, // unmapped types keep defaults
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want {
|
||||||
|
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
package gorm
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSQLTypeToGoType_Styles(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
style string
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{writers.NullableTypeSqlTypes, "integer", true, "int32"},
|
||||||
|
{writers.NullableTypeSqlTypes, "bigint", false, "sql_types.SqlInt64"},
|
||||||
|
{writers.NullableTypeSqlTypes, "text", false, "sql_types.SqlString"},
|
||||||
|
{writers.NullableTypeSqlTypes, "text[]", true, "sql_types.SqlStringArray"},
|
||||||
|
{writers.NullableTypeSqlTypes, "integer[]", false, "sql_types.SqlInt32Array"},
|
||||||
|
{writers.NullableTypeSqlTypes, "bigint[]", false, "sql_types.SqlInt64Array"},
|
||||||
|
{writers.NullableTypeSqlTypes, "smallint[]", false, "sql_types.SqlInt16Array"},
|
||||||
|
{writers.NullableTypeSqlTypes, "real[]", false, "sql_types.SqlFloat32Array"},
|
||||||
|
{writers.NullableTypeSqlTypes, "numeric[]", false, "sql_types.SqlFloat64Array"},
|
||||||
|
{writers.NullableTypeSqlTypes, "boolean[]", false, "sql_types.SqlBoolArray"},
|
||||||
|
{writers.NullableTypeSqlTypes, "uuid[]", false, "sql_types.SqlUUIDArray"},
|
||||||
|
{writers.NullableTypeSqlTypes, "weird[]", false, "sql_types.SqlStringArray"},
|
||||||
|
{writers.NullableTypeSqlTypes, "unknowntype", false, "sql_types.SqlString"},
|
||||||
|
{writers.NullableTypeStdlib, "integer", true, "int32"},
|
||||||
|
{writers.NullableTypeStdlib, "integer", false, "sql.NullInt32"},
|
||||||
|
{writers.NullableTypeStdlib, "smallint", false, "sql.NullInt16"},
|
||||||
|
{writers.NullableTypeStdlib, "bigint", false, "sql.NullInt64"},
|
||||||
|
{writers.NullableTypeStdlib, "boolean", false, "sql.NullBool"},
|
||||||
|
{writers.NullableTypeStdlib, "double precision", false, "sql.NullFloat64"},
|
||||||
|
{writers.NullableTypeStdlib, "varchar(10)", false, "sql.NullString"},
|
||||||
|
{writers.NullableTypeStdlib, "timestamptz", false, "sql.NullTime"},
|
||||||
|
{writers.NullableTypeStdlib, "bytea", false, "[]byte"},
|
||||||
|
{writers.NullableTypeStdlib, "mystery", false, "sql.NullString"},
|
||||||
|
{writers.NullableTypeBaselib, "integer", true, "int32"},
|
||||||
|
{writers.NullableTypeBaselib, "integer", false, "*int32"},
|
||||||
|
{writers.NullableTypeBaselib, "text", false, "*string"},
|
||||||
|
{"", "text", false, "*string"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.style+"/"+tt.sqlType, func(t *testing.T) {
|
||||||
|
got := NewTypeMapper(tt.style).SQLTypeToGoType(tt.sqlType, tt.notNull)
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("got %q want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStdlibArrayTypes(t *testing.T) {
|
||||||
|
tm := NewTypeMapper(writers.NullableTypeStdlib)
|
||||||
|
for _, sqlType := range []string{"text[]", "integer[]", "bigint[]", "boolean[]", "uuid[]", "numeric[]"} {
|
||||||
|
if got := tm.SQLTypeToGoType(sqlType, true); got == "" {
|
||||||
|
t.Errorf("%s: empty", sqlType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestImportHelpers(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
style string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{writers.NullableTypeStdlib, `"database/sql"`},
|
||||||
|
{writers.NullableTypeBaselib, ""},
|
||||||
|
{writers.NullableTypeSqlTypes, `sql_types "git.warky.dev/wdevs/relspecgo/pkg/sqltypes"`},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := NewTypeMapper(tt.style).GetNullableTypeImportLine(); got != tt.want {
|
||||||
|
t.Errorf("%s: got %q want %q", tt.style, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tm := NewTypeMapper("")
|
||||||
|
if !tm.NeedsFmtImport(true) || tm.NeedsFmtImport(false) {
|
||||||
|
t.Error("NeedsFmtImport should echo its argument")
|
||||||
|
}
|
||||||
|
if tm.GetSQLTypesImport() == "" {
|
||||||
|
t.Error("empty sqltypes import")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer {
|
|||||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
w.typeMapper.SetTypeMappings(options.TypeMappings)
|
||||||
|
|
||||||
// Initialize templates
|
// Initialize templates
|
||||||
tmpl, err := NewTemplates()
|
tmpl, err := NewTemplates()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+72
-74
@@ -44,26 +44,11 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
|||||||
return w.executeDatabaseSQL(db, connString)
|
return w.executeDatabaseSQL(db, connString)
|
||||||
}
|
}
|
||||||
|
|
||||||
var writer io.Writer
|
release, err := w.openOutput()
|
||||||
var file *os.File
|
|
||||||
var err error
|
|
||||||
|
|
||||||
// Use existing writer if already set (for testing)
|
|
||||||
if w.writer != nil {
|
|
||||||
writer = w.writer
|
|
||||||
} else if w.options.OutputPath != "" {
|
|
||||||
// Determine output destination
|
|
||||||
file, err = os.Create(w.options.OutputPath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create output file: %w", err)
|
return err
|
||||||
}
|
}
|
||||||
defer file.Close()
|
defer release()
|
||||||
writer = file
|
|
||||||
} else {
|
|
||||||
writer = os.Stdout
|
|
||||||
}
|
|
||||||
|
|
||||||
w.writer = writer
|
|
||||||
|
|
||||||
// Write header comment
|
// Write header comment
|
||||||
fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n")
|
fmt.Fprintf(w.writer, "-- MSSQL Database Schema\n")
|
||||||
@@ -80,11 +65,34 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// openOutput points w.writer at the configured destination (output file or
|
||||||
|
// stdout) when none is set, and returns a func that releases it again.
|
||||||
|
func (w *Writer) openOutput() (func(), error) {
|
||||||
|
if w.writer != nil {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
if w.options.OutputPath != "" {
|
||||||
|
file, err := os.Create(w.options.OutputPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create output file: %w", err)
|
||||||
|
}
|
||||||
|
w.writer = file
|
||||||
|
return func() {
|
||||||
|
file.Close()
|
||||||
|
w.writer = nil
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
w.writer = os.Stdout
|
||||||
|
return func() { w.writer = nil }, nil
|
||||||
|
}
|
||||||
|
|
||||||
// WriteSchema writes a single schema and all its tables
|
// WriteSchema writes a single schema and all its tables
|
||||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||||
if w.writer == nil {
|
release, err := w.openOutput()
|
||||||
w.writer = os.Stdout
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
defer release()
|
||||||
|
|
||||||
// Phase 1: Create schema (skip dbo schema and when flattening)
|
// Phase 1: Create schema (skip dbo schema and when flattening)
|
||||||
if schema.Name != "dbo" && !w.options.FlattenSchema {
|
if schema.Name != "dbo" && !w.options.FlattenSchema {
|
||||||
@@ -153,9 +161,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
|||||||
|
|
||||||
// WriteTable writes a single table with all its elements
|
// WriteTable writes a single table with all its elements
|
||||||
func (w *Writer) WriteTable(table *models.Table) error {
|
func (w *Writer) WriteTable(table *models.Table) error {
|
||||||
if w.writer == nil {
|
release, err := w.openOutput()
|
||||||
w.writer = os.Stdout
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
defer release()
|
||||||
|
|
||||||
// Create a temporary schema with just this table
|
// Create a temporary schema with just this table
|
||||||
schema := models.InitSchema(table.Schema)
|
schema := models.InitSchema(table.Schema)
|
||||||
@@ -481,18 +491,12 @@ func (w *Writer) writeComments(schema *models.Schema, table *models.Table) error
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// executeDatabaseSQL executes SQL statements directly on an MSSQL database
|
// executeDatabaseSQL executes the full generated schema (tables, keys,
|
||||||
|
// indexes, constraints and comments) directly on an MSSQL database.
|
||||||
func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error {
|
func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) error {
|
||||||
// Generate SQL statements
|
statements, err := w.generateStatements(db)
|
||||||
statements := []string{}
|
if err != nil {
|
||||||
statements = append(statements, "-- MSSQL Database Schema")
|
return err
|
||||||
statements = append(statements, fmt.Sprintf("-- Database: %s", db.Name))
|
|
||||||
statements = append(statements, "-- Generated by RelSpec")
|
|
||||||
|
|
||||||
for _, schema := range db.Schemas {
|
|
||||||
if err := w.generateSchemaStatements(schema, &statements); err != nil {
|
|
||||||
return fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Connect to database
|
// Connect to database
|
||||||
@@ -510,17 +514,9 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
|||||||
// Execute statements
|
// Execute statements
|
||||||
executedCount := 0
|
executedCount := 0
|
||||||
for i, stmt := range statements {
|
for i, stmt := range statements {
|
||||||
stmtTrimmed := strings.TrimSpace(stmt)
|
|
||||||
|
|
||||||
// Skip comments and empty statements
|
|
||||||
if strings.HasPrefix(stmtTrimmed, "--") || stmtTrimmed == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
fmt.Fprintf(os.Stderr, "Executing statement %d/%d...\n", i+1, len(statements))
|
fmt.Fprintf(os.Stderr, "Executing statement %d/%d...\n", i+1, len(statements))
|
||||||
|
|
||||||
_, execErr := dbConn.ExecContext(ctx, stmt)
|
if _, execErr := dbConn.ExecContext(ctx, stmt); execErr != nil {
|
||||||
if execErr != nil {
|
|
||||||
fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr)
|
fmt.Fprintf(os.Stderr, "⚠ Warning: Statement failed: %v\n", execErr)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -532,49 +528,51 @@ func (w *Writer) executeDatabaseSQL(db *models.Database, connString string) erro
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateSchemaStatements generates SQL statements for a schema
|
// generateStatements renders the same script WriteDatabase would produce and
|
||||||
func (w *Writer) generateSchemaStatements(schema *models.Schema, statements *[]string) error {
|
// splits it into individually executable statements (comments removed).
|
||||||
// Phase 1: Create schema
|
func (w *Writer) generateStatements(db *models.Database) ([]string, error) {
|
||||||
if schema.Name != "dbo" && !w.options.FlattenSchema {
|
var buf strings.Builder
|
||||||
*statements = append(*statements, fmt.Sprintf("-- Schema: %s", schema.Name))
|
saved := w.writer
|
||||||
*statements = append(*statements, fmt.Sprintf("CREATE SCHEMA [%s];", schema.Name))
|
w.writer = &buf
|
||||||
|
defer func() { w.writer = saved }()
|
||||||
|
|
||||||
|
for _, schema := range db.Schemas {
|
||||||
|
if err := w.WriteSchema(schema); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to generate statements for schema %s: %w", schema.Name, err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Phase 2: Create tables
|
// Every statement the writer emits ends with ";\n\n".
|
||||||
*statements = append(*statements, fmt.Sprintf("-- Tables for schema: %s", schema.Name))
|
var statements []string
|
||||||
for _, table := range schema.Tables {
|
for _, chunk := range strings.Split(buf.String(), ";\n\n") {
|
||||||
createTableSQL := fmt.Sprintf("CREATE TABLE %s (", w.qualTable(schema.Name, table.Name))
|
lines := make([]string, 0)
|
||||||
columnDefs := make([]string, 0)
|
for _, line := range strings.Split(chunk, "\n") {
|
||||||
|
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
|
||||||
columns := getSortedColumns(table.Columns)
|
lines = append(lines, line)
|
||||||
for _, col := range columns {
|
|
||||||
def := w.generateColumnDefinition(col)
|
|
||||||
columnDefs = append(columnDefs, " "+def)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
createTableSQL += "\n" + strings.Join(columnDefs, ",\n") + "\n)"
|
|
||||||
*statements = append(*statements, createTableSQL)
|
|
||||||
}
|
}
|
||||||
|
if stmt := strings.TrimSpace(strings.Join(lines, "\n")); stmt != "" {
|
||||||
// Phase 3-7: Constraints and indexes will be added by WriteSchema logic
|
statements = append(statements, stmt)
|
||||||
// For now, just create tables
|
}
|
||||||
return nil
|
}
|
||||||
|
return statements, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Helper functions
|
// Helper functions
|
||||||
|
|
||||||
// getSortedColumns returns columns sorted by sequence
|
// getSortedColumns returns columns sorted by sequence, then by name so that
|
||||||
|
// columns without a sequence still come out in a stable order.
|
||||||
func getSortedColumns(columns map[string]*models.Column) []*models.Column {
|
func getSortedColumns(columns map[string]*models.Column) []*models.Column {
|
||||||
names := make([]string, 0, len(columns))
|
|
||||||
for name := range columns {
|
|
||||||
names = append(names, name)
|
|
||||||
}
|
|
||||||
sort.Strings(names)
|
|
||||||
|
|
||||||
sorted := make([]*models.Column, 0, len(columns))
|
sorted := make([]*models.Column, 0, len(columns))
|
||||||
for _, name := range names {
|
for _, col := range columns {
|
||||||
sorted = append(sorted, columns[name])
|
sorted = append(sorted, col)
|
||||||
}
|
}
|
||||||
|
sort.Slice(sorted, func(i, j int) bool {
|
||||||
|
if sorted[i].Sequence != sorted[j].Sequence {
|
||||||
|
return sorted[i].Sequence < sorted[j].Sequence
|
||||||
|
}
|
||||||
|
return sorted[i].Name < sorted[j].Name
|
||||||
|
})
|
||||||
return sorted
|
return sorted
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,205 @@
|
|||||||
|
package mssql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func shopDB() *models.Database {
|
||||||
|
s := models.InitSchema("sales")
|
||||||
|
users := models.InitTable("users", "sales")
|
||||||
|
users.Description = "Registered users"
|
||||||
|
id := models.InitColumn("id", "users", "sales")
|
||||||
|
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "int", true, true, true, 1
|
||||||
|
email := models.InitColumn("email", "users", "sales")
|
||||||
|
email.Type, email.Length, email.NotNull, email.Sequence, email.Description = "string", 255, true, 2, "Login e-mail"
|
||||||
|
age := models.InitColumn("age", "users", "sales")
|
||||||
|
age.Type, age.Sequence, age.Default = "int", 3, 18
|
||||||
|
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
|
||||||
|
|
||||||
|
pk := models.InitConstraint("PK_users", models.PrimaryKeyConstraint)
|
||||||
|
pk.Columns = []string{"id"}
|
||||||
|
uq := models.InitConstraint("UQ_users_email", models.UniqueConstraint)
|
||||||
|
uq.Columns = []string{"email"}
|
||||||
|
ck := models.InitConstraint("CK_users_age", models.CheckConstraint)
|
||||||
|
ck.Expression = "[age] >= 0"
|
||||||
|
emptyCk := models.InitConstraint("CK_empty", models.CheckConstraint)
|
||||||
|
users.Constraints["PK_users"], users.Constraints["UQ_users_email"], users.Constraints["CK_users_age"], users.Constraints["CK_empty"] = pk, uq, ck, emptyCk
|
||||||
|
ix := models.InitIndex("IX_users_age", "users", "sales")
|
||||||
|
ix.Columns, ix.Unique = []string{"age"}, true
|
||||||
|
pkIx := models.InitIndex("pk_users_idx", "users", "sales")
|
||||||
|
pkIx.Columns = []string{"id"}
|
||||||
|
noCols := models.InitIndex("IX_nocols", "users", "sales")
|
||||||
|
users.Indexes["IX_users_age"], users.Indexes["pk_users_idx"], users.Indexes["IX_nocols"] = ix, pkIx, noCols
|
||||||
|
|
||||||
|
orders := models.InitTable("orders", "sales")
|
||||||
|
oid := models.InitColumn("id", "orders", "sales")
|
||||||
|
oid.Type, oid.IsPrimaryKey, oid.NotNull = "int", true, true
|
||||||
|
uid := models.InitColumn("user_id", "orders", "sales")
|
||||||
|
uid.Type, uid.NotNull = "int", true
|
||||||
|
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
|
||||||
|
fk := models.InitConstraint("FK_orders_users", models.ForeignKeyConstraint)
|
||||||
|
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
|
||||||
|
fk.OnDelete = "cascade"
|
||||||
|
badFk := models.InitConstraint("FK_bad", models.ForeignKeyConstraint)
|
||||||
|
orders.Constraints["FK_orders_users"], orders.Constraints["FK_bad"] = fk, badFk
|
||||||
|
|
||||||
|
s.Tables = append(s.Tables, users, orders)
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeToFile(t *testing.T, opts *writers.WriterOptions, db *models.Database) string {
|
||||||
|
t.Helper()
|
||||||
|
out := filepath.Join(t.TempDir(), "out.sql")
|
||||||
|
opts.OutputPath = out
|
||||||
|
if err := NewWriter(opts).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_FullScript(t *testing.T) {
|
||||||
|
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
|
||||||
|
for _, want := range []string{
|
||||||
|
"-- Database: shop", "CREATE SCHEMA [sales];",
|
||||||
|
"CREATE TABLE [sales].[users]", "[email] NVARCHAR(255) NOT NULL", "DEFAULT 18",
|
||||||
|
"ALTER TABLE [sales].[users] ADD CONSTRAINT [PK_users] PRIMARY KEY ([id]);",
|
||||||
|
"ALTER TABLE [sales].[orders] ADD CONSTRAINT [PK_sales_orders] PRIMARY KEY ([id]);", // generated PK name from IsPrimaryKey
|
||||||
|
"CREATE UNIQUE INDEX [IX_users_age] ON [sales].[users] ([age]);",
|
||||||
|
"ADD CONSTRAINT [UQ_users_email] UNIQUE ([email]);",
|
||||||
|
"ADD CONSTRAINT [CK_users_age] CHECK ([age] >= 0);",
|
||||||
|
"ADD CONSTRAINT [FK_orders_users] FOREIGN KEY ([user_id])",
|
||||||
|
"REFERENCES [sales].[users] ([id])", "ON DELETE CASCADE ON UPDATE NO ACTION;",
|
||||||
|
"@value = 'Registered users'", "@level2type = 'COLUMN', @level2name = 'email';",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("missing %q\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, unwanted := range []string{"pk_users_idx", "IX_nocols", "CK_empty", "FK_bad"} {
|
||||||
|
if strings.Contains(got, unwanted) {
|
||||||
|
t.Errorf("%q must be skipped\n%s", unwanted, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_PhaseOrder(t *testing.T) {
|
||||||
|
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
|
||||||
|
last := -1
|
||||||
|
for _, marker := range []string{"-- Schema: sales", "-- Tables for", "-- Primary keys", "-- Indexes", "-- Unique constraints", "-- Check constraints", "-- Foreign keys", "-- Comments"} {
|
||||||
|
i := strings.Index(got, marker)
|
||||||
|
if i < 0 || i < last {
|
||||||
|
t.Fatalf("marker %q out of order (index %d after %d)", marker, i, last)
|
||||||
|
}
|
||||||
|
last = i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||||
|
first := writeToFile(t, &writers.WriterOptions{}, shopDB())
|
||||||
|
for i := 0; i < 15; i++ {
|
||||||
|
if got := writeToFile(t, &writers.WriterOptions{}, shopDB()); got != first {
|
||||||
|
t.Fatalf("output differs on run %d", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_FlattenAndDbo(t *testing.T) {
|
||||||
|
flat := writeToFile(t, &writers.WriterOptions{FlattenSchema: true}, shopDB())
|
||||||
|
if strings.Contains(flat, "CREATE SCHEMA") || !strings.Contains(flat, "CREATE TABLE [users]") || strings.Contains(flat, "[sales].") {
|
||||||
|
t.Errorf("flatten:\n%s", flat)
|
||||||
|
}
|
||||||
|
|
||||||
|
db := shopDB()
|
||||||
|
db.Schemas[0].Name = "dbo"
|
||||||
|
dbo := writeToFile(t, &writers.WriterOptions{}, db)
|
||||||
|
if strings.Contains(dbo, "CREATE SCHEMA") {
|
||||||
|
t.Errorf("dbo schema must not be created:\n%s", dbo)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteTableAndSchema(t *testing.T) {
|
||||||
|
db := shopDB()
|
||||||
|
out := filepath.Join(t.TempDir(), "t.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, _ := os.ReadFile(out)
|
||||||
|
if !strings.Contains(string(b), "CREATE TABLE [sales].[users]") || strings.Contains(string(b), "CREATE TABLE [sales].[orders]") {
|
||||||
|
t.Errorf("WriteTable output:\n%s", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_OutputErrors(t *testing.T) {
|
||||||
|
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]any{
|
||||||
|
"connection_string": "sqlserver://u:p@127.0.0.1:1?database=none&connection+timeout=1",
|
||||||
|
}}
|
||||||
|
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "ping database") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGenerateStatements_CoversFullSchema(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
stmts, err := w.generateStatements(shopDB())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
joined := strings.Join(stmts, "\n---\n")
|
||||||
|
for _, want := range []string{
|
||||||
|
"CREATE SCHEMA [sales]", "CREATE TABLE [sales].[users]",
|
||||||
|
"PRIMARY KEY ([id])", "CREATE UNIQUE INDEX [IX_users_age]", "UNIQUE ([email])",
|
||||||
|
"CHECK ([age] >= 0)", "FOREIGN KEY ([user_id])", "EXEC sp_addextendedproperty",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(joined, want) {
|
||||||
|
t.Errorf("missing %q in:\n%s", want, joined)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, stmt := range stmts {
|
||||||
|
if strings.HasPrefix(stmt, "--") || strings.HasSuffix(stmt, ";") || stmt == "" {
|
||||||
|
t.Errorf("statement not clean: %q", stmt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if w.writer != nil {
|
||||||
|
t.Error("generateStatements must restore the writer")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDBMLFixtureProducesScript(t *testing.T) {
|
||||||
|
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := writeToFile(t, &writers.WriterOptions{}, db)
|
||||||
|
if !strings.Contains(got, "CREATE TABLE") || !strings.Contains(got, "-- Foreign keys") {
|
||||||
|
t.Errorf("script:\n%s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColumnsOrderedBySequence(t *testing.T) {
|
||||||
|
got := writeToFile(t, &writers.WriterOptions{}, shopDB())
|
||||||
|
id, email, age := strings.Index(got, "[id] INT"), strings.Index(got, "[email] NVARCHAR"), strings.Index(got, "[age] INT")
|
||||||
|
if !(id < email && email < age) {
|
||||||
|
t.Errorf("columns must follow Sequence (id, email, age): %d %d %d\n%s", id, email, age, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,209 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
_ "github.com/go-sql-driver/mysql"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Writer struct {
|
||||||
|
options *writers.WriterOptions
|
||||||
|
writer io.Writer
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewWriter(options *writers.WriterOptions) *Writer { return &Writer{options: options} }
|
||||||
|
func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||||
|
if w.options == nil {
|
||||||
|
return fmt.Errorf("writer options are required")
|
||||||
|
}
|
||||||
|
if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" {
|
||||||
|
return w.execute(db, conn)
|
||||||
|
}
|
||||||
|
release, err := w.openOutput()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer release()
|
||||||
|
return w.writeDatabaseDDL(db)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) writeDatabaseDDL(db *models.Database) error {
|
||||||
|
if _, err := fmt.Fprintf(w.writer, "-- MySQL Database Schema\n-- Database: %s\n-- Generated by RelSpec\n\n", db.Name); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, s := range db.Schemas {
|
||||||
|
if err := w.WriteSchema(s); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// openOutput points w.writer at the configured destination (output file or
|
||||||
|
// stdout) when none is set, and returns a func that releases it again.
|
||||||
|
func (w *Writer) openOutput() (func(), error) {
|
||||||
|
if w.writer != nil {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
if w.options != nil && w.options.OutputPath != "" {
|
||||||
|
f, err := os.Create(w.options.OutputPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
w.writer = f
|
||||||
|
return func() {
|
||||||
|
f.Close()
|
||||||
|
w.writer = nil
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
w.writer = os.Stdout
|
||||||
|
return func() { w.writer = nil }, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) WriteSchema(s *models.Schema) error {
|
||||||
|
release, err := w.openOutput()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer release()
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
if err := w.writeTable(s, t); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) WriteTable(t *models.Table) error {
|
||||||
|
release, err := w.openOutput()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer release()
|
||||||
|
return w.writeTable(nil, t)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
||||||
|
name := t.Name
|
||||||
|
if s != nil {
|
||||||
|
name = fmt.Sprintf("%s.%s", quote(s.Name), quote(t.Name))
|
||||||
|
} else {
|
||||||
|
name = quote(name)
|
||||||
|
}
|
||||||
|
cols := make([]*models.Column, 0, len(t.Columns))
|
||||||
|
for _, c := range t.Columns {
|
||||||
|
cols = append(cols, c)
|
||||||
|
}
|
||||||
|
sort.Slice(cols, func(i, j int) bool {
|
||||||
|
if cols[i].Sequence != cols[j].Sequence {
|
||||||
|
return cols[i].Sequence < cols[j].Sequence
|
||||||
|
}
|
||||||
|
return cols[i].Name < cols[j].Name
|
||||||
|
})
|
||||||
|
defs := []string{}
|
||||||
|
pk := []string{}
|
||||||
|
for _, c := range cols {
|
||||||
|
d := fmt.Sprintf(" %s %s", quote(c.Name), mariadb.ConvertCanonicalToMariaDB(c.Type))
|
||||||
|
if c.Length > 0 && strings.EqualFold(c.Type, "string") {
|
||||||
|
d = fmt.Sprintf(" %s VARCHAR(%d)", quote(c.Name), c.Length)
|
||||||
|
}
|
||||||
|
if c.NotNull {
|
||||||
|
d += " NOT NULL"
|
||||||
|
}
|
||||||
|
if c.AutoIncrement {
|
||||||
|
d += " AUTO_INCREMENT"
|
||||||
|
}
|
||||||
|
if c.Default != nil {
|
||||||
|
d += fmt.Sprintf(" DEFAULT %s", writers.QuoteDefaultValue(fmt.Sprint(c.Default), c.Type))
|
||||||
|
}
|
||||||
|
defs = append(defs, d)
|
||||||
|
if c.IsPrimaryKey {
|
||||||
|
pk = append(pk, quote(c.Name))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
constraintNames := make([]string, 0, len(t.Constraints))
|
||||||
|
for name := range t.Constraints {
|
||||||
|
constraintNames = append(constraintNames, name)
|
||||||
|
}
|
||||||
|
sort.Strings(constraintNames)
|
||||||
|
for _, name := range constraintNames {
|
||||||
|
c := t.Constraints[name]
|
||||||
|
if c.Type == models.PrimaryKeyConstraint {
|
||||||
|
pk = nil
|
||||||
|
for _, n := range c.Columns {
|
||||||
|
pk = append(pk, quote(n))
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(pk) > 0 {
|
||||||
|
defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")")
|
||||||
|
}
|
||||||
|
for _, name := range constraintNames {
|
||||||
|
c := t.Constraints[name]
|
||||||
|
if c.Type == models.UniqueConstraint {
|
||||||
|
defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sql := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n) ENGINE=InnoDB;\n\n", name, strings.Join(defs, ",\n"))
|
||||||
|
if _, err := io.WriteString(w.writer, sql); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) execute(dbm *models.Database, conn string) error {
|
||||||
|
var b strings.Builder
|
||||||
|
old := w.writer
|
||||||
|
w.writer = &b
|
||||||
|
if err := w.writeDatabaseDDL(dbm); err != nil {
|
||||||
|
w.writer = old
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
w.writer = old
|
||||||
|
db, err := sql.Open("mysql", conn)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to connect: %w", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
for _, stmt := range strings.Split(b.String(), ";\n") {
|
||||||
|
stmt = stripComments(strings.TrimSpace(stmt))
|
||||||
|
if stmt == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := db.ExecContext(context.Background(), stmt); err != nil {
|
||||||
|
return fmt.Errorf("failed to execute SQL: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func stripComments(sqlText string) string {
|
||||||
|
lines := strings.Split(sqlText, "\n")
|
||||||
|
kept := lines[:0]
|
||||||
|
for _, line := range lines {
|
||||||
|
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
|
||||||
|
kept = append(kept, line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(strings.Join(kept, "\n"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func quote(s string) string { return "`" + strings.ReplaceAll(s, "`", "``") + "`" }
|
||||||
|
func quoted(xs []string) string {
|
||||||
|
out := make([]string, len(xs))
|
||||||
|
for i, x := range xs {
|
||||||
|
out[i] = quote(x)
|
||||||
|
}
|
||||||
|
return strings.Join(out, ", ")
|
||||||
|
}
|
||||||
@@ -0,0 +1,159 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func shopDB() *models.Database {
|
||||||
|
s := models.InitSchema("shop")
|
||||||
|
t := models.InitTable("users", "shop")
|
||||||
|
add := func(name, typ string, mod func(*models.Column)) {
|
||||||
|
c := models.InitColumn(name, "users", "shop")
|
||||||
|
c.Type = typ
|
||||||
|
if mod != nil {
|
||||||
|
mod(c)
|
||||||
|
}
|
||||||
|
t.Columns[name] = c
|
||||||
|
}
|
||||||
|
add("id", "int", func(c *models.Column) { c.IsPrimaryKey, c.NotNull, c.AutoIncrement = true, true, true })
|
||||||
|
add("email", "string", func(c *models.Column) { c.Length, c.NotNull = 255, true })
|
||||||
|
add("nick", "string", nil)
|
||||||
|
add("age", "int", func(c *models.Column) { c.Default = 18 })
|
||||||
|
add("active", "boolean", func(c *models.Column) { c.Default = true })
|
||||||
|
add("zeta", "string", nil)
|
||||||
|
add("alpha", "string", nil)
|
||||||
|
for _, name := range []string{"uq_b", "uq_a", "uq_c"} {
|
||||||
|
u := models.InitConstraint(name, models.UniqueConstraint)
|
||||||
|
u.Columns = []string{"email"}
|
||||||
|
t.Constraints[name] = u
|
||||||
|
}
|
||||||
|
s.Tables = append(s.Tables, t)
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeFile(t *testing.T, db *models.Database) string {
|
||||||
|
t.Helper()
|
||||||
|
out := filepath.Join(t.TempDir(), "out.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_ToFile(t *testing.T) {
|
||||||
|
got := writeFile(t, shopDB())
|
||||||
|
for _, want := range []string{
|
||||||
|
"-- Database: shop", "CREATE TABLE IF NOT EXISTS `shop`.`users`",
|
||||||
|
"`id` ", "AUTO_INCREMENT", "`email` VARCHAR(255) NOT NULL", "DEFAULT 18",
|
||||||
|
"PRIMARY KEY (`id`)", "CONSTRAINT `uq_a` UNIQUE (`email`)", "ENGINE=InnoDB",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("missing %q\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||||
|
first := writeFile(t, shopDB())
|
||||||
|
for i := 0; i < 30; i++ {
|
||||||
|
if got := writeFile(t, shopDB()); got != first {
|
||||||
|
t.Fatalf("output differs on run %d:\n--- first\n%s\n--- got\n%s", i, first, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_UniqueConstraintsSorted(t *testing.T) {
|
||||||
|
got := writeFile(t, shopDB())
|
||||||
|
a, b, c := strings.Index(got, "`uq_a`"), strings.Index(got, "`uq_b`"), strings.Index(got, "`uq_c`")
|
||||||
|
if !(a < b && b < c) {
|
||||||
|
t.Errorf("unique constraints must be sorted by name: %d %d %d", a, b, c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
|
||||||
|
db := shopDB()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
sOut := filepath.Join(dir, "s.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `shop`.`users`") {
|
||||||
|
t.Errorf("schema output:\n%s", b)
|
||||||
|
}
|
||||||
|
|
||||||
|
tOut := filepath.Join(dir, "t.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE IF NOT EXISTS `users`") {
|
||||||
|
t.Errorf("table output (unqualified name):\n%s", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteSchema_WithoutWriterDoesNotPanic(t *testing.T) {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
t.Fatalf("panic: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
s := shopDB().Schemas[0]
|
||||||
|
s.Tables = nil // nothing to print to stdout
|
||||||
|
if err := NewWriter(&writers.WriterOptions{}).WriteSchema(s); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Errors(t *testing.T) {
|
||||||
|
if err := NewWriter(nil).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "options are required") {
|
||||||
|
t.Errorf("nil options: %v", err)
|
||||||
|
}
|
||||||
|
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil {
|
||||||
|
t.Error("bad output path must fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_ConnectionFailure(t *testing.T) {
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]any{
|
||||||
|
"connection_string": "u:p@tcp(127.0.0.1:1)/none?timeout=1s",
|
||||||
|
}}
|
||||||
|
if err := NewWriter(opts).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to execute SQL") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQuoteHelpers(t *testing.T) {
|
||||||
|
if got := quote("a`b"); got != "`a``b`" {
|
||||||
|
t.Errorf("quote: %q", got)
|
||||||
|
}
|
||||||
|
if got := quoted([]string{"a", "b"}); got != "`a`, `b`" {
|
||||||
|
t.Errorf("quoted: %q", got)
|
||||||
|
}
|
||||||
|
if got := quoted(nil); got != "" {
|
||||||
|
t.Errorf("quoted(nil): %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrimaryKeyConstraintOverridesColumnFlags(t *testing.T) {
|
||||||
|
db := shopDB()
|
||||||
|
tbl := db.Schemas[0].Tables[0]
|
||||||
|
pk := models.InitConstraint("pk", models.PrimaryKeyConstraint)
|
||||||
|
pk.Columns = []string{"email", "id"}
|
||||||
|
tbl.Constraints["pk"] = pk
|
||||||
|
if got := writeFile(t, db); !strings.Contains(got, "PRIMARY KEY (`email`, `id`)") {
|
||||||
|
t.Errorf("composite pk:\n%s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestStripComments(t *testing.T) {
|
||||||
|
got := stripComments("-- header\nCREATE TABLE `users` (\n `id` INT\n);")
|
||||||
|
if strings.HasPrefix(got, "--") || !strings.HasPrefix(got, "CREATE TABLE") {
|
||||||
|
t.Fatalf("stripComments() = %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterGeneratesMySQLDDL(t *testing.T) {
|
||||||
|
db := models.InitDatabase("app")
|
||||||
|
s := models.InitSchema("app")
|
||||||
|
table := models.InitTable("users", "app")
|
||||||
|
table.Columns["id"] = models.InitColumn("id", "users", "app")
|
||||||
|
table.Columns["id"].Type = "int"
|
||||||
|
table.Columns["id"].IsPrimaryKey = true
|
||||||
|
table.Columns["id"].NotNull = true
|
||||||
|
table.Columns["name"] = models.InitColumn("name", "users", "app")
|
||||||
|
table.Columns["name"].Type = "string"
|
||||||
|
table.Columns["name"].Length = 80
|
||||||
|
s.Tables = append(s.Tables, table)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
var out bytes.Buffer
|
||||||
|
w := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{}})
|
||||||
|
w.writer = &out
|
||||||
|
if err := w.WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := out.String()
|
||||||
|
for _, want := range []string{"CREATE TABLE IF NOT EXISTS `app`.`users`", "`id` INT NOT NULL", "`name` VARCHAR(80)", "PRIMARY KEY (`id`)"} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("DDL missing %q:\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,110 @@
|
|||||||
|
package pgsql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCurrentColumnHasDescription(t *testing.T) {
|
||||||
|
table := models.InitTable("users", "public")
|
||||||
|
c := models.InitColumn("Email", "users", "public")
|
||||||
|
c.Description = " the email "
|
||||||
|
table.Columns["Email"] = c
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
table *models.Table
|
||||||
|
col *models.Column
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"nil table", nil, &models.Column{Name: "email", Description: "x"}, false},
|
||||||
|
{"match ignoring case and whitespace", table, &models.Column{Name: "email", Description: "the email"}, true},
|
||||||
|
{"different description", table, &models.Column{Name: "email", Description: "other"}, false},
|
||||||
|
{"column missing", table, &models.Column{Name: "age", Description: "x"}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := currentColumnHasDescription(tt.table, tt.col); got != tt.want {
|
||||||
|
t.Errorf("got %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteCommentColumn(t *testing.T) {
|
||||||
|
te, err := NewTemplateExecutor(false)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got, err := te.ExecuteCommentColumn(CommentColumnData{
|
||||||
|
SchemaName: "public", TableName: "users", ColumnName: "email", Comment: "it''s",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "COMMENT ON COLUMN") || !strings.Contains(got, "public.users") ||
|
||||||
|
!strings.Contains(got, "email") || !strings.Contains(got, "IS 'it''s';") {
|
||||||
|
t.Errorf("unexpected output: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func migrationWithColumnDescription(t *testing.T, currentDesc string, withCurrentCol bool) string {
|
||||||
|
t.Helper()
|
||||||
|
newDB := func(desc string, include bool) *models.Database {
|
||||||
|
db := models.InitDatabase("testdb")
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
tbl := models.InitTable("users", "public")
|
||||||
|
id := models.InitColumn("id", "users", "public")
|
||||||
|
id.Type = "integer"
|
||||||
|
tbl.Columns["id"] = id
|
||||||
|
if include {
|
||||||
|
col := models.InitColumn("email", "users", "public")
|
||||||
|
col.Type = "text"
|
||||||
|
col.Description = desc
|
||||||
|
tbl.Columns["email"] = col
|
||||||
|
}
|
||||||
|
s.Tables = append(s.Tables, tbl)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
model := newDB("it's the email", true)
|
||||||
|
current := newDB(currentDesc, withCurrentCol)
|
||||||
|
|
||||||
|
var buf bytes.Buffer
|
||||||
|
w, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
w.writer = &buf
|
||||||
|
if err := w.WriteMigration(model, current); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteMigration_ColumnComments(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
currentDesc string
|
||||||
|
withCol bool
|
||||||
|
wantComment bool
|
||||||
|
}{
|
||||||
|
{"added", "", true, true},
|
||||||
|
{"changed", "old text", true, true},
|
||||||
|
{"unchanged", "it's the email", true, false},
|
||||||
|
{"new column", "", false, true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
out := migrationWithColumnDescription(t, tt.currentDesc, tt.withCol)
|
||||||
|
has := strings.Contains(out, "COMMENT ON COLUMN") && strings.Contains(out, "it''s the email")
|
||||||
|
if has != tt.wantComment {
|
||||||
|
t.Errorf("comment emitted = %v, want %v\n%s", has, tt.wantComment, out)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,243 @@
|
|||||||
|
package pgsql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func liveWriterConn(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||||
|
if conn == "" {
|
||||||
|
t.Skip("RELSPEC_TEST_PG_CONN not set")
|
||||||
|
}
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// liveWriterSchema returns a unique schema name and drops it on cleanup.
|
||||||
|
func liveWriterSchema(t *testing.T, connString string) (string, *pgx.Conn) {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
conn, err := pgx.Connect(ctx, connString)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect: %v", err)
|
||||||
|
}
|
||||||
|
name := fmt.Sprintf("pgw_test_%d", time.Now().UnixNano())
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
|
||||||
|
_ = conn.Close(ctx)
|
||||||
|
})
|
||||||
|
return name, conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func liveModel(schemaName string, columns map[string]string) *models.Database {
|
||||||
|
db := models.InitDatabase("live")
|
||||||
|
s := models.InitSchema(schemaName)
|
||||||
|
tbl := models.InitTable("accounts", schemaName)
|
||||||
|
id := models.InitColumn("id", "accounts", schemaName)
|
||||||
|
id.Type = "integer"
|
||||||
|
id.NotNull = true
|
||||||
|
id.IsPrimaryKey = true
|
||||||
|
tbl.Columns["id"] = id
|
||||||
|
for name, typ := range columns {
|
||||||
|
c := models.InitColumn(name, "accounts", schemaName)
|
||||||
|
c.Type = typ
|
||||||
|
tbl.Columns[name] = c
|
||||||
|
}
|
||||||
|
s.Tables = append(s.Tables, tbl)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func runLiveWrite(t *testing.T, connString string, db *models.Database, meta map[string]interface{}) (*ExecutionReport, error) {
|
||||||
|
t.Helper()
|
||||||
|
m := map[string]interface{}{"connection_string": connString}
|
||||||
|
for k, v := range meta {
|
||||||
|
m[k] = v
|
||||||
|
}
|
||||||
|
w := NewWriter(&writers.WriterOptions{Metadata: m})
|
||||||
|
err := w.WriteDatabase(db)
|
||||||
|
return w.executionReport, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func columnExists(t *testing.T, conn *pgx.Conn, schema, table, column string) bool {
|
||||||
|
t.Helper()
|
||||||
|
var ok bool
|
||||||
|
err := conn.QueryRow(context.Background(),
|
||||||
|
`SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema=$1 AND table_name=$2 AND column_name=$3)`,
|
||||||
|
schema, table, column).Scan(&ok)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_WriteDatabaseEmptyThenIdenticalThenDrifted(t *testing.T) {
|
||||||
|
connString := liveWriterConn(t)
|
||||||
|
schema, conn := liveWriterSchema(t, connString)
|
||||||
|
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||||
|
meta := map[string]interface{}{"report_path": reportPath}
|
||||||
|
|
||||||
|
// Empty database: schema and table are created.
|
||||||
|
rep, err := runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), meta)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if rep.FailedStatements != 0 || rep.ExecutedStatements == 0 {
|
||||||
|
t.Fatalf("first run report: %+v", rep)
|
||||||
|
}
|
||||||
|
if !columnExists(t, conn, schema, "accounts", "name") {
|
||||||
|
t.Fatal("column name not created")
|
||||||
|
}
|
||||||
|
data, err := os.ReadFile(reportPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("report not written: %v", err)
|
||||||
|
}
|
||||||
|
var onDisk ExecutionReport
|
||||||
|
if err := json.Unmarshal(data, &onDisk); err != nil || onDisk.TotalStatements != rep.TotalStatements {
|
||||||
|
t.Errorf("report on disk mismatch: %v %+v", err, onDisk)
|
||||||
|
}
|
||||||
|
created := false
|
||||||
|
for _, s := range rep.Schemas {
|
||||||
|
for _, tb := range s.Tables {
|
||||||
|
if tb.Name == "accounts" && tb.Created {
|
||||||
|
created = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !created {
|
||||||
|
t.Errorf("table creation not tracked: %+v", rep.Schemas)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Identical database: nothing to execute.
|
||||||
|
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if rep.TotalStatements != 0 {
|
||||||
|
t.Errorf("identical DB must produce no statements, got %d", rep.TotalStatements)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Drifted database: only the new column is added.
|
||||||
|
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text", "email": "text"}), nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if rep.FailedStatements != 0 || rep.TotalStatements == 0 {
|
||||||
|
t.Errorf("drift report: %+v", rep)
|
||||||
|
}
|
||||||
|
if !columnExists(t, conn, schema, "accounts", "email") {
|
||||||
|
t.Error("drifted column email not added")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_WriteDatabaseFailedStatementContinues(t *testing.T) {
|
||||||
|
connString := liveWriterConn(t)
|
||||||
|
schema, _ := liveWriterSchema(t, connString)
|
||||||
|
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||||
|
|
||||||
|
db := liveModel(schema, map[string]string{"bad": "no_such_type_xyz"})
|
||||||
|
rep, err := runLiveWrite(t, connString, db, map[string]interface{}{"full_ddl": true, "report_path": reportPath})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed statements must not abort the run: %v", err)
|
||||||
|
}
|
||||||
|
if rep.FailedStatements == 0 || len(rep.Errors) != rep.FailedStatements {
|
||||||
|
t.Fatalf("expected recorded failures: %+v", rep)
|
||||||
|
}
|
||||||
|
e := rep.Errors[0]
|
||||||
|
if e.StatementNumber == 0 || e.Statement == "" || e.Error == "" {
|
||||||
|
t.Errorf("incomplete error entry: %+v", e)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(reportPath); err != nil {
|
||||||
|
t.Errorf("report must be written even on failures: %v", err)
|
||||||
|
}
|
||||||
|
failedTable := false
|
||||||
|
for _, s := range rep.Schemas {
|
||||||
|
for _, tb := range s.Tables {
|
||||||
|
if tb.Name == "accounts" && !tb.Created && tb.Error != "" {
|
||||||
|
failedTable = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !failedTable {
|
||||||
|
t.Errorf("failed table creation not tracked: %+v", rep.Schemas)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_WriteDatabaseFlattenFallsBackToFullDDL(t *testing.T) {
|
||||||
|
connString := liveWriterConn(t)
|
||||||
|
schema, conn := liveWriterSchema(t, connString)
|
||||||
|
|
||||||
|
w := NewWriter(&writers.WriterOptions{
|
||||||
|
FlattenSchema: true,
|
||||||
|
Metadata: map[string]interface{}{"connection_string": connString},
|
||||||
|
})
|
||||||
|
if err := w.WriteDatabase(liveModel(schema, nil)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// Flattened output lands in public as <schema>_<table>.
|
||||||
|
flat := "public." + schema + "_accounts"
|
||||||
|
t.Cleanup(func() { _, _ = conn.Exec(context.Background(), "DROP TABLE IF EXISTS "+flat+" CASCADE") })
|
||||||
|
var ok bool
|
||||||
|
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", flat).Scan(&ok); err != nil || !ok {
|
||||||
|
t.Errorf("flattened table %s not created (ok=%v err=%v)", flat, ok, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGenerateLiveDiffStatements_FlattenRejected(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
|
||||||
|
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
|
||||||
|
t.Error("flatten must be rejected before connecting")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteStatements_ConnectFailure(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
w.executionReport = &ExecutionReport{}
|
||||||
|
err := w.executeStatements([]string{"SELECT 1"}, "postgres://nobody:x@127.0.0.1:1/none?connect_timeout=1")
|
||||||
|
if err == nil {
|
||||||
|
t.Error("expected connect failure")
|
||||||
|
}
|
||||||
|
if w.executionReport.TotalStatements != 1 {
|
||||||
|
t.Errorf("total not recorded: %+v", w.executionReport)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_ExecuteStatementsSkipsCommentsAndBlank(t *testing.T) {
|
||||||
|
connString := liveWriterConn(t)
|
||||||
|
schema, conn := liveWriterSchema(t, connString)
|
||||||
|
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
w.executionReport = &ExecutionReport{}
|
||||||
|
stmts := []string{
|
||||||
|
"-- Schema: " + schema,
|
||||||
|
" ",
|
||||||
|
"CREATE SCHEMA " + schema,
|
||||||
|
"CREATE TABLE " + schema + ".t (id int)",
|
||||||
|
"-- plain comment",
|
||||||
|
}
|
||||||
|
if err := w.executeStatements(stmts, connString); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
r := w.executionReport
|
||||||
|
if r.ExecutedStatements != 2 || r.FailedStatements != 0 || r.TotalStatements != 5 {
|
||||||
|
t.Errorf("counts: %+v", r)
|
||||||
|
}
|
||||||
|
if len(r.Schemas) != 1 || r.Schemas[0].Name != schema || len(r.Schemas[0].Tables) != 1 || !r.Schemas[0].Tables[0].Created {
|
||||||
|
t.Errorf("schema tracking: %+v", r.Schemas)
|
||||||
|
}
|
||||||
|
var ok bool
|
||||||
|
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", schema+".t").Scan(&ok); err != nil || !ok {
|
||||||
|
t.Errorf("table not created: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,288 @@
|
|||||||
|
package pgsql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExtractTableNameFromCreate(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, in, want string
|
||||||
|
}{
|
||||||
|
{"not create table", "SELECT 1", ""},
|
||||||
|
{"plain", "CREATE TABLE users (id int)", "users"},
|
||||||
|
{"qualified", "CREATE TABLE public.users (id int)", "users"},
|
||||||
|
{"if not exists", "CREATE TABLE IF NOT EXISTS public.users (id int)", "users"},
|
||||||
|
{"lowercase", "create table users(id int)", "users"},
|
||||||
|
{"newline", "CREATE TABLE\npublic.t\n(id int)", "t"},
|
||||||
|
{"no name", "CREATE TABLE", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := extractTableNameFromCreate(tt.in); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTruncateStatement(t *testing.T) {
|
||||||
|
short := strings.Repeat("a", 200)
|
||||||
|
if got := truncateStatement(short); got != short {
|
||||||
|
t.Errorf("200-char statement must not be truncated")
|
||||||
|
}
|
||||||
|
long := strings.Repeat("a", 201)
|
||||||
|
got := truncateStatement(long)
|
||||||
|
if got != strings.Repeat("a", 200)+"..." {
|
||||||
|
t.Errorf("unexpected truncation: len=%d", len(got))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetCurrentTimestamp(t *testing.T) {
|
||||||
|
ts := getCurrentTimestamp()
|
||||||
|
if len(ts) != len("2006-01-02 15:04:05") || ts[4] != '-' || ts[10] != ' ' || ts[13] != ':' {
|
||||||
|
t.Errorf("unexpected timestamp format %q", ts)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractStatementContext(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, in, want string
|
||||||
|
}{
|
||||||
|
{"do block", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' AND table_name = 'users' AND column_name = 'email') THEN NULL; END IF; END $$;`, "public.users (email)"},
|
||||||
|
{"do block constraint", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.table_constraints WHERE table_schema = 'public' AND table_name = 'users' AND constraint_name = 'uq_email') THEN NULL; END IF; END $$;`, "public.users [uq_email]"},
|
||||||
|
{"add column", `ALTER TABLE public.users ADD COLUMN "email" text`, "public.users (email)"},
|
||||||
|
{"alter column", `ALTER TABLE users ALTER COLUMN age SET NOT NULL`, "users (age)"},
|
||||||
|
{"add constraint", `ALTER TABLE public.users ADD CONSTRAINT uq_email UNIQUE (email)`, "public.users [uq_email]"},
|
||||||
|
{"drop constraint", `ALTER TABLE public.users DROP CONSTRAINT "uq_email"`, "public.users [uq_email]"},
|
||||||
|
{"alter table plain", `ALTER TABLE public.users RENAME TO people`, "public.users"},
|
||||||
|
{"create table", `CREATE TABLE public.users (id int)`, "public.users"},
|
||||||
|
{"create table if not exists", `CREATE TABLE IF NOT EXISTS "public"."users" (id int)`, "public.users"},
|
||||||
|
{"create schema", `CREATE SCHEMA IF_x;`, "IF_x"},
|
||||||
|
{"create index", `CREATE INDEX idx ON public.users (email)`, "public.users"},
|
||||||
|
{"create unique index", `CREATE UNIQUE INDEX idx ON users (email)`, "users"},
|
||||||
|
{"create index without on", `CREATE INDEX idx`, ""},
|
||||||
|
{"comment on table", `COMMENT ON TABLE public.users IS 'x'`, "public.users"},
|
||||||
|
{"comment on column", `COMMENT ON COLUMN public.users.email IS 'x'`, "public.users.email"},
|
||||||
|
{"unknown", `DROP TABLE users`, ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := extractStatementContext(tt.in); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractSQLStringValue(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, stmt, key, want string
|
||||||
|
}{
|
||||||
|
{"basic", "WHERE table_name = 'users'", "table_name", "users"},
|
||||||
|
{"case-insensitive key", "WHERE TABLE_NAME='users'", "table_name", "users"},
|
||||||
|
{"missing key", "WHERE a = 'b'", "table_name", ""},
|
||||||
|
{"no equals", "table_name is 'x'", "table_name", ""},
|
||||||
|
{"equals too far", "table_name abcdefgh = 'x'", "table_name", ""},
|
||||||
|
{"not quoted", "table_name = users", "table_name", ""},
|
||||||
|
{"unterminated", "table_name = 'users", "table_name", ""},
|
||||||
|
{"empty after key", "table_name", "table_name", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := extractSQLStringValue(tt.stmt, tt.key); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseQualifiedIdent(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
in, schema, name string
|
||||||
|
}{
|
||||||
|
{"users (id int)", "", "users"},
|
||||||
|
{"public.users (id int)", "public", "users"},
|
||||||
|
{`"public"."users" (id int)`, "public", "users"},
|
||||||
|
{`"users"`, "", "users"},
|
||||||
|
{"", "", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
s, n := parseQualifiedIdent(tt.in)
|
||||||
|
if s != tt.schema || n != tt.name {
|
||||||
|
t.Errorf("parseQualifiedIdent(%q) = (%q,%q), want (%q,%q)", tt.in, s, n, tt.schema, tt.name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFirstBareIdentAndHelpers(t *testing.T) {
|
||||||
|
bare := map[string]string{
|
||||||
|
"": "",
|
||||||
|
" ": "",
|
||||||
|
"abc": "abc",
|
||||||
|
"abc def": "abc",
|
||||||
|
"abc(def)": "abc",
|
||||||
|
"abc,def": "abc",
|
||||||
|
"abc;": "abc",
|
||||||
|
"\n abc\tdef": "abc",
|
||||||
|
`"a b" c`: `"a`,
|
||||||
|
" tbl (x int)": "tbl",
|
||||||
|
}
|
||||||
|
for in, want := range bare {
|
||||||
|
if got := firstBareIdent(in); got != want {
|
||||||
|
t.Errorf("firstBareIdent(%q) = %q, want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := stripQuotes(`"abc"`); got != "abc" {
|
||||||
|
t.Errorf("stripQuotes = %q", got)
|
||||||
|
}
|
||||||
|
if got := stripQuotes("abc"); got != "abc" {
|
||||||
|
t.Errorf("stripQuotes unquoted = %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
stmt := `ALTER TABLE t add column "c1" text`
|
||||||
|
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "ADD COLUMN"); got != "c1" {
|
||||||
|
t.Errorf("firstIdentAfterKeyword = %q", got)
|
||||||
|
}
|
||||||
|
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "DROP COLUMN"); got != "" {
|
||||||
|
t.Errorf("missing keyword must return empty, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildStmtContext(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
schema, table, column, constraint, want string
|
||||||
|
}{
|
||||||
|
{"", "", "", "", ""},
|
||||||
|
{"s", "t", "", "", "s.t"},
|
||||||
|
{"", "t", "", "", "t"},
|
||||||
|
{"s", "", "", "", ""},
|
||||||
|
{"s", "t", "c", "", "s.t (c)"},
|
||||||
|
{"s", "t", "", "k", "s.t [k]"},
|
||||||
|
{"s", "t", "c", "k", "s.t (c) [k]"},
|
||||||
|
{"", "", "c", "", "(c)"},
|
||||||
|
{"", "", "", "k", "[k]"},
|
||||||
|
{"", "", "c", "k", "(c) [k]"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := buildStmtContext(tt.schema, tt.table, tt.column, tt.constraint); got != tt.want {
|
||||||
|
t.Errorf("buildStmtContext(%q,%q,%q,%q) = %q, want %q", tt.schema, tt.table, tt.column, tt.constraint, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectStatementType(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, in, want string
|
||||||
|
}{
|
||||||
|
{"do unique", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT u UNIQUE (a); END $$", "ADD UNIQUE CONSTRAINT"},
|
||||||
|
{"do fk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id); END $$", "ADD FOREIGN KEY"},
|
||||||
|
{"do pk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a); END $$", "ADD PRIMARY KEY"},
|
||||||
|
{"do check", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c CHECK (a > 0); END $$", "ADD CHECK CONSTRAINT"},
|
||||||
|
{"do constraint", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a); END $$", "ADD CONSTRAINT"},
|
||||||
|
{"do add column", "DO $$ BEGIN ALTER TABLE t ADD COLUMN c int; END $$", "ADD COLUMN"},
|
||||||
|
{"do drop constraint", "DO $$ BEGIN DROP CONSTRAINT x; END $$", "DROP CONSTRAINT"},
|
||||||
|
{"do other", "DO $$ BEGIN NULL; END $$", "DO BLOCK"},
|
||||||
|
{"create schema", "create schema s", "CREATE SCHEMA"},
|
||||||
|
{"create sequence", "CREATE SEQUENCE s", "CREATE SEQUENCE"},
|
||||||
|
{"create table", "CREATE TABLE t ()", "CREATE TABLE"},
|
||||||
|
{"create index", "CREATE INDEX i ON t(a)", "CREATE INDEX"},
|
||||||
|
{"create unique index", "CREATE UNIQUE INDEX i ON t(a)", "CREATE UNIQUE INDEX"},
|
||||||
|
{"alter fk", "ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id)", "ADD FOREIGN KEY"},
|
||||||
|
{"alter pk", "ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a)", "ADD PRIMARY KEY"},
|
||||||
|
{"alter unique", "ALTER TABLE t ADD CONSTRAINT u UNIQUE (a)", "ADD UNIQUE CONSTRAINT"},
|
||||||
|
{"alter check", "ALTER TABLE t ADD CONSTRAINT c CHECK (a>0)", "ADD CHECK CONSTRAINT"},
|
||||||
|
{"alter constraint", "ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a)", "ADD CONSTRAINT"},
|
||||||
|
{"alter add column", "ALTER TABLE t ADD COLUMN c int", "ADD COLUMN"},
|
||||||
|
{"alter drop constraint", "ALTER TABLE t DROP CONSTRAINT c", "DROP CONSTRAINT"},
|
||||||
|
{"alter column", "ALTER TABLE t ALTER COLUMN c TYPE int", "ALTER COLUMN"},
|
||||||
|
{"alter table", "ALTER TABLE t RENAME TO u", "ALTER TABLE"},
|
||||||
|
{"comment table", "COMMENT ON TABLE t IS 'x'", "COMMENT ON TABLE"},
|
||||||
|
{"comment column", "COMMENT ON COLUMN t.c IS 'x'", "COMMENT ON COLUMN"},
|
||||||
|
{"drop table", "DROP TABLE t", "DROP TABLE"},
|
||||||
|
{"drop index", "DROP INDEX i", "DROP INDEX"},
|
||||||
|
{"default", "SELECT 1", "SQL"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := detectStatementType(tt.in); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteAndFinishReport(t *testing.T) {
|
||||||
|
report := &ExecutionReport{
|
||||||
|
TotalStatements: 3,
|
||||||
|
ExecutedStatements: 2,
|
||||||
|
FailedStatements: 1,
|
||||||
|
Schemas: []SchemaReport{{Name: "public", Tables: []TableReport{
|
||||||
|
{Name: "a", Created: true},
|
||||||
|
{Name: "b", Created: false, Error: "boom"},
|
||||||
|
}}},
|
||||||
|
Errors: []ExecutionError{{StatementNumber: 3, Statement: "CREATE TABLE b ()", Error: "boom"}},
|
||||||
|
StartTime: "s",
|
||||||
|
EndTime: "e",
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "report.json")
|
||||||
|
w := &Writer{
|
||||||
|
options: &writers.WriterOptions{Metadata: map[string]interface{}{"report_path": path}},
|
||||||
|
executionReport: report,
|
||||||
|
}
|
||||||
|
if err := w.finishReport(); err != nil {
|
||||||
|
t.Fatalf("finishReport: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("report not written: %v", err)
|
||||||
|
}
|
||||||
|
var got ExecutionReport
|
||||||
|
if err := json.Unmarshal(data, &got); err != nil {
|
||||||
|
t.Fatalf("invalid report JSON: %v", err)
|
||||||
|
}
|
||||||
|
if got.TotalStatements != 3 || got.FailedStatements != 1 || len(got.Errors) != 1 ||
|
||||||
|
len(got.Schemas) != 1 || len(got.Schemas[0].Tables) != 2 || got.Schemas[0].Tables[1].Error != "boom" {
|
||||||
|
t.Errorf("report round-trip mismatch: %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFinishReportNoPathAndSuccess(t *testing.T) {
|
||||||
|
w := &Writer{
|
||||||
|
options: &writers.WriterOptions{},
|
||||||
|
executionReport: &ExecutionReport{TotalStatements: 1, ExecutedStatements: 1},
|
||||||
|
}
|
||||||
|
if err := w.finishReport(); err != nil {
|
||||||
|
t.Errorf("finishReport without path: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteReportBadPath(t *testing.T) {
|
||||||
|
w := &Writer{options: &writers.WriterOptions{}, executionReport: &ExecutionReport{}}
|
||||||
|
if err := w.writeReport(filepath.Join(t.TempDir(), "missing", "r.json")); err == nil {
|
||||||
|
t.Error("expected error for unwritable path")
|
||||||
|
}
|
||||||
|
// finishReport must swallow the report error.
|
||||||
|
w.options.Metadata = map[string]interface{}{"report_path": filepath.Join(t.TempDir(), "missing", "r.json")}
|
||||||
|
if err := w.finishReport(); err != nil {
|
||||||
|
t.Errorf("finishReport must not fail on report write error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTemplateFilterAndMapFuncPassthrough(t *testing.T) {
|
||||||
|
in := []string{"a", "b"}
|
||||||
|
if got := filter(in, "X").([]string); len(got) != 2 {
|
||||||
|
t.Errorf("filter must return slice unchanged")
|
||||||
|
}
|
||||||
|
if got := mapFunc("v", "upper"); got != "v" {
|
||||||
|
t.Errorf("mapFunc must return value unchanged, got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
package prisma
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSQLTypeToPrisma(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
schema.Enums = append(schema.Enums, &models.Enum{Name: "Role", Values: []string{"A"}})
|
||||||
|
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"text", "String"}, {"varchar(255)", "String"}, {"character varying", "String"}, {"char(1)", "String"},
|
||||||
|
{"boolean", "Boolean"}, {"bool", "Boolean"},
|
||||||
|
{"integer", "Int"}, {"int", "Int"}, {"int4", "Int"},
|
||||||
|
{"bigint", "BigInt"}, {"int8", "BigInt"}, {"BIGINT", "BigInt"},
|
||||||
|
{"double precision", "Float"}, {"float8", "Float"},
|
||||||
|
{"numeric(10,2)", "Decimal"}, {"decimal", "Decimal"},
|
||||||
|
{"timestamp", "DateTime"}, {"timestamptz", "DateTime"}, {"date", "DateTime"},
|
||||||
|
{"jsonb", "Json"}, {"json", "Json"}, {"bytea", "Bytes"},
|
||||||
|
{"role", "Role"}, {"unknown_type", "String"},
|
||||||
|
}
|
||||||
|
// Repeat: the mapping used to depend on map iteration order.
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := w.sqlTypeToPrisma(tt.in, schema); got != tt.want {
|
||||||
|
t.Fatalf("sqlTypeToPrisma(%q) = %q, want %q (iteration %d)", tt.in, got, tt.want, i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -266,35 +266,37 @@ func (w *Writer) sqlTypeToPrisma(sqlType string, schema *models.Schema) string {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Standard type mapping
|
// Ordered so more specific patterns win (bigint/int8 before int); map
|
||||||
typeMap := map[string]string{
|
// iteration would make the result nondeterministic.
|
||||||
"text": "String",
|
typeMap := []struct{ pattern, prismaType string }{
|
||||||
"varchar": "String",
|
{"bigint", "BigInt"},
|
||||||
"character varying": "String",
|
{"int8", "BigInt"},
|
||||||
"char": "String",
|
{"text", "String"},
|
||||||
"boolean": "Boolean",
|
{"varchar", "String"},
|
||||||
"bool": "Boolean",
|
{"character varying", "String"},
|
||||||
"integer": "Int",
|
{"char", "String"},
|
||||||
"int": "Int",
|
{"boolean", "Boolean"},
|
||||||
"int4": "Int",
|
{"bool", "Boolean"},
|
||||||
"bigint": "BigInt",
|
{"integer", "Int"},
|
||||||
"int8": "BigInt",
|
{"int4", "Int"},
|
||||||
"double precision": "Float",
|
{"int", "Int"},
|
||||||
"float": "Float",
|
{"double precision", "Float"},
|
||||||
"float8": "Float",
|
{"float8", "Float"},
|
||||||
"decimal": "Decimal",
|
{"float", "Float"},
|
||||||
"numeric": "Decimal",
|
{"decimal", "Decimal"},
|
||||||
"timestamp": "DateTime",
|
{"numeric", "Decimal"},
|
||||||
"timestamptz": "DateTime",
|
{"timestamptz", "DateTime"},
|
||||||
"date": "DateTime",
|
{"timestamp", "DateTime"},
|
||||||
"jsonb": "Json",
|
{"date", "DateTime"},
|
||||||
"json": "Json",
|
{"jsonb", "Json"},
|
||||||
"bytea": "Bytes",
|
{"json", "Json"},
|
||||||
|
{"bytea", "Bytes"},
|
||||||
}
|
}
|
||||||
|
|
||||||
for sqlPattern, prismaType := range typeMap {
|
lower := strings.ToLower(sqlType)
|
||||||
if strings.Contains(strings.ToLower(sqlType), sqlPattern) {
|
for _, m := range typeMap {
|
||||||
return prismaType
|
if strings.Contains(lower, m.pattern) {
|
||||||
|
return m.prismaType
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,259 @@
|
|||||||
|
package prisma
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
rprisma "git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
const examplePrisma = "../../../tests/assets/prisma/example.prisma"
|
||||||
|
|
||||||
|
func readExample(t *testing.T) *models.Database {
|
||||||
|
t.Helper()
|
||||||
|
db, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: examplePrisma}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_ExampleToFile(t *testing.T) {
|
||||||
|
db := readExample(t)
|
||||||
|
out := filepath.Join(t.TempDir(), "schema.prisma")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := string(b)
|
||||||
|
for _, want := range []string{
|
||||||
|
"datasource db {", `provider = "postgresql"`, "generator client {",
|
||||||
|
"model User {", "model Post {", "model Category {", "model Profile {",
|
||||||
|
"enum Role {", " USER", " ADMIN",
|
||||||
|
"@id", "@unique", "@default(autoincrement())", "@default(now())", "@relation(",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("output missing %q\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||||
|
db := readExample(t)
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
first := w.databaseToPrisma(db)
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
if got := w.databaseToPrisma(db); got != first {
|
||||||
|
t.Fatalf("output differs on run %d", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_RoundTrip(t *testing.T) {
|
||||||
|
db := readExample(t)
|
||||||
|
out := filepath.Join(t.TempDir(), "schema.prisma")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
again, err := rprisma.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("re-read: %v", err)
|
||||||
|
}
|
||||||
|
names := func(d *models.Database) map[string]bool {
|
||||||
|
m := map[string]bool{}
|
||||||
|
for _, s := range d.Schemas {
|
||||||
|
for _, tb := range s.Tables {
|
||||||
|
m[tb.Name] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
a, b := names(db), names(again)
|
||||||
|
for n := range a {
|
||||||
|
if !b[n] {
|
||||||
|
t.Errorf("table %q lost in round trip (got %v)", n, b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteSchemaAndTable(t *testing.T) {
|
||||||
|
db := readExample(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
schemaOut := filepath.Join(dir, "s.prisma")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: schemaOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tableOut := filepath.Join(dir, "t.prisma")
|
||||||
|
tbl := db.Schemas[0].Tables[0]
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: tableOut}).WriteTable(tbl); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, _ := os.ReadFile(tableOut)
|
||||||
|
if !strings.Contains(string(b), "model "+tbl.Name+" {") {
|
||||||
|
t.Errorf("table output: %s", b)
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(schemaOut); err != nil || info.Size() == 0 {
|
||||||
|
t.Errorf("schema output: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "missing-dir", "x.prisma")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(models.InitDatabase("d")); err == nil {
|
||||||
|
t.Error("expected error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGenerateDatasource_Providers(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
tests := []struct {
|
||||||
|
dbType models.DatabaseType
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{models.PostgresqlDatabaseType, "postgresql"},
|
||||||
|
{models.MSSQLDatabaseType, "sqlserver"},
|
||||||
|
{models.SqlLiteDatabaseType, "sqlite"},
|
||||||
|
{"mysql", "mysql"},
|
||||||
|
{"", "postgresql"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
db.DatabaseType = tt.dbType
|
||||||
|
if got := w.generateDatasource(db); !strings.Contains(got, `provider = "`+tt.want+`"`) {
|
||||||
|
t.Errorf("%q: %s", tt.dbType, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatDefaultValue(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
tests := []struct {
|
||||||
|
in any
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"now()", "now()"}, {"gen_random_uuid()", "uuid()"}, {"uuid_generate_v4()", "uuid()"},
|
||||||
|
{"hello", `"hello"`}, {true, "true"}, {false, "false"},
|
||||||
|
{42, "42"}, {int64(7), "7"}, {1.5, "1.5"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := w.formatDefaultValue(tt.in); got != tt.want {
|
||||||
|
t.Errorf("formatDefaultValue(%v) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func joinTableSchema() *models.Schema {
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
mk := func(name string) *models.Table {
|
||||||
|
t := models.InitTable(name, "public")
|
||||||
|
id := models.InitColumn("id", name, "public")
|
||||||
|
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
|
||||||
|
t.Columns["id"] = id
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
post, cat := mk("Post"), mk("Category")
|
||||||
|
join := models.InitTable("_CategoryToPost", "public")
|
||||||
|
for _, c := range []string{"A", "B"} {
|
||||||
|
col := models.InitColumn(c, join.Name, "public")
|
||||||
|
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||||
|
join.Columns[c] = col
|
||||||
|
}
|
||||||
|
for name, target := range map[string]string{"fk_a": "Category", "fk_b": "Post"} {
|
||||||
|
col := "A"
|
||||||
|
if name == "fk_b" {
|
||||||
|
col = "B"
|
||||||
|
}
|
||||||
|
c := models.InitConstraint(name, models.ForeignKeyConstraint)
|
||||||
|
c.Columns, c.ReferencedTable, c.ReferencedSchema, c.ReferencedColumns = []string{col}, target, "public", []string{"id"}
|
||||||
|
join.Constraints[name] = c
|
||||||
|
}
|
||||||
|
s.Tables = append(s.Tables, post, cat, join)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentifyJoinTables(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
s := joinTableSchema()
|
||||||
|
got := w.identifyJoinTables(s)
|
||||||
|
if !got["_CategoryToPost"] || got["Post"] || got["Category"] {
|
||||||
|
t.Errorf("join tables: %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extra column disqualifies the join table.
|
||||||
|
extra := models.InitColumn("note", "_CategoryToPost", "public")
|
||||||
|
s.Tables[2].Columns["note"] = extra
|
||||||
|
if w.identifyJoinTables(s)["_CategoryToPost"] {
|
||||||
|
t.Error("table with extra column must not be a join table")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDatabaseToPrisma_SkipsJoinTables(t *testing.T) {
|
||||||
|
db := models.InitDatabase("d")
|
||||||
|
db.Schemas = append(db.Schemas, joinTableSchema())
|
||||||
|
out := NewWriter(&writers.WriterOptions{}).databaseToPrisma(db)
|
||||||
|
if strings.Contains(out, "model _CategoryToPost") {
|
||||||
|
t.Errorf("join table emitted as a model:\n%s", out)
|
||||||
|
}
|
||||||
|
if !strings.Contains(out, "model Post {") || !strings.Contains(out, "model Category {") {
|
||||||
|
t.Errorf("models missing:\n%s", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBlockAttributes(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
tbl := models.InitTable("Membership", "public")
|
||||||
|
for _, c := range []string{"user_id", "group_id"} {
|
||||||
|
col := models.InitColumn(c, "Membership", "public")
|
||||||
|
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||||
|
tbl.Columns[c] = col
|
||||||
|
}
|
||||||
|
u := models.InitConstraint("uq_pair", models.UniqueConstraint)
|
||||||
|
u.Columns = []string{"user_id", "group_id"}
|
||||||
|
tbl.Constraints["uq_pair"] = u
|
||||||
|
idx := models.InitIndex("idx_group", "Membership", "public")
|
||||||
|
idx.Columns = []string{"group_id"}
|
||||||
|
tbl.Indexes["idx_group"] = idx
|
||||||
|
|
||||||
|
got := w.generateBlockAttributes(tbl)
|
||||||
|
for _, want := range []string{"@@id(", "@@unique(", "@@index("} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("missing %q in:\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Composite PK columns must not carry a field-level @id.
|
||||||
|
if strings.Contains(w.generateFieldAttributes(tbl.Columns["user_id"], tbl), "@id") {
|
||||||
|
t.Error("composite pk column got @id")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFieldAttributes_UniqueAndUpdatedAt(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
tbl := models.InitTable("T", "public")
|
||||||
|
col := models.InitColumn("email", "T", "public")
|
||||||
|
col.Type = "text"
|
||||||
|
col.Comment = "@updatedAt"
|
||||||
|
col.Default = "x"
|
||||||
|
tbl.Columns["email"] = col
|
||||||
|
u := models.InitConstraint("uq", models.UniqueConstraint)
|
||||||
|
u.Columns = []string{"email"}
|
||||||
|
tbl.Constraints["uq"] = u
|
||||||
|
|
||||||
|
got := w.generateFieldAttributes(col, tbl)
|
||||||
|
for _, want := range []string{"@unique", `@default("x")`, "@updatedAt"} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("missing %q in %q", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if line := w.columnToField(col, tbl, models.InitSchema("public")); !strings.Contains(line, "String?") {
|
||||||
|
t.Errorf("nullable column must be optional: %q", line)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,223 @@
|
|||||||
|
package sqlexec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/assetloader"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWriter_Options(t *testing.T) {
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]interface{}{"k": "v"}}
|
||||||
|
if got := NewWriter(opts).Options(); got != opts {
|
||||||
|
t.Error("Options must return the same pointer")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriter_ConnectFailure(t *testing.T) {
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]interface{}{
|
||||||
|
"connection_string": "postgres://nobody:nopass@127.0.0.1:1/none?connect_timeout=1",
|
||||||
|
}}
|
||||||
|
w := NewWriter(opts)
|
||||||
|
scripts := []*models.Script{{Name: "s", SQL: "SELECT 1"}}
|
||||||
|
|
||||||
|
if err := w.WriteDatabase(&models.Database{Schemas: []*models.Schema{{Name: "public", Scripts: scripts}}}); err == nil ||
|
||||||
|
!strings.Contains(err.Error(), "failed to connect") {
|
||||||
|
t.Errorf("WriteDatabase: %v", err)
|
||||||
|
}
|
||||||
|
if err := w.WriteSchema(&models.Schema{Name: "public", Scripts: scripts}); err == nil ||
|
||||||
|
!strings.Contains(err.Error(), "failed to connect") {
|
||||||
|
t.Errorf("WriteSchema: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// liveConn returns a connection string for a live PostgreSQL or skips the test.
|
||||||
|
func liveConn(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||||
|
if conn == "" {
|
||||||
|
t.Skip("RELSPEC_TEST_PG_CONN not set")
|
||||||
|
}
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// liveSchema creates a throwaway schema and drops it on cleanup.
|
||||||
|
func liveSchema(t *testing.T, connString string) (string, *pgx.Conn) {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
conn, err := pgx.Connect(ctx, connString)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect: %v", err)
|
||||||
|
}
|
||||||
|
name := fmt.Sprintf("sqlexec_test_%d", time.Now().UnixNano())
|
||||||
|
if _, err := conn.Exec(ctx, "CREATE SCHEMA "+name); err != nil {
|
||||||
|
t.Fatalf("create schema: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
|
||||||
|
_ = conn.Close(ctx)
|
||||||
|
})
|
||||||
|
return name, conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func liveOptions(connString string, extra map[string]interface{}) *writers.WriterOptions {
|
||||||
|
meta := map[string]interface{}{"connection_string": connString}
|
||||||
|
for k, v := range extra {
|
||||||
|
meta[k] = v
|
||||||
|
}
|
||||||
|
return &writers.WriterOptions{Metadata: meta}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_ExecuteScriptsOrder(t *testing.T) {
|
||||||
|
connString := liveConn(t)
|
||||||
|
schema, conn := liveSchema(t, connString)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// Each script appends its own name; the resulting row order is the execution order.
|
||||||
|
mk := func(name string, prio int, seq uint) *models.Script {
|
||||||
|
return &models.Script{
|
||||||
|
Name: name, Priority: prio, Sequence: seq,
|
||||||
|
SQL: fmt.Sprintf("INSERT INTO %s.log(name) VALUES ('%s');", schema, name),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
scripts := []*models.Script{
|
||||||
|
{Name: "00_create", Priority: 0, SQL: fmt.Sprintf("CREATE TABLE %s.log(id serial primary key, name text);", schema)},
|
||||||
|
mk("c_late", 2, 1),
|
||||||
|
mk("b_prio1_seq2", 1, 2),
|
||||||
|
mk("a_prio1_seq1", 1, 1),
|
||||||
|
mk("a_same", 1, 3),
|
||||||
|
mk("b_same", 1, 3),
|
||||||
|
{Name: "empty", Priority: 1, Sequence: 0, SQL: ""},
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := liveOptions(connString, nil)
|
||||||
|
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
|
||||||
|
t.Fatalf("WriteSchema: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := conn.Query(ctx, fmt.Sprintf("SELECT name FROM %s.log ORDER BY id", schema))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var got []string
|
||||||
|
for rows.Next() {
|
||||||
|
var n string
|
||||||
|
if err := rows.Scan(&n); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got = append(got, n)
|
||||||
|
}
|
||||||
|
want := []string{"a_prio1_seq1", "b_prio1_seq2", "a_same", "b_same", "c_late"}
|
||||||
|
if strings.Join(got, ",") != strings.Join(want, ",") {
|
||||||
|
t.Errorf("execution order = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_total"] != 6 || opts.Metadata["execution_success"] != 6 || opts.Metadata["execution_failed"] != 0 {
|
||||||
|
t.Errorf("counts: %v", opts.Metadata)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_FailingScriptStops(t *testing.T) {
|
||||||
|
connString := liveConn(t)
|
||||||
|
schema, conn := liveSchema(t, connString)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
scripts := []*models.Script{
|
||||||
|
{Name: "01_ok", Priority: 1, SQL: fmt.Sprintf("CREATE TABLE %s.a(id int);", schema)},
|
||||||
|
{Name: "02_bad", Priority: 2, SQL: "SELECT * FROM definitely_missing_table;"},
|
||||||
|
{Name: "03_never", Priority: 3, SQL: fmt.Sprintf("CREATE TABLE %s.never(id int);", schema)},
|
||||||
|
}
|
||||||
|
err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: scripts})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "02_bad") {
|
||||||
|
t.Fatalf("expected failure naming 02_bad, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var exists bool
|
||||||
|
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".never").Scan(&exists); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if exists {
|
||||||
|
t.Error("script after the failure must not run")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_IgnoreErrorsContinues(t *testing.T) {
|
||||||
|
connString := liveConn(t)
|
||||||
|
schema, conn := liveSchema(t, connString)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
scripts := []*models.Script{
|
||||||
|
{Name: "01_bad", Priority: 1, SQL: "SELECT * FROM definitely_missing_table;"},
|
||||||
|
{Name: "02_ok", Priority: 2, SQL: fmt.Sprintf("CREATE TABLE %s.after(id int);", schema)},
|
||||||
|
}
|
||||||
|
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
|
||||||
|
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: scripts}); err != nil {
|
||||||
|
t.Fatalf("ignore_errors must not fail: %v", err)
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_total"] != 2 || opts.Metadata["execution_success"] != 1 || opts.Metadata["execution_failed"] != 1 {
|
||||||
|
t.Errorf("counts: %v", opts.Metadata)
|
||||||
|
}
|
||||||
|
var exists bool
|
||||||
|
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", schema+".after").Scan(&exists); err != nil || !exists {
|
||||||
|
t.Errorf("later script must run: exists=%v err=%v", exists, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_EmbedDirectiveErrorHandling(t *testing.T) {
|
||||||
|
connString := liveConn(t)
|
||||||
|
schema, _ := liveSchema(t, connString)
|
||||||
|
|
||||||
|
bad := models.InitScript("embed_bad")
|
||||||
|
bad.Priority = 1
|
||||||
|
bad.SQL = "-- @embed: path=missing.txt var=:body mode=text\nSELECT :body;"
|
||||||
|
bad.Metadata[assetloader.ScriptSourcePathMetadataKey] = filepath.Join(t.TempDir(), "s.sql")
|
||||||
|
if err := NewWriter(liveOptions(connString, nil)).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err == nil ||
|
||||||
|
!strings.Contains(err.Error(), "embed_bad") {
|
||||||
|
t.Errorf("expected error naming script, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
opts := liveOptions(connString, map[string]interface{}{"ignore_errors": true})
|
||||||
|
if err := NewWriter(opts).WriteSchema(&models.Schema{Name: schema, Scripts: []*models.Script{bad}}); err != nil {
|
||||||
|
t.Errorf("ignore_errors: %v", err)
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_failed"] != 1 {
|
||||||
|
t.Errorf("counts: %v", opts.Metadata)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLive_WriteDatabaseMultiSchema(t *testing.T) {
|
||||||
|
connString := liveConn(t)
|
||||||
|
s1, conn := liveSchema(t, connString)
|
||||||
|
s2, _ := liveSchema(t, connString)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
db := &models.Database{Schemas: []*models.Schema{
|
||||||
|
{Name: s1, Scripts: []*models.Script{{Name: "a", SQL: fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s.t(id int);", s1)}}},
|
||||||
|
{Name: s2, Scripts: []*models.Script{{Name: "b", SQL: fmt.Sprintf("CREATE TABLE %s.t(id int);", s2)}}},
|
||||||
|
}}
|
||||||
|
if err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, s := range []string{s1, s2} {
|
||||||
|
var ok bool
|
||||||
|
if err := conn.QueryRow(ctx, "SELECT to_regclass($1) IS NOT NULL", s+".t").Scan(&ok); err != nil || !ok {
|
||||||
|
t.Errorf("table in %s missing (err %v)", s, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A failure in one schema aborts and names that schema.
|
||||||
|
db.Schemas[1].Scripts[0].SQL = "SELECT * FROM definitely_missing_table;"
|
||||||
|
err := NewWriter(liveOptions(connString, nil)).WriteDatabase(db)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "schema "+s2) {
|
||||||
|
t.Errorf("expected error naming schema %s, got %v", s2, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -44,29 +44,37 @@ func (w *Writer) WriteDatabase(db *models.Database) error {
|
|||||||
return w.executeDatabaseSQL(db, dbPath)
|
return w.executeDatabaseSQL(db, dbPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
var writer io.Writer
|
release, err := w.openOutput()
|
||||||
var file *os.File
|
|
||||||
var err error
|
|
||||||
|
|
||||||
// Use existing writer if already set (for testing)
|
|
||||||
if w.writer != nil {
|
|
||||||
writer = w.writer
|
|
||||||
} else if w.options.OutputPath != "" {
|
|
||||||
// Determine output destination
|
|
||||||
file, err = os.Create(w.options.OutputPath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create output file: %w", err)
|
return err
|
||||||
}
|
|
||||||
defer file.Close()
|
|
||||||
writer = file
|
|
||||||
} else {
|
|
||||||
writer = os.Stdout
|
|
||||||
}
|
}
|
||||||
|
defer release()
|
||||||
|
|
||||||
w.writer = writer
|
|
||||||
return w.writeContent(db)
|
return w.writeContent(db)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// openOutput points w.writer at the configured destination (output file or
|
||||||
|
// stdout) when none is set, and returns a func that releases it again so the
|
||||||
|
// writer can be reused.
|
||||||
|
func (w *Writer) openOutput() (func(), error) {
|
||||||
|
if w.writer != nil {
|
||||||
|
return func() {}, nil
|
||||||
|
}
|
||||||
|
if w.options.OutputPath != "" {
|
||||||
|
file, err := os.Create(w.options.OutputPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create output file: %w", err)
|
||||||
|
}
|
||||||
|
w.writer = file
|
||||||
|
return func() {
|
||||||
|
file.Close()
|
||||||
|
w.writer = nil
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
w.writer = os.Stdout
|
||||||
|
return func() { w.writer = nil }, nil
|
||||||
|
}
|
||||||
|
|
||||||
// writeContent writes the header, pragma, and every schema's DDL to w.writer.
|
// writeContent writes the header, pragma, and every schema's DDL to w.writer.
|
||||||
func (w *Writer) writeContent(db *models.Database) error {
|
func (w *Writer) writeContent(db *models.Database) error {
|
||||||
// Write header comment
|
// Write header comment
|
||||||
@@ -184,6 +192,12 @@ func tableSchemaName(schema string) string {
|
|||||||
|
|
||||||
// WriteSchema writes a single schema as SQLite SQL
|
// WriteSchema writes a single schema as SQLite SQL
|
||||||
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
func (w *Writer) WriteSchema(schema *models.Schema) error {
|
||||||
|
release, err := w.openOutput()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer release()
|
||||||
|
|
||||||
tableSchema := tableSchemaName(schema.Name)
|
tableSchema := tableSchemaName(schema.Name)
|
||||||
|
|
||||||
if err := w.checkDirectives(schema); err != nil {
|
if err := w.checkDirectives(schema); err != nil {
|
||||||
@@ -229,6 +243,11 @@ func (w *Writer) WriteSchema(schema *models.Schema) error {
|
|||||||
|
|
||||||
// WriteTable writes a single table as SQLite SQL
|
// WriteTable writes a single table as SQLite SQL
|
||||||
func (w *Writer) WriteTable(table *models.Table) error {
|
func (w *Writer) WriteTable(table *models.Table) error {
|
||||||
|
release, err := w.openOutput()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer release()
|
||||||
return w.writeTable("", table)
|
return w.writeTable("", table)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,250 @@
|
|||||||
|
package sqlite
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
rdbml "git.warky.dev/wdevs/relspecgo/pkg/readers/dbml"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func shopDB() *models.Database {
|
||||||
|
s := models.InitSchema("public")
|
||||||
|
users := models.InitTable("users", "public")
|
||||||
|
id := models.InitColumn("id", "users", "public")
|
||||||
|
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement, id.Sequence = "integer", true, true, true, 1
|
||||||
|
email := models.InitColumn("email", "users", "public")
|
||||||
|
email.Type, email.NotNull, email.Sequence = "text", true, 2
|
||||||
|
age := models.InitColumn("age", "users", "public")
|
||||||
|
age.Type, age.Sequence, age.Default = "integer", 3, 18
|
||||||
|
users.Columns["id"], users.Columns["email"], users.Columns["age"] = id, email, age
|
||||||
|
|
||||||
|
uq := models.InitConstraint("uq_users_email", models.UniqueConstraint)
|
||||||
|
uq.Columns = []string{"email"}
|
||||||
|
ck := models.InitConstraint("ck_age", models.CheckConstraint)
|
||||||
|
ck.Expression = "age >= 0"
|
||||||
|
users.Constraints["uq_users_email"], users.Constraints["ck_age"] = uq, ck
|
||||||
|
ix := models.InitIndex("idx_users_age", "users", "public")
|
||||||
|
ix.Columns = []string{"age"}
|
||||||
|
uix := models.InitIndex("uidx_users_nick", "users", "public")
|
||||||
|
uix.Columns, uix.Unique = []string{"age", "email"}, true
|
||||||
|
pkIx := models.InitIndex("users_pkey", "users", "public")
|
||||||
|
pkIx.Columns = []string{"id"}
|
||||||
|
users.Indexes["idx_users_age"], users.Indexes["uidx_users_nick"], users.Indexes["users_pkey"] = ix, uix, pkIx
|
||||||
|
|
||||||
|
orders := models.InitTable("orders", "public")
|
||||||
|
oid := models.InitColumn("id", "orders", "public")
|
||||||
|
oid.Type, oid.IsPrimaryKey, oid.NotNull = "integer", true, true
|
||||||
|
uid := models.InitColumn("user_id", "orders", "public")
|
||||||
|
uid.Type, uid.NotNull = "integer", true
|
||||||
|
orders.Columns["id"], orders.Columns["user_id"] = oid, uid
|
||||||
|
fk := models.InitConstraint("fk_orders_users", models.ForeignKeyConstraint)
|
||||||
|
fk.Columns, fk.ReferencedTable, fk.ReferencedColumns = []string{"user_id"}, "users", []string{"id"}
|
||||||
|
orders.Constraints["fk_orders_users"] = fk
|
||||||
|
|
||||||
|
s.Tables = append(s.Tables, users, orders)
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func scriptFor(t *testing.T, db *models.Database) string {
|
||||||
|
t.Helper()
|
||||||
|
out := filepath.Join(t.TempDir(), "out.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Script(t *testing.T) {
|
||||||
|
got := scriptFor(t, shopDB())
|
||||||
|
for _, want := range []string{
|
||||||
|
"-- SQLite Database Schema", "-- Database: shop", "PRAGMA foreign_keys",
|
||||||
|
"CREATE TABLE", "users", "orders", "CREATE INDEX", "idx_users_age", "CREATE UNIQUE INDEX",
|
||||||
|
} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("missing %q\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if strings.Contains(got, "users_pkey") {
|
||||||
|
t.Errorf("pkey index must be skipped:\n%s", got)
|
||||||
|
}
|
||||||
|
if strings.Contains(got, "-- Schema: public") {
|
||||||
|
t.Errorf("default schema must not be announced:\n%s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_Deterministic(t *testing.T) {
|
||||||
|
first := scriptFor(t, shopDB())
|
||||||
|
for i := 0; i < 15; i++ {
|
||||||
|
if got := scriptFor(t, shopDB()); got != first {
|
||||||
|
t.Fatalf("output differs on run %d", i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriter_ReusableAfterFileOutput(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "o.sql")
|
||||||
|
w := NewWriter(&writers.WriterOptions{OutputPath: out})
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
if err := w.WriteDatabase(shopDB()); err != nil {
|
||||||
|
t.Fatalf("write %d: %v", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteSchemaAndTable_UseOutputPath(t *testing.T) {
|
||||||
|
db := shopDB()
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
sOut := filepath.Join(dir, "s.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if b, _ := os.ReadFile(sOut); !strings.Contains(string(b), "CREATE TABLE") {
|
||||||
|
t.Errorf("schema output:\n%s", b)
|
||||||
|
}
|
||||||
|
|
||||||
|
tOut := filepath.Join(dir, "t.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if b, _ := os.ReadFile(tOut); !strings.Contains(string(b), "CREATE TABLE") {
|
||||||
|
t.Errorf("table output:\n%s", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
||||||
|
bad := filepath.Join(t.TempDir(), "missing", "x.sql")
|
||||||
|
if err := NewWriter(&writers.WriterOptions{OutputPath: bad}).WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "failed to create output file") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteAgainstSQLiteFile(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "shop.db")
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
|
||||||
|
if err := NewWriter(opts).WriteDatabase(shopDB()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_failed"] != 0 || opts.Metadata["execution_success"].(int) == 0 {
|
||||||
|
t.Errorf("metadata: %+v", opts.Metadata)
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := sql.Open("sqlite", path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
for _, tbl := range []string{"users", "orders"} {
|
||||||
|
var n string
|
||||||
|
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name=?`, tbl).Scan(&n); err != nil {
|
||||||
|
t.Errorf("table %s not created: %v", tbl, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var idx int
|
||||||
|
if err := conn.QueryRow(`SELECT count(*) FROM sqlite_master WHERE type='index' AND name IN ('idx_users_age','uidx_users_nick','uq_users_email')`).Scan(&idx); err != nil || idx != 3 {
|
||||||
|
t.Errorf("indexes created: %d (%v)", idx, err)
|
||||||
|
}
|
||||||
|
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err != nil {
|
||||||
|
t.Errorf("insert: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := conn.Exec(`INSERT INTO users(email) VALUES('a@x')`); err == nil {
|
||||||
|
t.Error("unique constraint on email must be enforced")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecute_StopsOnErrorUnlessIgnored(t *testing.T) {
|
||||||
|
// Pre-create "users" so the first CREATE TABLE fails.
|
||||||
|
prepare := func(t *testing.T) string {
|
||||||
|
path := filepath.Join(t.TempDir(), "pre.db")
|
||||||
|
conn, err := sql.Open("sqlite", path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
if _, err := conn.Exec(`CREATE TABLE users (x int)`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
|
path := prepare(t)
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path}}
|
||||||
|
err := NewWriter(opts).WriteDatabase(shopDB())
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "failed to execute") {
|
||||||
|
t.Fatalf("expected failure, got %v", err)
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_failed"] != 1 {
|
||||||
|
t.Errorf("must stop at first failure: %+v", opts.Metadata)
|
||||||
|
}
|
||||||
|
|
||||||
|
path = prepare(t)
|
||||||
|
opts = &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
|
||||||
|
err = NewWriter(opts).WriteDatabase(shopDB())
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("errors are still reported when ignored")
|
||||||
|
}
|
||||||
|
if opts.Metadata["execution_success"].(int) == 0 || opts.Metadata["execution_failed"].(int) == 0 {
|
||||||
|
t.Errorf("ignore_errors must continue past failures: %+v", opts.Metadata)
|
||||||
|
}
|
||||||
|
conn, _ := sql.Open("sqlite", path)
|
||||||
|
defer conn.Close()
|
||||||
|
var n string
|
||||||
|
if err := conn.QueryRow(`SELECT name FROM sqlite_master WHERE name='orders'`).Scan(&n); err != nil {
|
||||||
|
t.Errorf("orders must still be created: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTruncateStatement(t *testing.T) {
|
||||||
|
if got := truncateStatement("CREATE TABLE\n x"); got != "CREATE TABLE x" {
|
||||||
|
t.Errorf("collapse: %q", got)
|
||||||
|
}
|
||||||
|
long := strings.Repeat("a", 200)
|
||||||
|
if got := truncateStatement(long); len(got) != 83 || !strings.HasSuffix(got, "...") {
|
||||||
|
t.Errorf("truncate: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTableSchemaName(t *testing.T) {
|
||||||
|
for in, want := range map[string]string{"public": "", "PUBLIC": "", "main": "", "auth": "auth", "": ""} {
|
||||||
|
if got := tableSchemaName(in); got != want {
|
||||||
|
t.Errorf("tableSchemaName(%q) = %q, want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckConstraintsWrittenAsComments(t *testing.T) {
|
||||||
|
w := NewWriter(&writers.WriterOptions{})
|
||||||
|
var sb strings.Builder
|
||||||
|
w.writer = &sb
|
||||||
|
if err := w.writeCheckConstraints("", shopDB().Schemas[0].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := sb.String(); !strings.Contains(got, "ck_age") || !strings.Contains(got, "age >= 0") {
|
||||||
|
t.Errorf("check output: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDBMLFixtureExecutes(t *testing.T) {
|
||||||
|
db, err := rdbml.NewReader(&readers.ReaderOptions{FilePath: "../../../tests/assets/dbml/complex.dbml"}).ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
path := filepath.Join(t.TempDir(), "complex.db")
|
||||||
|
opts := &writers.WriterOptions{Metadata: map[string]any{"connection_string": path, "ignore_errors": true}}
|
||||||
|
_ = NewWriter(opts).WriteDatabase(db)
|
||||||
|
if opts.Metadata["execution_success"].(int) == 0 {
|
||||||
|
t.Errorf("nothing executed: %+v", opts.Metadata)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTemplateError(t *testing.T) {
|
||||||
|
cause := errors.New("boom")
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
err *TemplateError
|
||||||
|
phase string
|
||||||
|
}{
|
||||||
|
{"load", NewTemplateLoadError("cannot read", cause), "load"},
|
||||||
|
{"parse", NewTemplateParseError("bad syntax", cause), "parse"},
|
||||||
|
{"execute", NewTemplateExecuteError("failed render", cause), "execute"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if tt.err.Phase != tt.phase {
|
||||||
|
t.Errorf("phase = %q", tt.err.Phase)
|
||||||
|
}
|
||||||
|
msg := tt.err.Error()
|
||||||
|
if !strings.Contains(msg, "template "+tt.phase+" error") || !strings.Contains(msg, "boom") {
|
||||||
|
t.Errorf("message = %q", msg)
|
||||||
|
}
|
||||||
|
if !errors.Is(tt.err, cause) {
|
||||||
|
t.Error("errors.Is must reach cause")
|
||||||
|
}
|
||||||
|
var te *TemplateError
|
||||||
|
if !errors.As(error(tt.err), &te) || te != tt.err {
|
||||||
|
t.Error("errors.As failed")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTemplateErrorWithoutCause(t *testing.T) {
|
||||||
|
e := NewTemplateParseError("only message", nil)
|
||||||
|
if got := e.Error(); got != "template parse error: only message" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if e.Unwrap() != nil {
|
||||||
|
t.Error("Unwrap must be nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sort"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func colNames(cols []*models.Column) []string {
|
||||||
|
out := make([]string, 0, len(cols))
|
||||||
|
for _, c := range cols {
|
||||||
|
out = append(out, c.Name)
|
||||||
|
}
|
||||||
|
sort.Strings(out)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func eqStrings(a, b []string) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for i := range a {
|
||||||
|
if a[i] != b[i] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func testColumns() map[string]*models.Column {
|
||||||
|
return map[string]*models.Column{
|
||||||
|
"id": {Name: "id", Type: "integer", IsPrimaryKey: true, NotNull: true},
|
||||||
|
"user_id": {Name: "user_id", Type: "bigint", NotNull: true},
|
||||||
|
"name": {Name: "name", Type: "varchar(50)"},
|
||||||
|
"email": {Name: "email", Type: "varchar(255)", NotNull: true},
|
||||||
|
"created_at": {Name: "created_at", Type: "timestamp"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterTables(t *testing.T) {
|
||||||
|
tables := []*models.Table{{Name: "user_profile"}, {Name: "user_settings"}, {Name: "orders"}}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in []*models.Table
|
||||||
|
pattern string
|
||||||
|
want []string
|
||||||
|
}{
|
||||||
|
{"empty pattern returns all", tables, "", []string{"user_profile", "user_settings", "orders"}},
|
||||||
|
{"glob", tables, "user_*", []string{"user_profile", "user_settings"}},
|
||||||
|
{"single char", tables, "order?", []string{"orders"}},
|
||||||
|
{"no match", tables, "zzz*", []string{}},
|
||||||
|
{"nil input", nil, "x*", []string{}},
|
||||||
|
{"invalid pattern falls back to exact", []*models.Table{{Name: "[a"}}, "[a", []string{"[a"}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := FilterTables(tt.in, tt.pattern)
|
||||||
|
names := []string{}
|
||||||
|
for _, tbl := range got {
|
||||||
|
names = append(names, tbl.Name)
|
||||||
|
}
|
||||||
|
if !eqStrings(names, tt.want) {
|
||||||
|
t.Errorf("got %v, want %v", names, tt.want)
|
||||||
|
}
|
||||||
|
byPattern := FilterTablesByPattern(tt.in, tt.pattern)
|
||||||
|
if len(byPattern) != len(got) {
|
||||||
|
t.Errorf("FilterTablesByPattern differs from FilterTables")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterColumns(t *testing.T) {
|
||||||
|
cols := testColumns()
|
||||||
|
tests := []struct {
|
||||||
|
pattern string
|
||||||
|
want []string
|
||||||
|
}{
|
||||||
|
{"", []string{"created_at", "email", "id", "name", "user_id"}},
|
||||||
|
{"*_id", []string{"user_id"}},
|
||||||
|
{"*", []string{"created_at", "email", "id", "name", "user_id"}},
|
||||||
|
{"nomatch", []string{}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := colNames(FilterColumns(cols, tt.pattern)); !eqStrings(got, tt.want) {
|
||||||
|
t.Errorf("pattern %q: got %v, want %v", tt.pattern, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := FilterColumns(nil, "*"); len(got) != 0 {
|
||||||
|
t.Errorf("nil map must yield empty result")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterColumnsByType(t *testing.T) {
|
||||||
|
cols := testColumns()
|
||||||
|
if got := colNames(FilterColumnsByType(cols, "varchar")); !eqStrings(got, []string{"email", "name"}) {
|
||||||
|
t.Errorf("varchar: got %v", got)
|
||||||
|
}
|
||||||
|
if got := colNames(FilterColumnsByType(cols, "varchar(10)")); !eqStrings(got, []string{"email", "name"}) {
|
||||||
|
t.Errorf("varchar(10) must match on base type, got %v", got)
|
||||||
|
}
|
||||||
|
if got := FilterColumnsByType(cols, "jsonb"); len(got) != 0 {
|
||||||
|
t.Errorf("jsonb: expected none, got %v", colNames(got))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterColumnFlags(t *testing.T) {
|
||||||
|
cols := testColumns()
|
||||||
|
if got := colNames(FilterPrimaryKeys(cols)); !eqStrings(got, []string{"id"}) {
|
||||||
|
t.Errorf("pks: %v", got)
|
||||||
|
}
|
||||||
|
if got := colNames(FilterNullable(cols)); !eqStrings(got, []string{"created_at", "name"}) {
|
||||||
|
t.Errorf("nullable: %v", got)
|
||||||
|
}
|
||||||
|
if got := colNames(FilterNotNull(cols)); !eqStrings(got, []string{"email", "id", "user_id"}) {
|
||||||
|
t.Errorf("notnull: %v", got)
|
||||||
|
}
|
||||||
|
for _, f := range []func(map[string]*models.Column) []*models.Column{FilterPrimaryKeys, FilterNullable, FilterNotNull} {
|
||||||
|
if got := f(nil); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("nil map must give non-nil empty slice")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterConstraints(t *testing.T) {
|
||||||
|
cons := map[string]*models.Constraint{
|
||||||
|
"pk": {Name: "pk", Type: models.PrimaryKeyConstraint},
|
||||||
|
"fk": {Name: "fk", Type: models.ForeignKeyConstraint},
|
||||||
|
"u1": {Name: "u1", Type: models.UniqueConstraint},
|
||||||
|
"u2": {Name: "u2", Type: models.UniqueConstraint},
|
||||||
|
"ck": {Name: "ck", Type: models.CheckConstraint},
|
||||||
|
}
|
||||||
|
count := func(f func(map[string]*models.Constraint) []*models.Constraint) int { return len(f(cons)) }
|
||||||
|
if n := count(FilterForeignKeys); n != 1 {
|
||||||
|
t.Errorf("fk count %d", n)
|
||||||
|
}
|
||||||
|
if n := count(FilterUniqueConstraints); n != 2 {
|
||||||
|
t.Errorf("unique count %d", n)
|
||||||
|
}
|
||||||
|
if n := count(FilterCheckConstraints); n != 1 {
|
||||||
|
t.Errorf("check count %d", n)
|
||||||
|
}
|
||||||
|
for _, f := range []func(map[string]*models.Constraint) []*models.Constraint{FilterForeignKeys, FilterUniqueConstraints, FilterCheckConstraints} {
|
||||||
|
if got := f(nil); got == nil || len(got) != 0 {
|
||||||
|
t.Errorf("nil map must give non-nil empty slice")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchPattern(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
s, pattern string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"user_profile", "user_*", true},
|
||||||
|
{"user", "user_*", false},
|
||||||
|
{"ab", "a?", true},
|
||||||
|
{"abc", "a?", false},
|
||||||
|
{"[a", "[A", true}, // invalid glob: case-insensitive exact
|
||||||
|
{"x", "[a", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := matchPattern(tt.s, tt.pattern); got != tt.want {
|
||||||
|
t.Errorf("matchPattern(%q,%q) = %v, want %v", tt.s, tt.pattern, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -30,7 +30,13 @@ func ToJSONPretty(v interface{}, indent string) string {
|
|||||||
|
|
||||||
// ToYAML converts a value to YAML string
|
// ToYAML converts a value to YAML string
|
||||||
// Usage: {{ .Database | toYAML }}
|
// Usage: {{ .Database | toYAML }}
|
||||||
func ToYAML(v interface{}) string {
|
func ToYAML(v interface{}) (out string) {
|
||||||
|
// yaml.v3 panics (rather than returning an error) for unsupported types such as channels.
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
out = fmt.Sprintf("error: failed to marshal: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
data, err := yaml.Marshal(v)
|
data, err := yaml.Marshal(v)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Sprintf("error: failed to marshal: %v", err)
|
return fmt.Sprintf("error: failed to marshal: %v", err)
|
||||||
|
|||||||
@@ -0,0 +1,118 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestToJSON(t *testing.T) {
|
||||||
|
if got := ToJSON(map[string]int{"a": 1}); got != `{"a":1}` {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := ToJSON(nil); got != "null" {
|
||||||
|
t.Errorf("nil: %q", got)
|
||||||
|
}
|
||||||
|
if got := ToJSON(math.Inf(1)); !strings.HasPrefix(got, `{"error": "failed to marshal`) {
|
||||||
|
t.Errorf("marshal failure: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToJSONPretty(t *testing.T) {
|
||||||
|
got := ToJSONPretty(map[string]int{"a": 1}, " ")
|
||||||
|
if got != "{\n \"a\": 1\n}" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := ToJSONPretty(make(chan int), " "); !strings.HasPrefix(got, `{"error"`) {
|
||||||
|
t.Errorf("marshal failure: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToYAML(t *testing.T) {
|
||||||
|
if got := ToYAML(map[string]int{"a": 1}); got != "a: 1\n" {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := ToYAML(make(chan int)); !strings.HasPrefix(got, "error: failed to marshal") {
|
||||||
|
// yaml.v3 panics-recovers into an error for unsupported types
|
||||||
|
t.Errorf("marshal failure: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIndent(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
in string
|
||||||
|
spaces int
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"", 4, ""},
|
||||||
|
{"a", 2, " a"},
|
||||||
|
{"a\nb", 2, " a\n b"},
|
||||||
|
{"a\n\nb", 2, " a\n\n b"},
|
||||||
|
{"a", 0, "a"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Indent(tt.in, tt.spaces); got != tt.want {
|
||||||
|
t.Errorf("Indent(%q,%d) = %q, want %q", tt.in, tt.spaces, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := IndentWith("", ">"); got != "" {
|
||||||
|
t.Errorf("IndentWith empty: %q", got)
|
||||||
|
}
|
||||||
|
if got := IndentWith("a\n\nb", "> "); got != "> a\n\n> b" {
|
||||||
|
t.Errorf("IndentWith: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEscape(t *testing.T) {
|
||||||
|
if got := Escape("a\"b\\c\nd\re\tf"); got != `a\"b\\c\nd\re\tf` {
|
||||||
|
t.Errorf("got %q", got)
|
||||||
|
}
|
||||||
|
if got := Escape(""); got != "" {
|
||||||
|
t.Errorf("empty: %q", got)
|
||||||
|
}
|
||||||
|
if got := EscapeQuotes(`a"b'c`); got != `a\"b\'c` {
|
||||||
|
t.Errorf("EscapeQuotes: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComment(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, in, style, want string
|
||||||
|
}{
|
||||||
|
{"empty", "", "//", ""},
|
||||||
|
{"slashes", "a\nb", "//", "// a\n// b"},
|
||||||
|
{"hash", "a", "#", "# a"},
|
||||||
|
{"sql", "a\nb", "--", "-- a\n-- b"},
|
||||||
|
{"block single", "a", "/* */", "/* a */"},
|
||||||
|
{"block single alt", "a", "/**/", "/* a */"},
|
||||||
|
{"block multi", "a\nb", "/* */", "/*\n * a\n * b\n */"},
|
||||||
|
{"default", "a", "weird", "// a"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := Comment(tt.in, tt.style); got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQuoteUnquote(t *testing.T) {
|
||||||
|
if got := QuoteString("a"); got != `"a"` {
|
||||||
|
t.Errorf("QuoteString: %q", got)
|
||||||
|
}
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{`"a"`, "a"},
|
||||||
|
{`'a'`, "a"},
|
||||||
|
{`""`, ""},
|
||||||
|
{`"a'`, `"a'`},
|
||||||
|
{`a`, `a`},
|
||||||
|
{`"`, `"`},
|
||||||
|
{"", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := UnquoteString(tt.in); got != tt.want {
|
||||||
|
t.Errorf("UnquoteString(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,75 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
"text/template"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBuildFuncMapEntriesAreFunctions(t *testing.T) {
|
||||||
|
fm := BuildFuncMap()
|
||||||
|
if len(fm) < 100 {
|
||||||
|
t.Errorf("unexpectedly small func map: %d", len(fm))
|
||||||
|
}
|
||||||
|
for name, fn := range fm {
|
||||||
|
if reflect.TypeOf(fn).Kind() != reflect.Func {
|
||||||
|
t.Errorf("%s is not a function", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, name := range []string{"toSnakeCase", "sqlToGo", "filterTables", "toJSON", "enumerate", "get", "sortTablesByName", "dict", "seq"} {
|
||||||
|
if _, ok := fm[name]; !ok {
|
||||||
|
t.Errorf("missing %s", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Must be accepted by text/template (valid names and signatures).
|
||||||
|
if _, err := template.New("x").Funcs(fm).Parse("ok"); err != nil {
|
||||||
|
t.Fatalf("funcmap rejected by text/template: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildFuncMapRender(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, tmpl, want string
|
||||||
|
}{
|
||||||
|
{"add", `{{add 2 3}}`, "5"},
|
||||||
|
{"sub", `{{sub 5 3}}`, "2"},
|
||||||
|
{"mul", `{{mul 2 3}}`, "6"},
|
||||||
|
{"div", `{{div 6 3}}`, "2"},
|
||||||
|
{"div zero", `{{div 6 0}}`, "0"},
|
||||||
|
{"mod", `{{mod 7 3}}`, "1"},
|
||||||
|
{"mod zero", `{{mod 7 0}}`, "0"},
|
||||||
|
{"default nil", `{{default "d" .Missing}}`, "d"},
|
||||||
|
{"default set", `{{default "d" "v"}}`, "v"},
|
||||||
|
{"dict", `{{get (dict "a" 1) "a"}}`, "1"},
|
||||||
|
{"dict odd", `{{if dict "a"}}set{{else}}nil{{end}}`, "nil"},
|
||||||
|
{"dict non-string key", `{{if dict 1 2}}set{{else}}nil{{end}}`, "nil"},
|
||||||
|
{"list", `{{len (list 1 2 3)}}`, "3"},
|
||||||
|
{"seq", `{{range seq 1 3}}{{.}}{{end}}`, "123"},
|
||||||
|
{"seq reversed", `{{len (seq 3 1)}}`, "0"},
|
||||||
|
{"snake", `{{toSnakeCase "UserName"}}`, "user_name"},
|
||||||
|
{"pluralize", `{{pluralize "category"}}`, "categories"},
|
||||||
|
{"sqlToGo", `{{sqlToGo "integer" true}}`, ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
tpl, err := template.New("t").Funcs(BuildFuncMap()).Parse(tt.tmpl)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse: %v", err)
|
||||||
|
}
|
||||||
|
var buf bytes.Buffer
|
||||||
|
if err := tpl.Execute(&buf, map[string]interface{}{}); err != nil {
|
||||||
|
t.Fatalf("execute: %v", err)
|
||||||
|
}
|
||||||
|
if tt.name == "sqlToGo" {
|
||||||
|
if buf.Len() == 0 {
|
||||||
|
t.Error("sqlToGo rendered nothing")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if buf.String() != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", buf.String(), tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,142 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
type loopItem struct {
|
||||||
|
Name string
|
||||||
|
Group string
|
||||||
|
N int
|
||||||
|
}
|
||||||
|
|
||||||
|
func ints(vs ...interface{}) []interface{} { return vs }
|
||||||
|
|
||||||
|
func TestEnumerate(t *testing.T) {
|
||||||
|
got := Enumerate([]string{"a", "b"})
|
||||||
|
want := []EnumeratedItem{{0, "a"}, {1, "b"}}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
if got := Enumerate([2]int{5, 6}); len(got) != 2 || got[1].Value != 6 {
|
||||||
|
t.Errorf("array: %v", got)
|
||||||
|
}
|
||||||
|
if got := Enumerate("nope"); len(got) != 0 {
|
||||||
|
t.Errorf("non-slice: %v", got)
|
||||||
|
}
|
||||||
|
if got := Enumerate(nil); len(got) != 0 {
|
||||||
|
t.Errorf("nil: %v", got)
|
||||||
|
}
|
||||||
|
if got := Enumerate([]int{}); len(got) != 0 {
|
||||||
|
t.Errorf("empty: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchChunk(t *testing.T) {
|
||||||
|
in := []int{1, 2, 3, 4, 5}
|
||||||
|
got := Batch(in, 2)
|
||||||
|
want := [][]interface{}{{1, 2}, {3, 4}, {5}}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
if got := Chunk(in, 10); len(got) != 1 || len(got[0]) != 5 {
|
||||||
|
t.Errorf("size > len: %v", got)
|
||||||
|
}
|
||||||
|
for _, size := range []int{0, -1} {
|
||||||
|
if got := Batch(in, size); len(got) != 0 {
|
||||||
|
t.Errorf("size %d: %v", size, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := Batch([]int{}, 2); len(got) != 0 {
|
||||||
|
t.Errorf("empty: %v", got)
|
||||||
|
}
|
||||||
|
if got := Batch("x", 2); len(got) != 0 {
|
||||||
|
t.Errorf("non-slice: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReverseFirstLastSkipTake(t *testing.T) {
|
||||||
|
in := []int{1, 2, 3, 4}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
got []interface{}
|
||||||
|
want []interface{}
|
||||||
|
}{
|
||||||
|
{"reverse", Reverse(in), ints(4, 3, 2, 1)},
|
||||||
|
{"reverse empty", Reverse([]int{}), ints()},
|
||||||
|
{"reverse non-slice", Reverse(5), ints()},
|
||||||
|
{"first 2", First(in, 2), ints(1, 2)},
|
||||||
|
{"first n>len", First(in, 9), ints(1, 2, 3, 4)},
|
||||||
|
{"first 0", First(in, 0), ints()},
|
||||||
|
{"first non-slice", First(5, 1), ints()},
|
||||||
|
{"last 2", Last(in, 2), ints(3, 4)},
|
||||||
|
{"last n>len", Last(in, 9), ints(1, 2, 3, 4)},
|
||||||
|
{"last neg", Last(in, -1), ints()},
|
||||||
|
{"last non-slice", Last(5, 1), ints()},
|
||||||
|
{"skip 1", Skip(in, 1), ints(2, 3, 4)},
|
||||||
|
{"skip neg", Skip(in, -3), ints(1, 2, 3, 4)},
|
||||||
|
{"skip all", Skip(in, 4), ints()},
|
||||||
|
{"skip n>len", Skip(in, 10), ints()},
|
||||||
|
{"skip non-slice", Skip(5, 1), ints()},
|
||||||
|
{"take", Take(in, 3), ints(1, 2, 3)},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if len(tt.got) != len(tt.want) || (len(tt.want) > 0 && !reflect.DeepEqual(tt.got, tt.want)) {
|
||||||
|
t.Errorf("got %v, want %v", tt.got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcatUnique(t *testing.T) {
|
||||||
|
got := Concat([]int{1, 2}, []string{"a"}, 5, nil, [1]int{9})
|
||||||
|
if !reflect.DeepEqual(got, ints(1, 2, "a", 9)) {
|
||||||
|
t.Errorf("concat: %v", got)
|
||||||
|
}
|
||||||
|
if got := Concat(); len(got) != 0 {
|
||||||
|
t.Errorf("concat none: %v", got)
|
||||||
|
}
|
||||||
|
if got := Unique([]int{1, 2, 1, 3, 2}); !reflect.DeepEqual(got, ints(1, 2, 3)) {
|
||||||
|
t.Errorf("unique: %v", got)
|
||||||
|
}
|
||||||
|
if got := Unique("x"); len(got) != 0 {
|
||||||
|
t.Errorf("unique non-slice: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortByGroupByCountIf(t *testing.T) {
|
||||||
|
items := []loopItem{{"c", "x", 3}, {"a", "y", 1}, {"b", "x", 2}}
|
||||||
|
|
||||||
|
sorted := SortBy(items, "Name")
|
||||||
|
if sorted[0].(loopItem).Name != "a" || sorted[2].(loopItem).Name != "c" {
|
||||||
|
t.Errorf("sortBy Name: %v", sorted)
|
||||||
|
}
|
||||||
|
sorted = SortBy(items, "N")
|
||||||
|
if sorted[0].(loopItem).N != 1 || sorted[2].(loopItem).N != 3 {
|
||||||
|
t.Errorf("sortBy N: %v", sorted)
|
||||||
|
}
|
||||||
|
if items[0].Name != "c" {
|
||||||
|
t.Errorf("SortBy must not mutate input")
|
||||||
|
}
|
||||||
|
if got := SortBy(5, "Name"); len(got) != 0 {
|
||||||
|
t.Errorf("sortBy non-slice")
|
||||||
|
}
|
||||||
|
|
||||||
|
groups := GroupBy(items, "Group")
|
||||||
|
if len(groups) != 2 || len(groups["x"]) != 2 || len(groups["y"]) != 1 {
|
||||||
|
t.Errorf("groupBy: %v", groups)
|
||||||
|
}
|
||||||
|
if got := GroupBy(5, "Group"); len(got) != 0 {
|
||||||
|
t.Errorf("groupBy non-slice")
|
||||||
|
}
|
||||||
|
|
||||||
|
n := CountIf(items, func(v interface{}) bool { return v.(loopItem).Group == "x" })
|
||||||
|
if n != 2 {
|
||||||
|
t.Errorf("countIf: %d", n)
|
||||||
|
}
|
||||||
|
if got := CountIf(5, func(interface{}) bool { return true }); got != 0 {
|
||||||
|
t.Errorf("countIf non-slice: %d", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -101,11 +101,8 @@ func Merge(maps ...interface{}) map[interface{}]interface{} {
|
|||||||
for _, m := range maps {
|
for _, m := range maps {
|
||||||
v := reflect.ValueOf(m)
|
v := reflect.ValueOf(m)
|
||||||
|
|
||||||
// Dereference pointers
|
// Dereference pointers; a nil pointer contributes nothing
|
||||||
for v.Kind() == reflect.Pointer {
|
for v.Kind() == reflect.Pointer && !v.IsNil() {
|
||||||
if v.IsNil() {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
v = v.Elem()
|
v = v.Elem()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,216 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
type accessItem struct {
|
||||||
|
Name string
|
||||||
|
ID int
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetAndGetOr(t *testing.T) {
|
||||||
|
m := map[string]interface{}{"a": 1, "nilv": nil}
|
||||||
|
if got := Get(m, "a"); got != 1 {
|
||||||
|
t.Errorf("Get: %v", got)
|
||||||
|
}
|
||||||
|
if got := Get(m, "missing"); got != nil {
|
||||||
|
t.Errorf("Get missing: %v", got)
|
||||||
|
}
|
||||||
|
if got := Get(nil, "a"); got != nil {
|
||||||
|
t.Errorf("Get nil map: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetOr(m, "missing", "def"); got != "def" {
|
||||||
|
t.Errorf("GetOr missing: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetOr(m, "nilv", "def"); got != "def" {
|
||||||
|
t.Errorf("GetOr nil value: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetOr(m, "a", "def"); got != 1 {
|
||||||
|
t.Errorf("GetOr present: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetPath(t *testing.T) {
|
||||||
|
cfg := map[string]interface{}{
|
||||||
|
"db": map[string]interface{}{"conn": map[string]interface{}{"host": "h"}},
|
||||||
|
}
|
||||||
|
if got := GetPath(cfg, "db.conn.host"); got != "h" {
|
||||||
|
t.Errorf("GetPath: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetPath(cfg, "db.nope.host"); got != nil {
|
||||||
|
t.Errorf("GetPath missing: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetPathOr(cfg, "db.nope", "dflt"); got != "dflt" {
|
||||||
|
t.Errorf("GetPathOr: %v", got)
|
||||||
|
}
|
||||||
|
if got := GetPathOr(cfg, "db.conn.host", "dflt"); got != "h" {
|
||||||
|
t.Errorf("GetPathOr present: %v", got)
|
||||||
|
}
|
||||||
|
if !HasPath(cfg, "db.conn") || HasPath(cfg, "db.x") || HasPath(nil, "a") {
|
||||||
|
t.Errorf("HasPath mismatch")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSafeIndex(t *testing.T) {
|
||||||
|
s := []string{"a", "b"}
|
||||||
|
if got := SafeIndex(s, 1); got != "b" {
|
||||||
|
t.Errorf("SafeIndex: %v", got)
|
||||||
|
}
|
||||||
|
for _, i := range []int{-1, 2, 99} {
|
||||||
|
if got := SafeIndex(s, i); got != nil {
|
||||||
|
t.Errorf("SafeIndex(%d) must be nil, got %v", i, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := SafeIndex("notslice", 0); got != nil {
|
||||||
|
t.Errorf("non-slice: %v", got)
|
||||||
|
}
|
||||||
|
if got := SafeIndexOr(s, 5, "d"); got != "d" {
|
||||||
|
t.Errorf("SafeIndexOr: %v", got)
|
||||||
|
}
|
||||||
|
if got := SafeIndexOr(s, 0, "d"); got != "a" {
|
||||||
|
t.Errorf("SafeIndexOr present: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHas(t *testing.T) {
|
||||||
|
m := map[string]int{"a": 1}
|
||||||
|
var nilPtr *map[string]int
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
m interface{}
|
||||||
|
key interface{}
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"present", m, "a", true},
|
||||||
|
{"missing", m, "b", false},
|
||||||
|
{"pointer to map", &m, "a", true},
|
||||||
|
{"nil pointer", nilPtr, "a", false},
|
||||||
|
{"non-map", []int{1}, 0, false},
|
||||||
|
{"nil", nil, "a", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := Has(tt.m, tt.key); got != tt.want {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeysValues(t *testing.T) {
|
||||||
|
m := map[string]int{"a": 1, "b": 2}
|
||||||
|
if got := Keys(m); len(got) != 2 {
|
||||||
|
t.Errorf("Keys: %v", got)
|
||||||
|
}
|
||||||
|
if got := Values(m); len(got) != 2 {
|
||||||
|
t.Errorf("Values: %v", got)
|
||||||
|
}
|
||||||
|
if got := Keys(nil); len(got) != 0 {
|
||||||
|
t.Errorf("Keys nil: %v", got)
|
||||||
|
}
|
||||||
|
if got := Values(5); len(got) != 0 {
|
||||||
|
t.Errorf("Values non-map: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMerge(t *testing.T) {
|
||||||
|
m1 := map[string]int{"a": 1, "b": 2}
|
||||||
|
m2 := map[string]int{"b": 3, "c": 4}
|
||||||
|
var nilPtr *map[string]int
|
||||||
|
got := Merge(m1, &m2, nilPtr, nil, 5)
|
||||||
|
want := map[interface{}]interface{}{"a": 1, "b": 3, "c": 4}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
if got := Merge(); len(got) != 0 {
|
||||||
|
t.Errorf("empty merge: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPickOmit(t *testing.T) {
|
||||||
|
m := map[string]int{"a": 1, "b": 2, "c": 3}
|
||||||
|
var nilPtr *map[string]int
|
||||||
|
|
||||||
|
if got := Pick(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"a": 1}) {
|
||||||
|
t.Errorf("Pick: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pick(&m, "b"); len(got) != 1 {
|
||||||
|
t.Errorf("Pick ptr: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pick(nilPtr, "a"); len(got) != 0 {
|
||||||
|
t.Errorf("Pick nil ptr: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pick(5, "a"); len(got) != 0 {
|
||||||
|
t.Errorf("Pick non-map: %v", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := Omit(m, "a", "z"); !reflect.DeepEqual(got, map[interface{}]interface{}{"b": 2, "c": 3}) {
|
||||||
|
t.Errorf("Omit: %v", got)
|
||||||
|
}
|
||||||
|
if got := Omit(&m); len(got) != 3 {
|
||||||
|
t.Errorf("Omit ptr: %v", got)
|
||||||
|
}
|
||||||
|
if got := Omit(nilPtr, "a"); len(got) != 0 {
|
||||||
|
t.Errorf("Omit nil ptr: %v", got)
|
||||||
|
}
|
||||||
|
if got := Omit("x", "a"); len(got) != 0 {
|
||||||
|
t.Errorf("Omit non-map: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSliceContainsIndexOf(t *testing.T) {
|
||||||
|
s := []string{"a", "b", "c"}
|
||||||
|
sp := &s
|
||||||
|
var nilPtr *[]string
|
||||||
|
if !SliceContains(s, "b") || SliceContains(s, "z") {
|
||||||
|
t.Errorf("SliceContains")
|
||||||
|
}
|
||||||
|
if !SliceContains(sp, "c") || !SliceContains([2]int{1, 2}, 2) {
|
||||||
|
t.Errorf("SliceContains ptr/array")
|
||||||
|
}
|
||||||
|
if SliceContains(nilPtr, "a") || SliceContains("str", "s") || SliceContains(nil, 1) {
|
||||||
|
t.Errorf("SliceContains invalid input")
|
||||||
|
}
|
||||||
|
if got := IndexOf(s, "c"); got != 2 {
|
||||||
|
t.Errorf("IndexOf: %d", got)
|
||||||
|
}
|
||||||
|
if got := IndexOf(sp, "a"); got != 0 {
|
||||||
|
t.Errorf("IndexOf ptr: %d", got)
|
||||||
|
}
|
||||||
|
for _, in := range []interface{}{s, nilPtr, "str", nil} {
|
||||||
|
if got := IndexOf(in, "zzz"); got != -1 {
|
||||||
|
t.Errorf("IndexOf miss %v: %d", in, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPluck(t *testing.T) {
|
||||||
|
items := []*accessItem{{"a", 1}, nil, {"c", 3}}
|
||||||
|
got := Pluck(items, "Name")
|
||||||
|
if !reflect.DeepEqual(got, []interface{}{"a", nil, "c"}) {
|
||||||
|
t.Errorf("struct ptrs: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pluck([]accessItem{{"a", 1}}, "Missing"); !reflect.DeepEqual(got, []interface{}{nil}) {
|
||||||
|
t.Errorf("missing field: %v", got)
|
||||||
|
}
|
||||||
|
maps := []map[string]int{{"k": 1}, {"x": 2}}
|
||||||
|
if got := Pluck(maps, "k"); !reflect.DeepEqual(got, []interface{}{1, nil}) {
|
||||||
|
t.Errorf("maps: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pluck([]int{1, 2}, "k"); !reflect.DeepEqual(got, []interface{}{nil, nil}) {
|
||||||
|
t.Errorf("scalars: %v", got)
|
||||||
|
}
|
||||||
|
var nilPtr *[]accessItem
|
||||||
|
if got := Pluck(nilPtr, "Name"); len(got) != 0 {
|
||||||
|
t.Errorf("nil ptr: %v", got)
|
||||||
|
}
|
||||||
|
if got := Pluck("str", "Name"); len(got) != 0 {
|
||||||
|
t.Errorf("non-slice: %v", got)
|
||||||
|
}
|
||||||
|
s := []accessItem{{"z", 9}}
|
||||||
|
if got := Pluck(&s, "ID"); !reflect.DeepEqual(got, []interface{}{9}) {
|
||||||
|
t.Errorf("ptr to slice: %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCaseConversions(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
in, camel, pascal, snake, kebab string
|
||||||
|
}{
|
||||||
|
{"", "", "", "", ""},
|
||||||
|
{"user_name", "userName", "UserName", "user_name", "user-name"},
|
||||||
|
{"http_request", "httpRequest", "HTTPRequest", "http_request", "http-request"},
|
||||||
|
{"user_id", "userID", "UserID", "user_id", "user-id"},
|
||||||
|
{"UserName", "username", "UserName", "user_name", "user-name"},
|
||||||
|
{"HTTPRequest", "httprequest", "HTTPRequest", "http_request", "http-request"},
|
||||||
|
{"userID", "userid", "UserID", "user_id", "user-id"},
|
||||||
|
{"name", "name", "Name", "name", "name"},
|
||||||
|
{"ÜberUser", "überuser", "ÜberUser", "über_user", "über-user"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.in, func(t *testing.T) {
|
||||||
|
if got := ToCamelCase(tt.in); got != tt.camel {
|
||||||
|
t.Errorf("ToCamelCase = %q, want %q", got, tt.camel)
|
||||||
|
}
|
||||||
|
if got := ToPascalCase(tt.in); got != tt.pascal {
|
||||||
|
t.Errorf("ToPascalCase = %q, want %q", got, tt.pascal)
|
||||||
|
}
|
||||||
|
if got := ToSnakeCase(tt.in); got != tt.snake {
|
||||||
|
t.Errorf("ToSnakeCase = %q, want %q", got, tt.snake)
|
||||||
|
}
|
||||||
|
if got := ToKebabCase(tt.in); got != tt.kebab {
|
||||||
|
t.Errorf("ToKebabCase = %q, want %q", got, tt.kebab)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPluralize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"user", "users"},
|
||||||
|
{"person", "people"},
|
||||||
|
{"Person", "people"},
|
||||||
|
{"status", "statuses"},
|
||||||
|
{"cats", "cats"},
|
||||||
|
{"bus", "buses"},
|
||||||
|
{"dress", "dresses"},
|
||||||
|
{"box", "boxes"},
|
||||||
|
{"quiz", "quizes"},
|
||||||
|
{"church", "churches"},
|
||||||
|
{"dish", "dishes"},
|
||||||
|
{"category", "categories"},
|
||||||
|
{"day", "days"},
|
||||||
|
{"leaf", "leaves"},
|
||||||
|
{"knife", "knives"},
|
||||||
|
{"hero", "heroes"},
|
||||||
|
{"video", "videos"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Pluralize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("Pluralize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSingularize(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"users", "user"},
|
||||||
|
{"people", "person"},
|
||||||
|
{"Children", "child"},
|
||||||
|
{"categories", "category"},
|
||||||
|
{"ies", "ie"},
|
||||||
|
{"leaves", "leaf"},
|
||||||
|
{"buses", "bus"},
|
||||||
|
{"boxes", "box"},
|
||||||
|
{"churches", "church"},
|
||||||
|
{"dishes", "dish"},
|
||||||
|
{"dress", "dress"},
|
||||||
|
{"user", "user"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := Singularize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("Singularize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPlainStringWrappers(t *testing.T) {
|
||||||
|
if ToUpper("aB") != "AB" || ToLower("aB") != "ab" {
|
||||||
|
t.Error("case")
|
||||||
|
}
|
||||||
|
if Title("hello world") != "Hello World" || Title("") != "" {
|
||||||
|
t.Errorf("Title: %q", Title("hello world"))
|
||||||
|
}
|
||||||
|
if Trim(" a \n") != "a" {
|
||||||
|
t.Error("Trim")
|
||||||
|
}
|
||||||
|
if TrimPrefix("foobar", "foo") != "bar" || TrimPrefix("bar", "foo") != "bar" {
|
||||||
|
t.Error("TrimPrefix")
|
||||||
|
}
|
||||||
|
if TrimSuffix("foobar", "bar") != "foo" || TrimSuffix("foo", "bar") != "foo" {
|
||||||
|
t.Error("TrimSuffix")
|
||||||
|
}
|
||||||
|
if Replace("aaa", "a", "b", 2) != "bba" || Replace("aaa", "a", "b", -1) != "bbb" {
|
||||||
|
t.Error("Replace")
|
||||||
|
}
|
||||||
|
if !StringContains("abc", "b") || StringContains("abc", "z") {
|
||||||
|
t.Error("StringContains")
|
||||||
|
}
|
||||||
|
if !HasPrefix("abc", "ab") || HasPrefix("abc", "bc") {
|
||||||
|
t.Error("HasPrefix")
|
||||||
|
}
|
||||||
|
if !HasSuffix("abc", "bc") || HasSuffix("abc", "ab") {
|
||||||
|
t.Error("HasSuffix")
|
||||||
|
}
|
||||||
|
if got := Split("a,b", ","); !reflect.DeepEqual(got, []string{"a", "b"}) {
|
||||||
|
t.Errorf("Split: %v", got)
|
||||||
|
}
|
||||||
|
if Join([]string{"a", "b"}, "-") != "a-b" || Join(nil, "-") != "" {
|
||||||
|
t.Error("Join")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCapitalizeAndIsVowel(t *testing.T) {
|
||||||
|
tests := []struct{ in, want string }{
|
||||||
|
{"", ""},
|
||||||
|
{"id", "ID"},
|
||||||
|
{"Uuid", "UUID"},
|
||||||
|
{"http", "HTTP"},
|
||||||
|
{"name", "Name"},
|
||||||
|
{"élan", "Élan"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := capitalize(tt.in); got != tt.want {
|
||||||
|
t.Errorf("capitalize(%q) = %q, want %q", tt.in, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, c := range []byte("aeiouAEIOU") {
|
||||||
|
if !isVowel(c) {
|
||||||
|
t.Errorf("%c should be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, c := range []byte("bcxyz") {
|
||||||
|
if isVowel(c) {
|
||||||
|
t.Errorf("%c should not be vowel", c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func sampleDB() (*models.Database, *models.Schema, *models.Table) {
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
table := models.InitTable("users", "public")
|
||||||
|
col := models.InitColumn("id", "users", "public")
|
||||||
|
col.Type = "integer"
|
||||||
|
col.IsPrimaryKey = true
|
||||||
|
table.Columns["id"] = col
|
||||||
|
schema.Tables = append(schema.Tables, table)
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
return db, schema, table
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTemplateDataConstructors(t *testing.T) {
|
||||||
|
db, schema, table := sampleDB()
|
||||||
|
meta := map[string]interface{}{"k": "v"}
|
||||||
|
|
||||||
|
dd := NewDatabaseData(db, meta)
|
||||||
|
if dd.Database != db || dd.ParentDatabase != db || dd.Summary == nil || len(dd.FlatColumns) != 1 || len(dd.FlatTables) != 1 || dd.Metadata["k"] != "v" {
|
||||||
|
t.Errorf("database data: %+v", dd)
|
||||||
|
}
|
||||||
|
if dd.Name() != "shop" {
|
||||||
|
t.Errorf("name: %q", dd.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
sd := NewSchemaData(schema, meta)
|
||||||
|
if sd.Schema != schema || sd.ParentDatabase == nil || sd.ParentDatabase.Name != "public" || len(sd.FlatColumns) != 1 {
|
||||||
|
t.Errorf("schema data: %+v", sd)
|
||||||
|
}
|
||||||
|
if sd.Name() != "public" {
|
||||||
|
t.Errorf("name: %q", sd.Name())
|
||||||
|
}
|
||||||
|
|
||||||
|
td := NewTableData(table, schema, db, meta)
|
||||||
|
if td.Table != table || td.ParentSchema != schema || td.ParentDatabase != db || td.Name() != "users" {
|
||||||
|
t.Errorf("table data: %+v", td)
|
||||||
|
}
|
||||||
|
|
||||||
|
dom := &models.Domain{Name: "billing"}
|
||||||
|
dmd := NewDomainData(dom, db, meta)
|
||||||
|
if dmd.Domain != dom || dmd.ParentDatabase != db || dmd.Name() != "billing" {
|
||||||
|
t.Errorf("domain data: %+v", dmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
sc := &models.Script{Name: "seed"}
|
||||||
|
scd := NewScriptData(sc, schema, db, meta)
|
||||||
|
if scd.Script != sc || scd.ParentSchema != schema || scd.Name() != "seed" {
|
||||||
|
t.Errorf("script data: %+v", scd)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := (&TemplateData{}).Name(); got != "output" {
|
||||||
|
t.Errorf("empty name: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTypeMappersDelegate(t *testing.T) {
|
||||||
|
if got := SQLToGo("integer", false); got == "" {
|
||||||
|
t.Error("SQLToGo")
|
||||||
|
}
|
||||||
|
if got := SQLToTypeScript("integer", false); got == "" {
|
||||||
|
t.Error("SQLToTypeScript")
|
||||||
|
}
|
||||||
|
if got := SQLToJava("integer", false); got == "" {
|
||||||
|
t.Error("SQLToJava")
|
||||||
|
}
|
||||||
|
if got := SQLToPython("integer"); got == "" {
|
||||||
|
t.Error("SQLToPython")
|
||||||
|
}
|
||||||
|
if got := SQLToRust("integer", false); got == "" {
|
||||||
|
t.Error("SQLToRust")
|
||||||
|
}
|
||||||
|
if got := SQLToCSharp("integer", false); got == "" {
|
||||||
|
t.Error("SQLToCSharp")
|
||||||
|
}
|
||||||
|
if got := SQLToPhp("integer", false); got == "" {
|
||||||
|
t.Error("SQLToPhp")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,219 @@
|
|||||||
|
package template
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func writeTemplateFile(t *testing.T, body string) string {
|
||||||
|
t.Helper()
|
||||||
|
p := filepath.Join(t.TempDir(), "t.tmpl")
|
||||||
|
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func modeDB() *models.Database {
|
||||||
|
db := models.InitDatabase("shop")
|
||||||
|
for _, sn := range []string{"a", "b"} {
|
||||||
|
s := models.InitSchema(sn)
|
||||||
|
for _, tn := range []string{"t1", "t2"} {
|
||||||
|
s.Tables = append(s.Tables, models.InitTable(tn, sn))
|
||||||
|
}
|
||||||
|
s.Scripts = append(s.Scripts, &models.Script{Name: "seed_" + sn})
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
}
|
||||||
|
db.Domains = append(db.Domains, &models.Domain{Name: "billing"})
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestWriter(t *testing.T, body, mode, pattern, out string) (*Writer, error) {
|
||||||
|
t.Helper()
|
||||||
|
meta := map[string]interface{}{"template_path": writeTemplateFile(t, body)}
|
||||||
|
if mode != "" {
|
||||||
|
meta["mode"] = mode
|
||||||
|
}
|
||||||
|
if pattern != "" {
|
||||||
|
meta["filename_pattern"] = pattern
|
||||||
|
}
|
||||||
|
return NewWriter(&writers.WriterOptions{OutputPath: out, Metadata: meta})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewWriterErrors(t *testing.T) {
|
||||||
|
if _, err := NewWriter(&writers.WriterOptions{}); err == nil {
|
||||||
|
t.Error("expected error for missing template path")
|
||||||
|
}
|
||||||
|
_, err := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{"template_path": "/no/such/file"}})
|
||||||
|
var te *TemplateError
|
||||||
|
if !errors.As(err, &te) || te.Phase != "load" {
|
||||||
|
t.Errorf("load error: %v", err)
|
||||||
|
}
|
||||||
|
_, err = newTestWriter(t, "{{ .Unclosed ", "", "", "")
|
||||||
|
if !errors.As(err, &te) || te.Phase != "parse" {
|
||||||
|
t.Errorf("parse error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterModes(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name, mode, body, pattern string
|
||||||
|
wantFiles []string
|
||||||
|
}{
|
||||||
|
{"database", "database", "{{.Database.Name}}", "", []string{"out.txt"}},
|
||||||
|
{"schema", "schema", "{{.Schema.Name}}", "{{.Name}}.txt", []string{"a.txt", "b.txt"}},
|
||||||
|
{"table", "table", "{{.Table.Name}}", "{{.ParentSchema.Name}}_{{.Name}}.txt", []string{"a_t1.txt", "a_t2.txt", "b_t1.txt", "b_t2.txt"}},
|
||||||
|
{"script", "script", "{{.Script.Name}}", "{{.Name}}.sql", []string{"seed_a.sql", "seed_b.sql"}},
|
||||||
|
{"domain", "domain", "{{.Domain.Name}}", "{{.Name}}.md", []string{"billing.md"}},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
outDir := t.TempDir()
|
||||||
|
out := outDir
|
||||||
|
if tt.mode == "database" {
|
||||||
|
out = filepath.Join(outDir, "out.txt")
|
||||||
|
}
|
||||||
|
w, err := newTestWriter(t, tt.body, tt.mode, tt.pattern, out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := w.WriteDatabase(modeDB()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, f := range tt.wantFiles {
|
||||||
|
if _, err := os.Stat(filepath.Join(outDir, f)); err != nil {
|
||||||
|
t.Errorf("missing %s: %v", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
entries, _ := os.ReadDir(outDir)
|
||||||
|
if len(entries) != len(tt.wantFiles) {
|
||||||
|
t.Errorf("got %d files, want %d", len(entries), len(tt.wantFiles))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterDatabaseModeContent(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "sub", "dir", "o.txt")
|
||||||
|
w, err := newTestWriter(t, "{{.Database.Name}}:{{len .Database.Schemas}}", "", "", out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := w.WriteDatabase(modeDB()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
data, err := os.ReadFile(out)
|
||||||
|
if err != nil || string(data) != "shop:2" {
|
||||||
|
t.Errorf("content %q err %v", data, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterUnknownMode(t *testing.T) {
|
||||||
|
w, err := newTestWriter(t, "x", "bogus", "", "")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := w.WriteDatabase(modeDB()); err == nil || !strings.Contains(err.Error(), "unknown entrypoint mode") {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterExecuteErrors(t *testing.T) {
|
||||||
|
// Execution failure: field does not exist on TemplateData.
|
||||||
|
for _, mode := range []string{"database", "schema", "table", "script", "domain"} {
|
||||||
|
t.Run(mode, func(t *testing.T) {
|
||||||
|
w, err := newTestWriter(t, "{{.NoSuchField}}", mode, "", t.TempDir())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
err = w.WriteDatabase(modeDB())
|
||||||
|
var te *TemplateError
|
||||||
|
if !errors.As(err, &te) || te.Phase != "execute" {
|
||||||
|
t.Errorf("got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterBadFilenamePattern(t *testing.T) {
|
||||||
|
for _, pattern := range []string{"{{.Unclosed", "{{.NoSuchField}}"} {
|
||||||
|
for _, mode := range []string{"schema", "table", "script", "domain"} {
|
||||||
|
w, err := newTestWriter(t, "x", mode, pattern, t.TempDir())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := w.WriteDatabase(modeDB()); err == nil {
|
||||||
|
t.Errorf("mode %s pattern %q: expected error", mode, pattern)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterWriteOutputFailure(t *testing.T) {
|
||||||
|
// Output path whose parent is a regular file cannot be created.
|
||||||
|
blocker := filepath.Join(t.TempDir(), "file")
|
||||||
|
if err := os.WriteFile(blocker, nil, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
w, err := newTestWriter(t, "x", "database", "", filepath.Join(blocker, "child", "o.txt"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := w.WriteDatabase(modeDB()); err == nil {
|
||||||
|
t.Error("expected write failure")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterGenerateFilenameOutputPathForms(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
data := NewTableData(models.InitTable("users", "public"), nil, nil, nil)
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name, out, want string
|
||||||
|
}{
|
||||||
|
{"no output path", "", "users.txt"},
|
||||||
|
{"existing dir", dir, filepath.Join(dir, "users.txt")},
|
||||||
|
{"trailing separator", filepath.Join(dir, "new") + string(filepath.Separator), filepath.Join(dir, "new", "users.txt")},
|
||||||
|
{"file path uses its dir", filepath.Join(dir, "x.out"), filepath.Join(dir, "users.txt")},
|
||||||
|
{"bare file name", "x.out", "users.txt"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
w, err := newTestWriter(t, "x", "table", "{{.Name}}.txt", tt.out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got, err := w.generateFilename(data)
|
||||||
|
if err != nil || got != tt.want {
|
||||||
|
t.Errorf("got %q err %v, want %q", got, err, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterWriteSchemaAndTable(t *testing.T) {
|
||||||
|
out := filepath.Join(t.TempDir(), "o.txt")
|
||||||
|
w, err := newTestWriter(t, "{{range .Database.Schemas}}{{.Name}}:{{len .Tables}};{{end}}", "", "", out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db := modeDB()
|
||||||
|
if err := w.WriteSchema(db.Schemas[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if data, _ := os.ReadFile(out); string(data) != "a:2;" {
|
||||||
|
t.Errorf("WriteSchema: %q", data)
|
||||||
|
}
|
||||||
|
if err := w.WriteTable(db.Schemas[1].Tables[0]); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if data, _ := os.ReadFile(out); string(data) != "b:1;" {
|
||||||
|
t.Errorf("WriteTable: %q", data)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package writers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ParseTypeMappings parses "sqltype=gotype" entries (as given to --type-map)
|
||||||
|
// into a map keyed by the canonical lower-case SQL base type, so that
|
||||||
|
// "VARCHAR", "character varying" and "varchar(50)" all address one entry.
|
||||||
|
// It returns nil for empty input.
|
||||||
|
func ParseTypeMappings(entries []string) (map[string]string, error) {
|
||||||
|
if len(entries) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
out := make(map[string]string, len(entries))
|
||||||
|
for _, entry := range entries {
|
||||||
|
sqlType, goType, ok := strings.Cut(entry, "=")
|
||||||
|
sqlType, goType = strings.TrimSpace(sqlType), strings.TrimSpace(goType)
|
||||||
|
if !ok || sqlType == "" || goType == "" {
|
||||||
|
return nil, fmt.Errorf("invalid type mapping %q: expected sqltype=gotype", entry)
|
||||||
|
}
|
||||||
|
out[pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))] = goType
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LookupTypeMapping returns the user-configured override for baseType, which
|
||||||
|
// the caller must already have canonicalized.
|
||||||
|
func LookupTypeMapping(mappings map[string]string, baseType string) (string, bool) {
|
||||||
|
goType, ok := mappings[baseType]
|
||||||
|
return goType, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyTypeMapping wraps an overridden Go type for nullability: NOT NULL uses
|
||||||
|
// the type verbatim; nullable columns get a pointer prefix unless the type is
|
||||||
|
// already a pointer, slice, map or interface.
|
||||||
|
func ApplyTypeMapping(goType string, notNull bool) string {
|
||||||
|
if notNull || strings.HasPrefix(goType, "*") || strings.HasPrefix(goType, "[]") ||
|
||||||
|
strings.HasPrefix(goType, "map[") || goType == "any" || goType == "interface{}" {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
return "*" + goType
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user