Compare commits

..
28 Commits
Author SHA1 Message Date
warkanum bed80b046b Merge pull request 'Fix/writers readers determinism and tests' (#55) from fix/writers-readers-determinism-and-tests into master
Reviewed-on: #55
2026-10-03 19:34:59 +00:00
warkanum 495a21b67b test: expand coverage across readers, writers, cmd, ui, diff and merge
Implements tests/_plans and previously deferred packages; updates plan
README with new coverage numbers.
2026-10-03 21:33:59 +02:00
warkanum a32647ee16 fix: writer/reader correctness and determinism issues
- merge: cloneTable keeps relationships; skip-tables applies to new schemas
- diff: detect schema description/owner changes
- prisma: reader no longer turns relation fields into columns, enum
  detection via declared names; writer type mapping is ordered
- typeorm writer: keep explicit SQL types that cannot be inferred
- drizzle: enum columns call the enum constant; reader resolves them
- mysql/mssql/sqlite writers: honour OutputPath, deterministic column
  and constraint order; mssql live execute covers full schema
- template: ToYAML recovers from panics; Merge nil-pointer loop
- regenerate drizzle fixtures
2026-10-03 21:33:59 +02:00
warkanum 08e1417393 Merge pull request 'fix(ui): resolve lint issues in file browser' (#54) from fix/lint-filebrowser into master
Reviewed-on: #54
2026-10-03 18:17:30 +00:00
warkanum 70282fff73 fix(mysql): resolve lint issues in reader and writer 2026-10-03 20:17:16 +02:00
warkanum 43265dac0f fix(ui): resolve lint issues in file browser 2026-10-03 20:16:20 +02:00
warkanum 66b90ca54b Merge pull request 'feat(cli): add batch command for converting multiple inputs (#38)' (#49) from issue-38-batch-processing into master
Reviewed-on: #49
2026-10-03 18:14:30 +00:00
warkanum 47108809aa Merge remote-tracking branch 'origin/master' into issue-38-batch-processing
# Conflicts:
#	README.md
2026-10-03 20:13:22 +02:00
warkanum 720476fd6e Merge pull request 'docs(ui): plan TUI mouse support' (#53) from issue-46-mouse-support-plan into master
Reviewed-on: #53
2026-10-03 18:12:04 +00:00
warkanum 572d03fe42 Merge pull request 'feat: add MySQL reader and writer' (#52) from issue-34-mysql-driver into master
Reviewed-on: #52
2026-10-03 18:11:50 +00:00
warkanum f1b9079b2d Merge pull request 'test(sqltypes): cover scalar conversions and constructors' (#51) from issue-45-test-coverage into master
Reviewed-on: #51
2026-10-03 18:11:40 +00:00
warkanum bb671c3680 Merge pull request 'feat(ui): file browser and connection string builder dialogs' (#50) from issue-44-tui-input-dialogs into master
Reviewed-on: #50
2026-10-03 18:11:28 +00:00
warkanum bc8284db25 Merge pull request 'feat(cli): watch mode for convert (#39)' (#48) from issue-39-watch-mode into master
Reviewed-on: #48
2026-10-03 18:11:10 +00:00
warkanum 29e747393d Merge pull request 'feat: custom SQL-to-Go type mapping (--type-map) for bun/gorm' (#47) from issue-36-type-mapping into master
Reviewed-on: #47
2026-10-03 18:10:59 +00:00
SG Command df980a3434 docs(ui): plan TUI mouse support 2026-10-03 13:24:46 +02:00
SG Command 778379538b fix: make MySQL writer execute generated DDL 2026-10-03 12:42:07 +02:00
SG Command 7a9219b6e3 feat: add MySQL reader and writer 2026-10-03 12:36:39 +02:00
SG Command fd9c37cd25 test(sqltypes): cover scalar conversions and constructors 2026-10-03 12:30:26 +02:00
SG CommandandClaude Sonnet 5.5 0235a28add feat(ui): file browser and connection string builder dialogs (#44)
Enter on File Path inputs opens a file browser (load/save, extension
filter, hidden toggle, overwrite confirm). Enter on Connection String
inputs opens a builder for PostgreSQL, MSSQL and SQLite with masked
password/preview, parsing and optional connection test.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-10-03 11:35:12 +02:00
SG CommandandClaude Sonnet 5.5 d961536186 feat(cli): add batch command to convert many inputs in one run
Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-10-03 11:28:57 +02:00
Hermes AgentandClaude Sonnet 5.5 d36806047b feat(cli): add --watch mode to convert (#39)
Poll source files/directories and regenerate output on change.
Output path is excluded from watching; errors don't stop the loop.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-10-03 10:29:42 +02:00
Hermes AgentandClaude Sonnet 5.5 948419ffd3 feat: add --type-map to override SQL-to-Go types in bun and gorm writers
Adds WriterOptions.TypeMappings and a repeatable --type-map sqltype=gotype
flag. Defaults are unchanged when no mapping is given. Closes #36 (bun/gorm).

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
2026-10-03 10:28:58 +02:00
warkanum b38f53c603 docs(tests): add test coverage plans 2026-10-03 10:01:48 +02:00
warkanum ccba53c494 test: add podman/docker dbtest tool for postgres, mssql and mysql 2026-10-03 10:01:48 +02:00
sgcommand 53327b9a5a feat(ui): TUI indexes, views, sequences, scripts, domain/table assignment (#40) (#43)
Co-authored-by: SG Command <sgcommand@warky.dev>
2026-10-03 07:27:25 +00:00
sgcommand 734b14d48d docs: add usage examples for each format combination (#42)
Co-authored-by: SG Command <sgcommand@warky.dev>
2026-10-03 07:27:03 +00:00
sgcommand 938f0ed51f feat(cli): add --dry-run to convert, merge and split (#41)
Co-authored-by: SG Command <sgcommand@warky.dev>
2026-10-03 07:26:54 +00:00
warkanum 6e2e7eb19e feat(release): add rerelease target to move latest tag 2026-10-02 23:41:42 +02:00
168 changed files with 25463 additions and 180 deletions
+9 -1
View File
@@ -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}'
+39
View File
@@ -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)
+214
View File
@@ -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
}
+143
View File
@@ -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
View File
@@ -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))
+312
View File
@@ -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)
}
}
+158
View File
@@ -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)
}
}
+35
View File
@@ -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)
}
}
+1
View File
@@ -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,
+10
View File
@@ -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:, …)")
+86
View File
@@ -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)
}
}
+23
View File
@@ -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)
+152
View File
@@ -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
}
+92
View File
@@ -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")
}
+103
View File
@@ -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).
+212
View File
@@ -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.
+2
View File
@@ -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
+4
View File
@@ -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=
+17
View File
@@ -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) {
+337
View File
@@ -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)
}
}
+63
View File
@@ -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")
}
}
+1
View File
@@ -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"`
+262
View File
@@ -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)
}
}
}
+19
View File
@@ -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
} }
+277
View File
@@ -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")
}
}
+58
View File
@@ -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")
}
}
+1
View File
@@ -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
+232
View File
@@ -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")
}
}
+170
View File
@@ -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")
}
}
+249
View File
@@ -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)
}
})
}
}
+105
View File
@@ -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)
}
}
+27
View File
@@ -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
+114
View File
@@ -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")
}
}
+152
View 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)
}
})
}
}
+254
View File
@@ -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
}
+22
View File
@@ -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")
}
}
+113
View File
@@ -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)
}
}
}
+17 -13
View File
@@ -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
+348
View File
@@ -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)
}
}
}
+375
View File
@@ -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)
}
}
+157
View File
@@ -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)
}
}
+38
View File
@@ -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)
}
}
+186
View File
@@ -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
}
+62
View File
@@ -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)
}
+210
View File
@@ -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
})
}
+143
View File
@@ -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")
}
}
+198
View 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)
}
}
+4
View File
@@ -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)
}) })
+134
View File
@@ -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, ""
}
+363
View File
@@ -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)
}
+124
View File
@@ -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")
}
}
+294
View File
@@ -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")
}
}
+63
View File
@@ -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)
}
}
+12
View File
@@ -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
}) })
+12
View File
@@ -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()
}). }).
+263
View File
@@ -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
}
+136
View File
@@ -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)
}
}
+476
View File
@@ -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)
}
+113
View File
@@ -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])
})
}
}
+12
View File
@@ -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
+83
View File
@@ -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)
}
}
}
+25
View File
@@ -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")
}
}
+2
View File
@@ -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 {
+2 -2
View File
@@ -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)
} }
+300
View File
@@ -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)
}
}
+83
View File
@@ -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)
}
}
}
+25
View File
@@ -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")
}
}
+2
View File
@@ -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
View File
@@ -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
} }
+205
View File
@@ -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)
}
}
+209
View File
@@ -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, ", ")
}
+159
View File
@@ -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)
}
}
+44
View File
@@ -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)
}
}
}
+110
View File
@@ -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)
}
})
}
}
+243
View File
@@ -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)
}
}
+288
View File
@@ -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)
}
}
+34
View File
@@ -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)
}
}
}
}
+29 -27
View File
@@ -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
} }
} }
+259
View File
@@ -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)
}
}
+223
View File
@@ -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)
}
}
+36 -17
View File
@@ -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)
} }
+250
View File
@@ -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)
}
}
+48
View File
@@ -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")
}
}
+168
View File
@@ -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)
}
}
}
+7 -1
View File
@@ -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)
+118
View File
@@ -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)
}
}
}
+75
View File
@@ -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)
}
})
}
}
+142
View File
@@ -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)
}
}
+2 -5
View File
@@ -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()
} }
+216
View File
@@ -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)
}
}
+151
View File
@@ -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")
}
}
+219
View File
@@ -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)
}
}
+46
View File
@@ -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