Compare commits

...
57 Commits
Author SHA1 Message Date
warkanum 2796ebe381 chore(release): update package version to 1.0.86
Release / test (push) Failing after 2m29s
Release / release (push) Skipped
Release / pkg-aur (push) Skipped
Release / pkg-deb (push) Skipped
Release / pkg-rpm (push) Skipped
2026-10-03 23:23:20 +02:00
warkanum 9447d18323 docs: README covers mysql, update command and Windows installer 2026-10-03 23:22:03 +02:00
warkanum 9da69f4234 docs(tests): update coverage plans status 2026-10-03 23:21:26 +02:00
warkanum 2c304dc363 Merge pull request 'feat: NSIS Windows installer and update check (#32)' (#56) from issue-32-nsis-installer into master
Reviewed-on: #56
2026-10-03 21:19:58 +00:00
warkanum 02a394c99d feat: NSIS Windows installer and update check (#32)
- windows/installer.nsi: installs to Program Files, manages system PATH,
  registers uninstaller
- relspec update [--check] [--yes]: compares against the latest release and
  prompts; on Windows downloads and runs the installer
- release workflow builds and uploads relspec-setup-windows-amd64.exe
- make installer-windows
2026-10-03 21:44:08 +02:00
warkanum b49859537b style: reformat test tables 2026-10-03 21:41:54 +02:00
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
warkanum 82e86cdf39 fix(release): run lint and format check before version release
Release / test (push) Successful in 2m7s
Release / release (push) Successful in 1m49s
Release / pkg-aur (push) Successful in 50s
Release / pkg-deb (push) Successful in 1m30s
Release / pkg-rpm (push) Successful in 1m48s
2026-10-02 23:38:28 +02:00
warkanum b7ae00ab52 fix(test): simplify parameter list in identity column builder 2026-10-02 23:35:56 +02:00
warkanum 8f8664357e chore(release): update package version to 1.0.85
Release / test (push) Failing after 2m1s
Release / release (push) Skipped
Release / pkg-aur (push) Skipped
Release / pkg-deb (push) Skipped
Release / pkg-rpm (push) Skipped
2026-10-02 23:32:07 +02:00
warkanum 7bd61a21b5 fix: resolve gofumpt and staticcheck lint errors 2026-10-02 23:31:41 +02:00
warkanum 6f6b9834ca docs(CONTRIBUTING): update setup instructions and add make targets 2026-10-02 23:29:25 +02:00
warkanum 9cc10715e3 feat: generated and identity columns for mssql, sqlite, drizzle, typeorm, bun and gorm readers
- mssql: write computed columns (AS (expr) PERSISTED) and identity; read computed definition and identity
- sqlite: write GENERATED ALWAYS AS (expr) STORED; read generated columns via table_xinfo and parse the expression
- drizzle: generatedAlwaysAs / generated(Always|ByDefault)AsIdentity, read and write
- typeorm: asExpression/generatedType and identity decorators, read and write; decorator scan is now quote-aware
- bun, gorm readers: read generated and identity markers
- readmes updated
2026-10-02 23:04:04 +02:00
warkanum bbab5ce936 feat: mark generated and identity columns across writers and DBML
- bun: scanonly+generated for GENERATED columns; scanonly+identity for ALWAYS identity (PK gets identity only)
- gorm: <-:false+generated / identity, same primary key rule
- dbml: write GENERATED/IDENTITY column notes and read them back into the model
- pgsql: emit generation/identity clauses in migration create-table and add-column, no DEFAULT
- readmes updated
2026-10-02 22:58:00 +02:00
warkanum 275424c605 feat(pgsql): diff against the live database by default for direct output
Direct pgsql output (job output.conn_env, merge --output-conn) now reads the live
schema and executes only the differences. full_ddl: true (job option) or --full-ddl
(merge) restores the full idempotent DDL; file output is unchanged.

- compare PKs by columns, skip constraint-backed indexes, normalize index method,
  FK actions, serial/numeric types, default literals/casts and truncated names
- diff table and column comments instead of re-emitting them
- remove leftover ZZDUMP debug code
2026-10-02 22:52:59 +02:00
warkanum b6c0cd3d1b fix(pgsql-reader): pair composite foreign key columns by position
key_column_usage and constraint_column_usage were joined on constraint name only,
so an N-column foreign key returned N*N column pairs. Read from pg_constraint with
unnest(conkey, confkey) instead.
2026-10-02 22:52:53 +02:00
warkanum d7d1d99ebc fix(pgsql): tie PK sequence to nextval default, setval past data, keep serial defaults
- derive the primary key sequence from the column's nextval() default instead of an unused identity_<table>_<pk> sequence
- setval after table creation (full and diff paths), forward-only, MAX+1 with is_called=false
- do not DROP DEFAULT on serial/bigserial columns without a model default
2026-10-02 22:37:37 +02:00
warkanum 43849324ce chore(release): update package version to 1.0.84
Release / test (push) Successful in 1m25s
Release / release (push) Successful in 2m7s
Release / pkg-aur (push) Successful in 55s
Release / pkg-deb (push) Successful in 1m42s
Release / pkg-rpm (push) Successful in 1m49s
2026-09-24 15:41:46 +02:00
warkanum e1df3b7cee fix(writers): drop version stamps from generated output
Generated files carried the RelSpec version and build date, so every
rebuild produced a diff even when the schema was unchanged.

- bun, gorm, drizzle: remove the GeneratedBy template field and line
- pgsql, mssql, sqlite: emit a constant "-- Generated by RelSpec"
- bun: replace the version header test with a version-free assertion

pkg/buildinfo is unchanged and still backs the CLI version.
2026-09-24 15:40:37 +02:00
warkanum f7b5d5f054 chore(release): update package version to 1.0.83
Release / pkg-deb (push) Successful in 1m47s
Release / test (push) Successful in 59s
Release / release (push) Successful in 2m0s
Release / pkg-aur (push) Successful in 54s
Release / pkg-rpm (push) Successful in 1m56s
2026-09-23 18:52:48 +02:00
warkanum 99fc4b0944 Merge pull request 'Fix/dbml commented refs duplicate indexes' (#31) from fix/dbml-commented-refs-duplicate-indexes into master
Reviewed-on: #31
2026-09-23 16:48:54 +00:00
warkanum 278d488363 chore(lint): use strings.EqualFold in jobs output format check 2026-09-23 18:48:07 +02:00
warkanum 2f69205aa0 fix(dbml): resolve commented cross-file // Ref: lines
Commented refs are collected per file and resolved against the combined
model after all inputs are loaded (directory, --from-list, merge, jobs).
Matched refs become FKs and relationships; duplicates of existing FKs are
skipped; missing targets are skipped with a warning; column type
mismatches warn.

Also keep reused index names within a DBML table instead of overwriting,
give a second FK to the same table a distinct relationship name, and make
the pgsql writer match relationships to FKs by name first.
2026-09-23 18:48:07 +02:00
warkanum b91985c493 feat(inspector): fail on duplicate index names per schema
PostgreSQL index names share the schema-wide relation namespace, so a
reused name makes CREATE INDEX IF NOT EXISTS silently skip the duplicate.
Add duplicate_index_name rule (enforce by default) covering indexes and
PK/unique constraint backing indexes.
2026-09-23 18:48:07 +02:00
warkanum 80a3453233 Merge pull request 'fix(jobs): preserve writer and template output options' (#30) from issue-29-job-manifest-options into master
Reviewed-on: #30
Reviewed-by: Warky <2+warkanum@noreply@warky.dev>
2026-09-21 04:00:50 +00:00
SG Command ade71b75fa fix(jobs): preserve writer and template output options 2026-09-20 23:16:41 +02:00
warkanum 799b7feb53 ci(rpm): update RPM builder to avoid transient mirrors 2026-09-20 21:41:52 +02:00
warkanum eb40201b59 chore(release): update package version to 1.0.82
Release / release (push) Successful in 1m50s
Release / pkg-aur (push) Successful in 1m0s
Release / pkg-deb (push) Successful in 1m17s
Release / test (push) Successful in 53s
Release / pkg-rpm (push) Failing after 16s
2026-09-20 20:37:41 +02:00
warkanum 1461f8f69f chore(release): update version to 1.0.81 and metadata 2026-09-20 20:37:09 +02:00
warkanum ea70e19a46 feat(jobs): expand environment variables in paths 2026-09-20 20:08:20 +02:00
259 changed files with 30026 additions and 540 deletions
+19 -3
View File
@@ -63,6 +63,8 @@ jobs:
- name: Build release binaries
run: |
VERSION="${{ github.event.inputs.tag || github.ref_name }}"
BUILD_DATE="$(date -u +"%Y-%m-%d %H:%M:%S UTC")"
LDFLAGS="-X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=${VERSION}' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=${BUILD_DATE}'"
for target in "linux/amd64" "linux/arm64" "darwin/amd64" "darwin/arm64" "windows/amd64"; do
GOOS="${target%/*}"
GOARCH="${target#*/}"
@@ -71,11 +73,18 @@ jobs:
NAME="relspec-${GOOS}-${GOARCH}${EXT}"
GOOS="$GOOS" GOARCH="$GOARCH" go build \
-trimpath \
-ldflags "-X git.warky.dev/wdevs/relspecgo/cmd/relspec.version=${VERSION}" \
-ldflags "$LDFLAGS" \
-o "$NAME" ./cmd/relspec
echo "Built $NAME"
done
- name: Build Windows installer
run: |
sudo apt-get update
sudo apt-get install -y nsis
VERSION="${{ github.event.inputs.tag || github.ref_name }}"
makensis -DVERSION="${VERSION#v}" -DEXE="$PWD/relspec-windows-amd64.exe" -DOUT="$PWD/relspec-setup-windows-amd64.exe" windows/installer.nsi
- name: Create release and upload assets
run: |
TAG="${{ github.event.inputs.tag || github.ref_name }}"
@@ -234,11 +243,12 @@ jobs:
run: |
VERSION="${{ github.event.inputs.tag || github.ref_name }}"
PKGVER="${VERSION#v}"
BUILD_DATE="$(date -u +"%Y-%m-%d %H:%M:%S UTC")"
for GOARCH in amd64 arm64; do
GOOS=linux GOARCH=$GOARCH go build \
-trimpath \
-ldflags "-X git.warky.dev/wdevs/relspecgo/cmd/relspec.version=${PKGVER}" \
-ldflags "-X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=v${PKGVER}' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=${BUILD_DATE}'" \
-o relspec ./cmd/relspec
PKGDIR="relspec_${PKGVER}_${GOARCH}"
@@ -309,7 +319,13 @@ jobs:
rockylinux:9 \
bash -lc "
set -euo pipefail
dnf install -y rpm-build git &&
# Avoid transient/out-of-sync mirrors while bootstrapping the RPM builder.
sed -i \
-e 's|^mirrorlist=|#mirrorlist=|' \
-e 's|^#baseurl=http://dl.rockylinux.org|baseurl=https://dl.rockylinux.org|' \
/etc/yum.repos.d/rocky*.repo
dnf clean all
dnf -y --setopt=timeout=60 --setopt=retries=5 install rpm-build git
curl -fsSL https://go.dev/dl/go\${GO_VER}.linux-amd64.tar.gz | tar -C /usr/local -xz &&
export PATH=\$PATH:/usr/local/go/bin &&
mkdir -p ~/rpmbuild/{BUILD,BUILDROOT,RPMS,SOURCES,SPECS,SRPMS} &&
+59 -113
View File
@@ -1,145 +1,91 @@
# Contributing to RelSpec
Thank you for your interest in contributing to RelSpec.
## Setup
## Development Setup
- Go 1.25+ (see `go.mod`), Git
- Optional: golangci-lint, Docker/Podman (PostgreSQL integration tests)
### Prerequisites
- Go 1.21 or higher
- Git
- (Optional) golangci-lint for linting
- (Optional) Docker for database testing
### Getting Started
1. Clone the repository:
```bash
git clone https://github.com/wdevs/relspecgo.git
git clone git@git.warky.dev:wdevs/relspecgo.git
cd relspecgo
make deps
make build # outputs build/relspec
```
2. Install dependencies:
```bash
go mod download
```
## Make Targets
3. Run tests:
```bash
go test ./...
```
| Target | Purpose |
|-------------------------|------------------------------------------------------|
| `make test` | Unit tests (race detection, coverage) |
| `make test-integration` | Integration tests (needs `RELSPEC_TEST_PG_CONN`) |
| `make docker-test` | PostgreSQL integration tests via Docker/Podman |
| `make lint` | golangci-lint |
| `make fmt` / `fmt-check`| gofumpt + goimports |
| `make check` | vet, fmt-check, staticcheck, govulncheck |
| `make coverage` | Coverage report |
4. Build the project:
```bash
go build -o relspec ./cmd/relspec
```
Single test: `go test -run TestName ./pkg/readers/dbml`
## Project Structure
## Layout
```
relspecgo/
├── cmd/ # CLI application entry point
├── pkg/
│ ├── readers/ # Input format readers (XML, JSON, DCTX, DB, GORM, Bun)
│ ├── writers/ # Output format writers (GORM, Bun, JSON, YAML)
│ ├── models/ # Internal data models for relations
│ └── transform/ # Transformation and validation logic
├── examples/ # Usage examples and sample files
├── tests/ # Integration tests
└── .claude/ # Claude Code configuration and commands
cmd/relspec/ CLI commands (convert, diff, merge, split, edit, inspect, job, ...)
pkg/models/ Core model: Database > Schema > Table > Column/Constraint/Index/Relationship
pkg/readers/<fmt> One reader per format
pkg/writers/<fmt> One writer per format
pkg/diff, merge, inspector, jobs, transform, ui, pgsql, sqltypes ...
examples/ Sample files
tests/ Integration tests and assets
docs/ Feature docs
```
## Adding New Readers
## Adding a Reader
To add a new input format reader:
1. Create `pkg/readers/<format>/reader.go` with `NewReader(options *readers.ReaderOptions)`.
2. Implement `readers.Reader`: `ReadDatabase`, `ReadSchema`, `ReadTable`.
3. Add `reader_test.go` in the same package.
4. Register the format in the CLI switches (`cmd/relspec/convert.go`, `diff.go`, `edit.go`, ...).
5. Add a `README.md` in the reader directory.
1. Create a new file in `pkg/readers/` (e.g., `myformat_reader.go`)
2. Implement the `Reader` interface:
```go
type Reader interface {
Read(source string) (*models.Schema, error)
}
```
3. Add tests in `pkg/readers/myformat_reader_test.go`
4. Register the reader in the CLI
## Adding a Writer
## Adding New Writers
1. Create `pkg/writers/<format>/writer.go` with `NewWriter(options *writers.WriterOptions)`.
2. Implement `writers.Writer`: `WriteDatabase`, `WriteSchema`, `WriteTable`.
3. Add `writer_test.go` in the same package.
4. Register the format in the CLI switches.
5. Add a `README.md` in the writer directory.
To add a new output format writer:
## Code Rules
1. Create a new file in `pkg/writers/` (e.g., `myformat_writer.go`)
2. Implement the `Writer` interface:
```go
type Writer interface {
Write(schema *models.Schema, destination string) error
}
```
3. Add tests in `pkg/writers/myformat_writer_test.go`
4. Register the writer in the CLI
## Code Style
- Follow standard Go conventions
- Use `gofmt` for formatting
- Run `go vet` to check for issues
- Use meaningful variable and function names
- Add comments for exported functions and types
- Format with gofumpt/goimports (`make fmt`); `make check` must pass.
- Iterate `Table.Columns`, `Constraints`, `Indexes`, `Relationships` in sorted order (maps are unordered; output must be deterministic).
- Every writer stamps `buildinfo.GeneratedComment()` in its file header.
- Comment exported functions and types.
## Testing
- Write unit tests for all new functionality
- Aim for >80% code coverage
- Use table-driven tests where appropriate
- Include both positive and negative test cases
- Tests live in the same package as the code.
- Table-driven tests; cover positive and negative cases.
- Reuse existing test data in `tests/` and `examples/` before adding new data.
- Tests run with `-race`.
### Running Tests
## Commits
```bash
# All tests
go test ./...
- Format: `type(scope): description`
- Types: `feat`, `fix`, `docs`, `test`, `refactor`, `chore`, `ci`
- Keep commits focused. Reference issues where applicable.
# With coverage
go test -cover ./...
## Pull Requests
# Verbose output
go test -v ./...
1. Branch from `master`.
2. Add tests; `make test` and `make check` pass.
3. Update docs/README if behaviour changes.
4. Open a PR with a clear description.
# Specific package
go test ./pkg/readers/...
```
## Security
## Committing Changes
- Write clear, descriptive commit messages
- Follow conventional commits format: `type(scope): description`
- Types: feat, fix, docs, test, refactor, chore
- Example: `feat(readers): add PostgreSQL support`
- Keep commits focused and atomic
- Reference issues in commit messages when applicable
## Pull Request Process
1. Create a feature branch from `master`
2. Make your changes
3. Add tests for new functionality
4. Ensure all tests pass
5. Update documentation if needed
6. Submit a pull request with a clear description
## Claude Code Commands
This project includes Claude Code slash commands for common tasks:
- `/test` - Run all tests
- `/build` - Build the binary
- `/lint` - Run linters
- `/coverage` - Generate coverage report
## Questions or Issues?
- Open an issue for bugs or feature requests
- Start a discussion for questions or ideas
- Check existing issues before creating new ones
Do not report vulnerabilities in public issues. See [SECURITY.md](SECURITY.md).
## License
By contributing to RelSpec, you agree that your contributions will be licensed under the Apache License 2.0.
Contributions are licensed under the Apache License 2.0.
+35 -10
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: installer-windows 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=relspec
@@ -20,7 +20,9 @@ STATICCHECK = go run honnef.co/go/tools/cmd/staticcheck@latest
GOVULNCHECK = go run golang.org/x/vuln/cmd/govulncheck@latest
# Version information
VERSION := $(shell git describe --tags --always --dirty 2>/dev/null || echo "dev")
# VERSION file is the single source of truth (bumped by `make release-version`).
# Falls back to git describe for ad-hoc builds where VERSION hasn't been committed yet.
VERSION := $(shell [ -f VERSION ] && echo "v$$(cat VERSION | tr -d '[:space:]')" || git describe --tags --always --dirty 2>/dev/null || echo "dev")
BUILD_DATE := $(shell date -u +"%Y-%m-%d %H:%M:%S UTC")
LDFLAGS := -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=$(VERSION)' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=$(BUILD_DATE)'
@@ -223,6 +225,15 @@ release: lint fmt-check test build ## Run lint, format check, tests, build, then
echo "Creating new release: $$version"; \
commit_logs=$$(git log "$${latest_tag}..HEAD" --pretty=format:"- %s" --no-merges); \
fi; \
PKGVER=$${version#v}; \
echo "Updating package metadata to $$PKGVER..."; \
echo "$$PKGVER" > VERSION; \
sed -i "s/^pkgver=.*/pkgver=$$PKGVER/" linux/arch/PKGBUILD; \
sed -i "s/^Version:.*/Version: $$PKGVER/" linux/centos/relspec.spec; \
git add VERSION linux/arch/PKGBUILD linux/centos/relspec.spec; \
if ! git diff --cached --quiet; then \
git commit -m "chore(release): update package version to $$PKGVER"; \
fi; \
if [ -z "$$commit_logs" ]; then \
tag_message="Release $$version"; \
else \
@@ -232,21 +243,35 @@ release: lint fmt-check test build ## Run lint, format check, tests, build, then
git push origin "$$version"; \
echo "Tag $$version created and pushed to remote repository."
release-version: ## Auto-increment patch version, update package files, commit, tag, and push
@CURRENT=$$(git describe --tags --abbrev=0 2>/dev/null || echo "v0.0.0"); \
MAJOR=$$(echo $$CURRENT | sed 's/v\([0-9]*\)\.\([0-9]*\)\.\([0-9]*\).*/\1/'); \
MINOR=$$(echo $$CURRENT | sed 's/v\([0-9]*\)\.\([0-9]*\)\.\([0-9]*\).*/\2/'); \
PATCH=$$(echo $$CURRENT | sed 's/v\([0-9]*\)\.\([0-9]*\)\.\([0-9]*\).*/\3/'); \
NEXT="v$$MAJOR.$$MINOR.$$((PATCH + 1))"; \
release-version: lint fmt-check ## Run lint and format check, then auto-increment patch version, update VERSION/package files, commit, tag, and push
@CURRENT=$$([ -f VERSION ] && cat VERSION | tr -d '[:space:]' || (git describe --tags --abbrev=0 2>/dev/null | sed 's/^v//') || echo "0.0.0"); \
MAJOR=$$(echo $$CURRENT | cut -d. -f1); \
MINOR=$$(echo $$CURRENT | cut -d. -f2); \
PATCH=$$(echo $$CURRENT | cut -d. -f3); \
PKGVER="$$MAJOR.$$MINOR.$$((PATCH + 1))"; \
echo "Current: $$CURRENT → Next: $$NEXT"; \
NEXT="v$$PKGVER"; \
echo "Current: v$$CURRENT → Next: $$NEXT"; \
echo "$$PKGVER" > VERSION; \
sed -i "s/^pkgver=.*/pkgver=$$PKGVER/" linux/arch/PKGBUILD; \
sed -i "s/^Version:.*/Version: $$PKGVER/" linux/centos/relspec.spec; \
git add linux/arch/PKGBUILD linux/centos/relspec.spec; \
git add VERSION linux/arch/PKGBUILD linux/centos/relspec.spec; \
git commit -m "chore(release): update package version to $$PKGVER"; \
git tag -a "$$NEXT" -m "Release $$NEXT"; \
git push origin HEAD "$$NEXT"; \
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
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
installer-windows: ## Build the Windows binary and NSIS installer (requires makensis)
@echo "Building Windows installer..."
GOOS=windows GOARCH=amd64 $(GOBUILD) -trimpath -ldflags "$(LDFLAGS)" -o $(BUILD_DIR)/relspec-windows-amd64.exe ./cmd/relspec
makensis -DVERSION=$$(cat VERSION | tr -d '[:space:]') -DEXE=$(CURDIR)/$(BUILD_DIR)/relspec-windows-amd64.exe -DOUT=$(CURDIR)/$(BUILD_DIR)/relspec-setup-windows-amd64.exe windows/installer.nsi
+62 -4
View File
@@ -16,12 +16,18 @@
go install -v git.warky.dev/wdevs/relspecgo/cmd/relspec@latest
```
Windows: download `relspec-setup-windows-amd64.exe` from the
[latest release](https://git.warky.dev/wdevs/relspecgo/releases/latest)
(installs to Program Files and adds `relspec` to `PATH`). See [windows/README.md](windows/README.md).
## Supported Formats
| Direction | Formats |
|-----------|---------|
| **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` |
| **Readers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `mysql` `pgsql` `prisma` `sqldir` `sqlite` `typeorm` `yaml` |
| **Writers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `mysql` `pgsql` `prisma` `sqlexec` `sqlite` `template` `typeorm` `yaml` |
See [docs/FORMAT_EXAMPLES.md](docs/FORMAT_EXAMPLES.md) for usage examples covering every format.
## Commands
@@ -40,6 +46,26 @@ relspec convert --from pgsql --from-conn "postgres://..." --to sqlite --to-path
# Multiple input files merged
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
@@ -126,7 +152,7 @@ relspec job run build-schema
version: 1
jobs:
build-schema:
command: convert # closed allow-list: convert | merge | scripts-list
command: convert # closed allow-list: convert | merge | split | scripts-list | scripts-exec | templ | inspect | diff
description: Merge the DBML sources and emit PostgreSQL DDL
inputs:
- path: schema/core.dbml
@@ -147,7 +173,9 @@ resolved relative to the job file and may not escape it, and remote database
credentials are referenced by environment-variable name (`conn_env:`) and
redacted from logs. The whole plan — unknown commands/formats, duplicate job
names, missing inputs, path traversal, dependency cycles — is validated before
any job runs. See [docs/JOB_FILES.md](docs/JOB_FILES.md).
any job runs. Path-like fields also support `${NAME}` environment-variable
references, which are expanded and safety-checked during pre-flight. See
[docs/JOB_FILES.md](docs/JOB_FILES.md).
### `edit` — Interactive TUI editor
@@ -170,6 +198,16 @@ relspec edit --from pgsql --from-conn "postgres://user:pass@localhost/mydb" \
<img src="./assets/image/screenshots/edit_column.jpg">
</p>
### `update` — Check for a newer release
```bash
relspec update # check, prompt, then download + run installer (Windows)
relspec update --check # report only
relspec update --yes # skip the prompt
```
Non-Windows platforms print the release URL. `dev` and commit-hash builds are never reported as outdated.
## Development
**Prerequisites:** Go 1.24.0+
@@ -180,6 +218,7 @@ make test # race detection + coverage
make lint # requires golangci-lint
make coverage # → coverage.html
make install # → $GOPATH/bin
make installer-windows # → build/relspec-setup-windows-amd64.exe (needs makensis)
```
## Project Structure
@@ -194,6 +233,8 @@ pkg/merge/ Schema merging
pkg/models/ Internal data models
pkg/transform/ Transformation logic
pkg/pgsql/ PostgreSQL utilities
pkg/updatecheck/ Latest-release lookup and version comparison
windows/ NSIS installer script
pkg/sqltypes/ Nullable SQL types for generated/hand-written models (see below)
```
@@ -214,6 +255,23 @@ see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
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
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
+12
View File
@@ -0,0 +1,12 @@
# Security Policy
## Supported versions
Security fixes are released for the latest minor version of RelSpec.
## Reporting a vulnerability
Please do not open a public issue for security problems.
Report privately by email: warkydevs@gmail.com
Reports are reviewed on a best-effort basis. No response time or fix timeline is
guaranteed. Reporters may be credited in the release notes if they wish.
+20 -9
View File
@@ -4,20 +4,25 @@
- [✔️] **Database Inspector**
- [✔️] PostgreSQL driver (reader + writer)
- [ ] MySQL driver
- [ ] MySQL driver (only MariaDB datatype conversion in pkg/mariadb, no reader/writer)
- [✔️] SQLite driver (reader + writer with automatic schema flattening)
- [ ] MSSQL driver
- [✔️] MSSQL driver (reader + writer, generated and identity columns)
- [✔️] Foreign key detection
- [✔️] Index extraction
- [✔️] .sql file generation (PostgreSQL, SQLite)
- [✔️] .dbml: Database Markup Language (DBML) for textual schema representation.
- [✔️] Prisma schema support (PSL format) .prisma
- [ ] Sequelize (Typescript/Javascript) (Use templates, 💲 Someone can do this, not me)
- [✔️] Drizzle ORM support .ts (TypeScript / JavaScript) (Mr. Edd wanted to move from Prisma to Drizzle. If you are bugs, you are welcome to do pull requests or issues)
- [☠️] Entity Framework (.NET) model .edmx (Fuck no, EDMX files were bloated, verbose XML nightmares—hard to merge, error-prone, and a pain in teams. Microsoft wisely ditched them in EF Core for code-first. Classic overkill from old MS era.)
- [✔️] TypeORM support
- [] .hbm.xml / schema.xml: Hibernate/Propel mappings (Java/PHP) (💲 Someone can do this, not me)
- [ ] Django models.py (Python classes), Sequelize migrations (JS) (💲 Someone can do this, not me)
- [] .avsc: Avro schema (JSON format for data serialization) (💲 Someone can do this, not me)
- [ ] .hbm.xml / schema.xml: Hibernate/Propel mappings (Java/PHP) (Use templates, 💲 Someone can do this, not me)
- [ ] Django models.py (Python classes), Sequelize migrations (JS) (Use templates, 💲 Someone can do this, not me)
- [ ] SQLAlchemy, Python (Use templates, 💲 Someone can do this, not me)
- [ ] .avsc: Avro schema (JSON format for data serialization) (Use templates, 💲 Someone can do this, not me)
- [ ] Rails schema in db/schema.rb (Use templates, 💲 Someone can do this, not me)
- [ ] Laravel migration (Use templates, 💲 Someone can do this, not me)
- [ ] Hibernate / JPA (Use templates, 💲 Someone can do this, not me)
- [✔️] GraphQL schema generation
## UI
@@ -39,13 +44,19 @@
## Advanced Features
- [ ] Dry-run mode for validation
- [x] Diff tool for comparing specifications
- [ ] Migration script generation
- [ ] Dry-run mode for validation (only `job run --dry-run`; not on convert/merge/split)
- [✔️] Diff tool for comparing specifications
- [✔️] Migration script generation (PostgreSQL diff, live database by default for direct output)
- [ ] Custom type mapping configuration
- [ ] Batch processing support
- [ ] Batch processing support (partial: job files run named workflows via `relspec job run`)
- [ ] Watch mode for auto-regeneration
## Distribution
- [ ] NSIS Windows installer
- [ ] Check https://git.warky.dev/wdevs/relspecgo/releases for new releases and prompt to update
- [ ] CI pipeline to build and publish the Windows installer
## Future Considerations
- [ ] Web UI for visual editing
+1
View File
@@ -0,0 +1 @@
1.0.86
+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)
}
})
}
}
+153 -15
View File
@@ -3,6 +3,7 @@ package main
import (
stdjson "encoding/json"
"fmt"
"io"
"os"
"strings"
"time"
@@ -21,6 +22,7 @@ import (
"git.warky.dev/wdevs/relspecgo/pkg/readers/graphql"
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
"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/prisma"
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
@@ -36,6 +38,7 @@ import (
wgraphql "git.warky.dev/wdevs/relspecgo/pkg/writers/graphql"
wjson "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
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"
wprisma "git.warky.dev/wdevs/relspecgo/pkg/writers/prisma"
wsqlite "git.warky.dev/wdevs/relspecgo/pkg/writers/sqlite"
@@ -57,6 +60,9 @@ var (
convertNullableArrays string
convertContinueOnError bool
convertExtraFields string
convertDryRun bool
convertWatch bool
convertWatchInterval time.Duration
)
var convertCmd = &cobra.Command{
@@ -165,7 +171,11 @@ Examples:
# Convert SQLite to PostgreSQL SQL
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,
}
@@ -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().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")
if err != nil {
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 {
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, "Started at: %s\n\n", getCurrentTimestamp())
@@ -229,6 +251,7 @@ func runConvert(cmd *cobra.Command, args []string) error {
if err != nil {
return fmt.Errorf("failed to read source: %w", err)
}
finalizeCommentedRefs(db, stderrWarn)
fmt.Fprintf(os.Stderr, " ✓ Successfully read database '%s'\n", db.Name)
fmt.Fprintf(os.Stderr, " Found: %d schema(s)\n", len(db.Schemas))
@@ -239,6 +262,22 @@ func runConvert(cmd *cobra.Command, args []string) error {
}
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
fmt.Fprintf(os.Stderr, "[2/2] Writing to target format...\n")
fmt.Fprintf(os.Stderr, " Format: %s\n", convertTargetType)
@@ -261,6 +300,18 @@ func runConvert(cmd *cobra.Command, args []string) error {
return nil
}
// finalizeCommentedRefs resolves DBML `// Ref:` comments against the fully
// loaded model and warns about refs whose target is not loaded.
func finalizeCommentedRefs(db *models.Database, warn func(string)) {
for _, w := range dbml.ResolveCommentedRefs(db, true) {
warn(w)
}
}
func stderrWarn(msg string) {
fmt.Fprintf(os.Stderr, " ⚠ %s\n", msg)
}
func readDatabaseListForConvert(dbType string, files []string) (*models.Database, error) {
if len(files) == 0 {
return nil, fmt.Errorf("file list is empty")
@@ -367,6 +418,12 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
}
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":
// SQLite can use either file path or connection string
dbPath := filePath
@@ -395,23 +452,12 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
if extraFields != "" {
if !strings.EqualFold(dbType, "bun") {
return fmt.Errorf("--extra-fields is only supported for Bun output")
}
extraFieldsJSON, err := os.ReadFile(extraFields)
extraFieldsJSON, err := loadExtraFields(dbType, extraFields)
if err != nil {
return fmt.Errorf("failed to read --extra-fields file %q: %w", extraFields, 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")
return err
}
writerOpts.Metadata = map[string]interface{}{
"extra_fields": string(extraFieldsJSON),
"extra_fields": extraFieldsJSON,
}
}
@@ -452,6 +498,9 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
case "mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022":
writer = wmssql.NewWriter(writerOpts)
case "mysql", "mariadb":
writer = wmysql.NewWriter(writerOpts)
case "sqlite", "sqlite3":
writer = wsqlite.NewWriter(writerOpts)
@@ -516,6 +565,95 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
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
func getSchemaNames(db *models.Database) []string {
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)
}
}
+57 -19
View File
@@ -60,8 +60,9 @@ inputs, output and options:
logfile: .relspec/log/build-schema.log
Rules and guarantees:
- command is a closed allow-list (convert, merge, scripts-list, templ). Arbitrary
shell strings are never executed.
- command is a closed allow-list (convert, merge, split, scripts-list,
scripts-exec, templ, inspect, diff). Arbitrary shell strings are never
executed.
- Every path is relative to the directory holding the job file and may not
escape it. Absolute and home-relative paths are rejected.
- Remote database credentials are referenced by environment-variable name
@@ -246,22 +247,37 @@ type resolvedInput struct {
func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolvedJob, error) {
root := j.Dir()
rj := &resolvedJob{job: j, root: root, logPolicy: j.ResolvedLogPolicy()}
expandPath := func(label, value string) (string, error) {
expanded, err := jobs.ExpandEnv(value)
if err != nil {
return "", fmt.Errorf("%s: %w", label, err)
}
return expanded, nil
}
if j.Logfile != "" {
p, err := jobs.SafeJoin(root, j.Logfile)
logfile, err := expandPath("logfile", j.Logfile)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, logfile)
if err != nil {
return nil, fmt.Errorf("logfile: %w", err)
}
rj.logPath = p
}
if j.Template != "" {
p, err := jobs.SafeJoin(root, j.Template)
templatePath, err := expandPath("template", j.Template)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, templatePath)
if err != nil {
return nil, fmt.Errorf("template: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("template %q: not found or is a directory", j.Template)
return nil, fmt.Errorf("template %q: not found or is a directory", templatePath)
}
rj.templatePath = p
}
@@ -291,16 +307,20 @@ func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolve
ri.connEnv = in.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, in.Path)
inputPath, err := expandPath(fmt.Sprintf("input[%d]", i), in.Path)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, inputPath)
if err != nil {
return nil, fmt.Errorf("input[%d]: %w", i, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("input[%d]: %s: file not found", i, in.Path)
return nil, fmt.Errorf("input[%d]: %s: file not found", i, inputPath)
}
if info.IsDir() {
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, in.Path)
return nil, fmt.Errorf("input[%d]: %s: is a directory, not a file", i, inputPath)
}
ri.path = p
}
@@ -308,13 +328,17 @@ func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolve
}
for _, d := range j.ScriptDirs {
p, err := jobs.SafeJoin(root, d)
scriptDir, err := expandPath("script_dir", d)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, scriptDir)
if err != nil {
return nil, fmt.Errorf("script_dir %q: %w", d, err)
}
info, err := os.Stat(p)
if err != nil {
return nil, fmt.Errorf("script_dir %q: not found", d)
return nil, fmt.Errorf("script_dir %q: not found", scriptDir)
}
if !info.IsDir() {
return nil, fmt.Errorf("script_dir %q: not a directory", d)
@@ -332,25 +356,33 @@ func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolve
rj.outputConnEnv = j.Output.ConnEnv
rj.secrets = append(rj.secrets, v)
} else {
p, err := jobs.SafeJoin(root, j.Output.Path)
outputPath, err := expandPath("output", j.Output.Path)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, outputPath)
if err != nil {
return nil, fmt.Errorf("output: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Output.Overwrite {
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", j.Output.Path)
return nil, fmt.Errorf("output %s already exists (set output.overwrite: true to replace it)", outputPath)
}
rj.outputPath = p
}
}
if j.Rules != "" {
p, err := jobs.SafeJoin(root, j.Rules)
rulesPath, err := expandPath("rules", j.Rules)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, rulesPath)
if err != nil {
return nil, fmt.Errorf("rules: %w", err)
}
info, err := os.Stat(p)
if err != nil || info.IsDir() {
return nil, fmt.Errorf("rules %q: not found or is a directory", j.Rules)
return nil, fmt.Errorf("rules %q: not found or is a directory", rulesPath)
}
rj.rulesPath = p
}
@@ -358,12 +390,16 @@ func preflightJob(j *jobs.Job, resolvedByName map[string]*resolvedJob) (*resolve
if j.Report != nil {
rj.reportFormat = strings.ToLower(j.Report.Format)
if j.Report.Path != "" {
p, err := jobs.SafeJoin(root, j.Report.Path)
reportPath, err := expandPath("report", j.Report.Path)
if err != nil {
return nil, err
}
p, err := jobs.SafeJoin(root, reportPath)
if err != nil {
return nil, fmt.Errorf("report: %w", err)
}
if _, err := os.Stat(p); err == nil && !j.Report.Overwrite {
return nil, fmt.Errorf("report %s already exists (set report.overwrite: true to replace it)", j.Report.Path)
return nil, fmt.Errorf("report %s already exists (set report.overwrite: true to replace it)", reportPath)
}
rj.reportPath = p
}
@@ -542,6 +578,7 @@ func runMergeJob(rj *resolvedJob, lg *jobLogger) error {
lg.logf("merging: %s", inputLabel(ri))
merge.MergeDatabases(base, db, opts)
}
finalizeCommentedRefs(base, func(w string) { lg.logf("warning: %s", w) })
base.UpdateDate()
return writeJobOutput(rj, base, lg)
}
@@ -778,6 +815,7 @@ func readJobInputs(rj *resolvedJob, lg *jobLogger) (*models.Database, error) {
if base == nil {
return nil, fmt.Errorf("no inputs produced a database")
}
finalizeCommentedRefs(base, func(w string) { lg.logf("warning: %s", w) })
return base, nil
}
@@ -805,8 +843,8 @@ func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
return fmt.Errorf("database output is only supported for pgsql (got %q)", rj.job.Output.Format)
}
lg.logf("writing output to database env:%s", rj.outputConnEnv)
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, "", "", o.ContinueOnError)
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, o.Types, o.ArrayNullable, o.ContinueOnError)
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn, "full_ddl": o.FullDDL}
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
}
@@ -816,7 +854,7 @@ func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
lg.logf("writing output: %s (%s)", rj.outputPath, format)
write := func(target string) error {
return writeDatabase(db, format, target, o.Package, o.Schema, o.FlattenSchema, "", "", o.ContinueOnError, "")
return writeDatabase(db, format, target, o.Package, o.Schema, o.FlattenSchema, o.Types, o.ArrayNullable, o.ContinueOnError, "")
}
// Single-file formats are written to a temp file and renamed into place so
// a failure never leaves a partial or truncated output. Directory-emitting
+59
View File
@@ -666,6 +666,7 @@ jobs:
format: dbml
template: templates/schema.tmpl
output:
format: text
path: build/schema.txt
overwrite: true
`)
@@ -682,3 +683,61 @@ jobs:
t.Fatalf("templ output missing users table: %s", out)
}
}
func TestJobRun_BunSqlTypesOption(t *testing.T) {
dir := jobFixture(t, `version: 1
jobs:
models:
command: convert
inputs:
- path: schema/core.dbml
format: dbml
output:
format: bun
path: build/models
options:
package: models
types: sqltypes
`)
set := mustLoadSet(t, filepath.Join(dir, "relspec.yml"))
if err := executeJobPlan(set, "models", false, false, &bytes.Buffer{}); err != nil {
t.Fatalf("execute Bun job: %v", err)
}
generated, err := os.ReadFile(filepath.Join(dir, "build", "models"))
if err != nil {
t.Fatalf("read Bun output: %v", err)
}
if !strings.Contains(string(generated), "git.warky.dev/wdevs/relspecgo/pkg/sqltypes") {
t.Fatalf("Bun output did not preserve options.types: %s", generated)
}
}
func TestJobValidation_TemplTextFormat(t *testing.T) {
validDir := jobFixture(t, `version: 1
jobs:
docs:
command: templ
inputs: [{path: schema/core.dbml, format: dbml}]
template: schema.tmpl
output: {format: text, path: build/docs.txt}
`)
valid := mustLoadSet(t, filepath.Join(validDir, "relspec.yml"))
if err := valid.Validate(); err != nil {
t.Fatalf("text output should be valid for templ: %v", err)
}
invalidDir := jobFixture(t, `version: 1
jobs:
docs:
command: templ
inputs: [{path: schema/core.dbml, format: dbml}]
template: schema.tmpl
output: {format: json, path: build/docs.txt}
`)
invalid, err := jobs.Load([]string{filepath.Join(invalidDir, "relspec.yml")})
if err != nil {
t.Fatalf("load invalid manifest: %v", err)
}
if err := invalid.Validate(); err == nil || !strings.Contains(err.Error(), "output.format: text") {
t.Fatalf("expected templ format rejection, got %v", err)
}
}
+39
View File
@@ -59,7 +59,9 @@ var (
mergeSkipTables string // Comma-separated table names to skip
mergeVerbose bool
mergeReportPath string // Path to write merge report
mergeFullDDL bool // Execute full DDL instead of diffing the live pgsql database
mergeFlattenSchema bool
mergeDryRun bool
)
var mergeCmd = &cobra.Command{
@@ -127,7 +129,9 @@ func init() {
mergeCmd.Flags().BoolVar(&mergeSkipSequences, "skip-sequences", false, "Skip sequences during merge")
mergeCmd.Flags().StringVar(&mergeSkipTables, "skip-tables", "", "Comma-separated list of table names to skip during merge")
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().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)")
}
@@ -248,6 +252,7 @@ func runMerge(cmd *cobra.Command, args []string) error {
}
result := merge.MergeDatabases(targetDB, sourceDB, opts)
finalizeCommentedRefs(targetDB, stderrWarn)
// Update timestamp
targetDB.UpdateDate()
@@ -261,6 +266,29 @@ func runMerge(cmd *cobra.Command, args []string) error {
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
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
@@ -282,6 +310,16 @@ func runMerge(cmd *cobra.Command, args []string) error {
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) {
var reader readers.Reader
@@ -442,6 +480,7 @@ func writeDatabaseForMerge(dbType, filePath, connString string, db *models.Datab
if connString != "" {
writerOpts.Metadata = map[string]interface{}{
"connection_string": connString,
"full_ddl": mergeFullDDL,
}
// Add report path if merge report is enabled
if mergeReportPath != "" {
@@ -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,
NullableTypes: nullableTypes,
NullableArrays: nullableArrays,
TypeMappings: typeMappings,
Prisma7: prisma7,
ContinueOnError: continueOnError,
StrictDirectives: strictDirectives,
+11
View File
@@ -6,6 +6,7 @@ import (
"github.com/spf13/cobra"
"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.
@@ -17,6 +18,8 @@ var (
noVersion bool
silent bool
strictDirectives bool
typeMapFlags []string
typeMappings map[string]string
)
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,
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
JSON, YAML, SQL, etc.).`,
PersistentPreRunE: func(cmd *cobra.Command, args []string) error {
var err error
typeMappings, err = writers.ParseTypeMappings(typeMapFlags)
return err
},
}
func init() {
rootCmd.AddCommand(convertCmd)
rootCmd.AddCommand(batchCmd)
rootCmd.AddCommand(diffCmd)
rootCmd.AddCommand(inspectCmd)
rootCmd.AddCommand(scriptsCmd)
@@ -42,8 +51,10 @@ func init() {
rootCmd.AddCommand(mergeCmd)
rootCmd.AddCommand(splitCmd)
rootCmd.AddCommand(versionCmd)
rootCmd.AddCommand(updateCmd)
rootCmd.AddCommand(reportCmd)
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(&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:, …)")
+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
splitNullableTypes string
splitNullableArrays string
splitDryRun bool
)
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(&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")
if err != nil {
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)
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
fmt.Fprintf(os.Stderr, "[3/3] Writing to target format...\n")
fmt.Fprintf(os.Stderr, " Format: %s\n", splitTargetType)
+1
View File
@@ -114,6 +114,7 @@ func runTempl(cmd *cobra.Command, args []string) error {
if err != nil {
return fmt.Errorf("failed to read source: %w", err)
}
finalizeCommentedRefs(db, stderrWarn)
// Print database stats
schemaCount := len(db.Schemas)
+148
View File
@@ -0,0 +1,148 @@
package main
import (
"bufio"
"context"
"fmt"
"io"
"net/http"
"os"
"os/exec"
"path/filepath"
"runtime"
"strings"
"time"
"github.com/spf13/cobra"
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
"git.warky.dev/wdevs/relspecgo/pkg/updatecheck"
)
// windowsInstallerAsset is the release asset built by the NSIS installer step.
const windowsInstallerAsset = "relspec-setup-windows-amd64.exe"
var (
updateCheckOnly bool
updateAssumeYes bool
updateAPIURL = updatecheck.DefaultAPIURL
)
var updateCmd = &cobra.Command{
Use: "update",
Short: "Check for a newer release and offer to install it",
Long: `Check the project releases for a version newer than this binary.
If one exists you are prompted to update. On Windows the installer is
downloaded and started; on other platforms the release page URL is shown.
Use --check to only report, and --yes to skip the prompt.`,
RunE: func(cmd *cobra.Command, args []string) error {
ctx, cancel := context.WithTimeout(cmd.Context(), 10*time.Minute)
defer cancel()
return runUpdate(ctx, updateOptions{
current: buildinfo.Version,
apiURL: updateAPIURL,
goos: runtime.GOOS,
checkOnly: updateCheckOnly,
assumeYes: updateAssumeYes,
in: cmd.InOrStdin(),
out: cmd.OutOrStdout(),
install: downloadAndRunInstaller,
})
},
}
func init() {
updateCmd.Flags().BoolVar(&updateCheckOnly, "check", false, "Only report whether an update is available")
updateCmd.Flags().BoolVarP(&updateAssumeYes, "yes", "y", false, "Update without prompting")
}
type updateOptions struct {
current string
apiURL string
goos string
checkOnly bool
assumeYes bool
in io.Reader
out io.Writer
// install downloads and starts the installer found at url.
install func(ctx context.Context, url string, out io.Writer) error
}
func runUpdate(ctx context.Context, o updateOptions) error {
rel, err := updatecheck.Latest(ctx, nil, o.apiURL)
if err != nil {
return err
}
if !updatecheck.IsNewer(o.current, rel.Tag) {
_, _ = fmt.Fprintf(o.out, "RelSpec %s is up to date (latest release: %s)\n", o.current, rel.Tag)
return nil
}
_, _ = fmt.Fprintf(o.out, "A newer RelSpec is available: %s (installed: %s)\n", rel.Tag, o.current)
if o.checkOnly {
_, _ = fmt.Fprintf(o.out, "Release: %s\n", rel.URL)
return nil
}
installer, canInstall := rel.FindAsset(windowsInstallerAsset)
if o.goos != "windows" || !canInstall {
_, _ = fmt.Fprintf(o.out, "Download it from: %s\n", rel.URL)
return nil
}
if !o.assumeYes && !confirm(o.in, o.out, "Download and run the installer now?") {
_, _ = fmt.Fprintf(o.out, "Skipped. Release: %s\n", rel.URL)
return nil
}
return o.install(ctx, installer.URL, o.out)
}
// confirm asks a yes/no question, defaulting to no.
func confirm(in io.Reader, out io.Writer, question string) bool {
_, _ = fmt.Fprintf(out, "%s [y/N]: ", question)
line, _ := bufio.NewReader(in).ReadString('\n')
switch strings.ToLower(strings.TrimSpace(line)) {
case "y", "yes":
return true
}
return false
}
// downloadAndRunInstaller saves the installer to a temp directory and starts
// it detached so this process can exit and release relspec.exe.
func downloadAndRunInstaller(ctx context.Context, url string, out io.Writer) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return err
}
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("downloading installer: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("downloading installer: unexpected status %s", resp.Status)
}
dir, err := os.MkdirTemp("", "relspec-update-")
if err != nil {
return err
}
path := filepath.Join(dir, windowsInstallerAsset)
f, err := os.Create(path)
if err != nil {
return err
}
if _, err := io.Copy(f, resp.Body); err != nil {
_ = f.Close()
return fmt.Errorf("downloading installer: %w", err)
}
if err := f.Close(); err != nil {
return err
}
_, _ = fmt.Fprintf(out, "Starting installer: %s\n", path)
return exec.Command(path).Start()
}
+110
View File
@@ -0,0 +1,110 @@
package main
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func releaseServer(t *testing.T, tag string, withInstaller bool) *httptest.Server {
t.Helper()
assets := ""
if withInstaller {
assets = fmt.Sprintf(`{"name":%q,"browser_download_url":"https://x/setup.exe"}`, windowsInstallerAsset)
}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = fmt.Fprintf(w, `{"tag_name":%q,"html_url":"https://x/release","assets":[%s]}`, tag, assets)
}))
t.Cleanup(srv.Close)
return srv
}
func TestRunUpdate(t *testing.T) {
tests := []struct {
name string
current string
latest string
withInstaller bool
goos string
checkOnly bool
assumeYes bool
stdin string
wantOut []string
wantInstall bool
}{
{name: "up to date", current: "v1.0.5", latest: "v1.0.5", goos: "windows", withInstaller: true, wantOut: []string{"up to date"}},
{name: "dev build never prompts", current: "dev", latest: "v9.0.0", goos: "windows", withInstaller: true, wantOut: []string{"up to date"}},
{name: "check only", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, checkOnly: true, wantOut: []string{"v1.0.6", "https://x/release"}},
{name: "non-windows shows url", current: "v1.0.5", latest: "v1.0.6", goos: "linux", withInstaller: true, wantOut: []string{"Download it from: https://x/release"}},
{name: "windows without installer asset", current: "v1.0.5", latest: "v1.0.6", goos: "windows", wantOut: []string{"Download it from"}},
{name: "windows prompt yes", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, stdin: "y\n", wantOut: []string{"[y/N]"}, wantInstall: true},
{name: "windows prompt no", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, stdin: "n\n", wantOut: []string{"Skipped"}},
{name: "windows prompt empty", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, wantOut: []string{"Skipped"}},
{name: "windows --yes", current: "v1.0.5", latest: "v1.0.6", goos: "windows", withInstaller: true, assumeYes: true, wantInstall: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
srv := releaseServer(t, tt.latest, tt.withInstaller)
var out bytes.Buffer
var installedURL string
err := runUpdate(context.Background(), updateOptions{
current: tt.current, apiURL: srv.URL, goos: tt.goos,
checkOnly: tt.checkOnly, assumeYes: tt.assumeYes,
in: strings.NewReader(tt.stdin), out: &out,
install: func(_ context.Context, url string, _ io.Writer) error {
installedURL = url
return nil
},
})
if err != nil {
t.Fatal(err)
}
for _, want := range tt.wantOut {
if !strings.Contains(out.String(), want) {
t.Errorf("output missing %q:\n%s", want, out.String())
}
}
if (installedURL != "") != tt.wantInstall {
t.Errorf("installer called = %v, want %v", installedURL != "", tt.wantInstall)
}
if tt.wantInstall && installedURL != "https://x/setup.exe" {
t.Errorf("installer url = %q", installedURL)
}
})
}
}
func TestRunUpdate_Errors(t *testing.T) {
bad := httptest.NewServer(http.NotFoundHandler())
defer bad.Close()
if err := runUpdate(context.Background(), updateOptions{current: "v1.0.0", apiURL: bad.URL, out: io.Discard}); err == nil {
t.Error("expected lookup error")
}
srv := releaseServer(t, "v1.0.6", true)
want := errors.New("boom")
err := runUpdate(context.Background(), updateOptions{
current: "v1.0.5", apiURL: srv.URL, goos: "windows", assumeYes: true, out: io.Discard,
install: func(context.Context, string, io.Writer) error { return want },
})
if !errors.Is(err, want) {
t.Errorf("err = %v", err)
}
}
func TestDownloadAndRunInstaller_BadStatus(t *testing.T) {
srv := httptest.NewServer(http.NotFoundHandler())
defer srv.Close()
if err := downloadAndRunInstaller(context.Background(), srv.URL, io.Discard); err == nil {
t.Error("expected error for non-200")
}
if err := downloadAndRunInstaller(context.Background(), "://bad", io.Discard); err == nil {
t.Error("expected error for bad url")
}
}
+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).
+13 -2
View File
@@ -55,6 +55,10 @@ already forbid duplicate keys within a single file.
`script_dirs[]`, `template`, `logfile`) is **relative to the directory
containing the job file that declared the job**, not the process working
directory.
* Those path fields support `${NAME}` environment-variable references. They
are expanded immediately before execution; a missing variable is an error.
The expanded value is still subject to all relative-path and symlink safety
checks.
* Absolute paths, `~`-relative paths and any path that resolves outside the job
file directory (`../`, `a/../../b`, …) are **rejected during validation** —
before anything runs.
@@ -183,12 +187,17 @@ jobs:
format: pgsql
path: build/schema.sql # file output, OR:
conn_env: TARGET_DB_URL # execute against DB (pgsql only)
# default: reads live DB, executes only the diff
# (file output = full DDL); set options.full_ddl to force full DDL
overwrite: false # default false
options:
flatten_schema: false
schema: public
package: models # for gorm/bun output
types: sqltypes # Bun/GORM nullable types: baselib|stdlib|sqltypes
array_nullable: pointer_slice # Bun: slice|pointer_slice
continue_on_error: false # pgsql / scripts-exec output
full_ddl: false # pgsql database output: true = run full DDL; default diffs against live DB
skip_relations: false # merge only
skip_enums: false
skip_views: false
@@ -201,8 +210,10 @@ jobs:
For `templ`, `inputs` use the same file or `pgsql`/`conn_env` source forms as
schema conversion. `output` is optional (empty means stdout); when present it
contains only `path` and `overwrite`, because templates do not select a schema
writer format.
contains `path` and `overwrite`, plus the explicit `format: text` marker. The
marker is templ-only: it describes arbitrary text produced by a Go template,
not a database schema format. It is optional for compatibility with older
manifests and has no effect on template execution.
A `from_job` input takes no `path`, `format` or `conn_env`: it resolves to the
named job's `output.path` and inherits its format, and implies a dependency on
+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.
+80
View File
@@ -0,0 +1,80 @@
# RelSpec job examples
Run these commands from this directory:
```bash
cd examples/jobs
# List every job from relspec.yml and the additional relspec.*.yml files.
relspec job list
# Validate a job and print its resolved plan without reading or writing data.
relspec job run build-schema --plan
# Build PostgreSQL DDL from the two DBML inputs.
relspec job run build-schema
# Run a job's dependency chain. This runs build-schema first, then build-json.
relspec job run build-json
# Inspect build-json's output. The from_job input adds the dependency automatically.
relspec job run lint-schema
# Extract only the posts table.
relspec job run posts-only
# List migration scripts in deterministic order without connecting to a database.
relspec job run migration-order
# Render one TypeScript file per table using the example template.
relspec job run generate-typescript
# Compare the two example schemas and write the summary to the logfile.
relspec job run compare-schemas
# Use an environment variable in an output path.
export VST=build/typescript
relspec job run generate-typescript --plan
relspec job run generate-typescript
```
Database jobs use environment variables for connection strings. Set them in
the shell before running the jobs; do not put connection strings in YAML:
```bash
export SOURCE_DB_URL='postgres://user:password@localhost/source_db'
export TARGET_DB_URL='postgres://user:password@localhost/target_db'
# Read the source database and write a local DBML snapshot.
relspec job run snapshot-database
# Preview migration execution without connecting to PostgreSQL.
relspec job run apply-migrations --plan
# Execute the migrations against TARGET_DB_URL.
relspec job run apply-migrations
```
The `apply-migrations` job demonstrates execution from multiple directories:
```yaml
script_dirs:
- migrations/core
- migrations/tenant
output:
format: pgsql
conn_env: TARGET_DB_URL
```
RelSpec combines the SQL files from both folders in deterministic order before
executing them against the configured PostgreSQL database.
To use a different directory or manifest, pass `--dir` or `--file`:
```bash
relspec job list --dir /path/to/project
relspec job run build-schema --file /path/to/project/relspec.yml --plan
```
See [../../docs/JOB_FILES.md](../../docs/JOB_FILES.md) for the complete job
file reference.
+29
View File
@@ -0,0 +1,29 @@
# Database-oriented examples. Set the referenced environment variables before
# running these jobs; connection strings never belong in a job file.
version: 1
jobs:
snapshot-database:
command: convert
description: Read a PostgreSQL database and write a DBML snapshot
inputs:
- format: pgsql
conn_env: SOURCE_DB_URL
output:
format: dbml
path: build/database-snapshot.dbml
overwrite: true
apply-migrations:
command: scripts-exec
description: Execute migrations from core and tenant folders against PostgreSQL
# All SQL files from these directories are discovered in deterministic order.
script_dirs:
- migrations/core
- migrations/tenant
output:
format: pgsql
conn_env: TARGET_DB_URL
options:
continue_on_error: false
logfile: .relspec/log/apply-migrations.log
+32
View File
@@ -0,0 +1,32 @@
# Reporting and templating examples. These jobs are discovered together with
# relspec.yml when running `relspec job list` from this directory.
version: 1
jobs:
generate-typescript:
command: templ
description: Render one TypeScript file per table from the merged schema
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
template: templates/schema.tmpl
mode: table
filename_pattern: "{{.Name}}.ts"
output:
format: text
path: "${VST}"
overwrite: true
compare-schemas:
command: diff
description: Compare the two example schemas and print a summary
inputs:
- path: schema/core.dbml
format: dbml
- path: schema/tenant.dbml
format: dbml
report:
format: summary
logfile: .relspec/log/compare-schemas.log
+6
View File
@@ -0,0 +1,6 @@
// Code generated by RelSpec. DO NOT EDIT.
export interface {{.Name}} {
{{- range values .Table.Columns}}
{{.Name}}: {{.Type}};
{{- end}}
}
+2
View File
@@ -4,6 +4,7 @@ go 1.25.13
require (
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/jackc/pgx/v5 v5.9.2
github.com/microsoft/go-mssqldb v1.10.0
@@ -18,6 +19,7 @@ require (
)
require (
filippo.io/edwards25519 v1.1.0 // indirect
github.com/davecgh/go-spew v1.1.1 // indirect
github.com/dustin/go-humanize 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/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
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/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
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/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
+2 -2
View File
@@ -1,6 +1,6 @@
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
pkgname=relspec
pkgver=1.0.74
pkgver=1.0.86
pkgrel=1
pkgdesc="RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs."
arch=('x86_64' 'aarch64')
@@ -15,7 +15,7 @@ build() {
export CGO_ENABLED=0
go build \
-trimpath \
-ldflags "-X git.warky.dev/wdevs/relspecgo/cmd/relspec.version=$pkgver" \
-ldflags "-X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=v$pkgver' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=$(date -u +"%Y-%m-%d %H:%M:%S UTC")'" \
-o "$pkgname" ./cmd/relspec
}
+2 -2
View File
@@ -1,5 +1,5 @@
Name: relspec
Version: 1.0.74
Version: 1.0.86
Release: 1%{?dist}
Summary: RelSpec is a comprehensive database relations management tool that reads, transforms, and writes database table specifications across multiple formats and ORMs.
@@ -25,7 +25,7 @@ DBML, GraphQL, and more.
export CGO_ENABLED=0
go build \
-trimpath \
-ldflags "-X git.warky.dev/wdevs/relspecgo/cmd/relspec.version=%{version}" \
-ldflags "-X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=v%{version}' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=$(date -u +"%%Y-%%m-%%d %%H:%%M:%%S UTC")'" \
-o %{name} ./cmd/relspec
%install
+17
View File
@@ -71,6 +71,13 @@ func compareSchemas(source, target []*models.Schema) *SchemaDiff {
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 {
change := &SchemaChange{
Name: source.Name,
@@ -78,6 +85,16 @@ func compareSchemaDetails(source, target *models.Schema) *SchemaChange {
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
tableDiff := compareTables(source.Tables, target.Tables)
if !isEmpty(tableDiff) {
+345
View File
@@ -0,0 +1,345 @@
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")
}
}
+6 -5
View File
@@ -18,11 +18,12 @@ type SchemaDiff struct {
// SchemaChange represents changes within a schema
type SchemaChange struct {
Name string `json:"name"`
Tables *TableDiff `json:"tables,omitempty"`
Views *ViewDiff `json:"views,omitempty"`
Sequences *SequenceDiff `json:"sequences,omitempty"`
Scripts *ScriptDiff `json:"scripts,omitempty"`
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"`
Views *ViewDiff `json:"views,omitempty"`
Sequences *SequenceDiff `json:"sequences,omitempty"`
Scripts *ScriptDiff `json:"scripts,omitempty"`
}
// TableDiff represents differences in tables
+1
View File
@@ -123,6 +123,7 @@ rules:
| `missing_primary_key` | `have_primary_key` | Ensure tables have primary keys |
| `orphaned_foreign_key` | `orphaned_foreign_key` | Detect FKs referencing non-existent tables |
| `circular_dependency` | `circular_dependency` | Detect circular FK dependencies |
| `duplicate_index_name` | `duplicate_index_name` | Index / PK / unique names must be unique per schema (default: `enforce`) |
## Rule Configuration
+1
View File
@@ -161,6 +161,7 @@ func getValidator(functionName string) (validatorFunc, bool) {
"have_primary_key": validateMissingPrimaryKey,
"orphaned_foreign_key": validateOrphanedForeignKey,
"circular_dependency": validateCircularDependency,
"duplicate_index_name": validateDuplicateIndexName,
}
fn, exists := validators[functionName]
+5
View File
@@ -154,6 +154,11 @@ func GetDefaultConfig() *Config {
Function: "circular_dependency",
Message: "Circular foreign key dependency detected",
},
"duplicate_index_name": {
Enabled: "enforce",
Function: "duplicate_index_name",
Message: "Index name is reused within the schema; PostgreSQL skips the duplicate CREATE INDEX IF NOT EXISTS",
},
},
}
}
+1
View File
@@ -37,6 +37,7 @@ func TestGetDefaultConfig(t *testing.T) {
"missing_primary_key",
"orphaned_foreign_key",
"circular_dependency",
"duplicate_index_name",
}
for _, ruleName := range expectedRules {
+61
View File
@@ -643,3 +643,64 @@ func contains(slice []string, value string) bool {
}
return false
}
// validateDuplicateIndexName checks that index names are unique per schema.
// PostgreSQL keeps indexes in the schema-wide relation namespace, so a name
// reused on another table makes CREATE INDEX IF NOT EXISTS silently skip it.
// Primary key and unique constraints create backing indexes and share that
// namespace too. An index and a constraint with the same name on the same
// table describe one object and are not reported.
func validateDuplicateIndexName(db *models.Database, rule Rule, ruleName string) []ValidationResult {
results := []ValidationResult{}
for _, schema := range db.Schemas {
// lowercased name -> "table" entries, one per distinct object
owners := make(map[string][]string)
display := make(map[string]string)
order := []string{}
add := func(name, table string, sameTableMerges bool) {
key := strings.ToLower(name)
if _, seen := owners[key]; !seen {
order = append(order, key)
display[key] = name
}
if sameTableMerges && contains(owners[key], table) {
return
}
owners[key] = append(owners[key], table)
}
for _, table := range schema.Tables {
for _, key := range sortedKeys(table.Indexes) {
if name := table.Indexes[key].Name; name != "" {
add(name, table.Name, false)
}
}
for _, c := range sortConstraints(table.Constraints) {
if c.Name == "" || (c.Type != models.PrimaryKeyConstraint && c.Type != models.UniqueConstraint) {
continue
}
add(c.Name, table.Name, true)
}
}
for _, key := range order {
tables := owners[key]
results = append(results, createResult(
ruleName,
len(tables) == 1,
rule.Message,
formatLocation(schema.Name, display[key], "")+" on "+strings.Join(tables, ", "),
map[string]interface{}{
"schema": schema.Name,
"index": display[key],
"tables": tables,
"occurrences": len(tables),
},
))
}
}
return results
}
+81
View File
@@ -835,3 +835,84 @@ func TestFormatLocation(t *testing.T) {
})
}
}
func TestValidateDuplicateIndexName(t *testing.T) {
db := &models.Database{
Name: "testdb",
Schemas: []*models.Schema{
{
Name: "entity",
Tables: []*models.Table{
{
Name: "actor_phone",
Indexes: map[string]*models.Index{
"idx_actor": {Name: "idx_actor", Columns: []string{"rid_actor"}},
"uk_phone": {Name: "uk_phone", Columns: []string{"phone"}, Unique: true},
},
Constraints: map[string]*models.Constraint{
// Same name as the index on the same table: one object.
"uk_phone": {Name: "uk_phone", Type: models.UniqueConstraint, Columns: []string{"phone"}},
},
},
{
Name: "actor_email",
Indexes: map[string]*models.Index{
"idx_actor": {Name: "idx_actor", Columns: []string{"rid_actor"}},
},
Constraints: map[string]*models.Constraint{
"UK_Phone": {Name: "UK_Phone", Type: models.UniqueConstraint, Columns: []string{"email"}},
},
},
{
Name: "actor_address",
Indexes: map[string]*models.Index{
"idx_actor": {Name: "idx_actor", Columns: []string{"rid_actor"}},
"idx_actor#2": {Name: "idx_actor", Columns: []string{"rid_actor", "kind"}},
"idx_address": {Name: "idx_address", Columns: []string{"line1"}},
},
},
},
},
{
// Same names in another schema do not collide.
Name: "org",
Tables: []*models.Table{
{Name: "api_provider", Indexes: map[string]*models.Index{"idx_actor": {Name: "idx_actor"}}},
},
},
},
}
results := validateDuplicateIndexName(db, Rule{Message: "dup"}, "duplicate_index_name")
got := map[string]bool{}
occ := map[string]int{}
for _, r := range results {
key := r.Context["schema"].(string) + "." + r.Context["index"].(string)
got[key] = r.Passed
occ[key] = r.Context["occurrences"].(int)
}
want := map[string]struct {
passed bool
occ int
}{
"entity.idx_actor": {false, 4},
"entity.uk_phone": {false, 2},
"entity.idx_address": {true, 1},
"org.idx_actor": {true, 1},
}
if len(got) != len(want) {
t.Fatalf("got %d results %v, want %d", len(got), got, len(want))
}
for k, w := range want {
p, ok := got[k]
if !ok {
t.Errorf("missing result for %s", k)
continue
}
if p != w.passed || occ[k] != w.occ {
t.Errorf("%s: passed=%v occurrences=%d, want passed=%v occurrences=%d", k, p, occ[k], w.passed, w.occ)
}
}
}
+42 -9
View File
@@ -19,6 +19,7 @@ import (
"fmt"
"os"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
@@ -276,15 +277,22 @@ func parseHumanSize(s string) (int64, error) {
// Options carries the subset of command flags a job file may set.
type Options struct {
FlattenSchema bool `yaml:"flatten_schema"`
Schema string `yaml:"schema"`
Package string `yaml:"package"`
FlattenSchema bool `yaml:"flatten_schema"`
Schema string `yaml:"schema"`
Package string `yaml:"package"`
// Types and ArrayNullable mirror the existing convert/split writer flags.
// They are passed through only to writers that already support them.
Types string `yaml:"types"`
ArrayNullable string `yaml:"array_nullable"`
ContinueOnError bool `yaml:"continue_on_error"`
SkipRelations bool `yaml:"skip_relations"`
SkipEnums bool `yaml:"skip_enums"`
SkipViews bool `yaml:"skip_views"`
SkipDomains bool `yaml:"skip_domains"`
SkipSequences bool `yaml:"skip_sequences"`
// FullDDL makes direct pgsql output execute the full idempotent DDL instead of
// diffing against the live database first.
FullDDL bool `yaml:"full_ddl"`
SkipRelations bool `yaml:"skip_relations"`
SkipEnums bool `yaml:"skip_enums"`
SkipViews bool `yaml:"skip_views"`
SkipDomains bool `yaml:"skip_domains"`
SkipSequences bool `yaml:"skip_sequences"`
}
// Dir returns the directory that a job's relative paths resolve against:
@@ -568,7 +576,9 @@ func (j *Job) validate() []string {
e = append(e, "command \"templ\" does not support database output")
}
if j.Output != nil && j.Output.Format != "" {
e = append(e, "output.format is not valid for command \"templ\"")
if !strings.EqualFold(j.Output.Format, "text") {
e = append(e, "command \"templ\" accepts only output.format: text")
}
}
case CommandSplit:
if len(j.Inputs) < 1 {
@@ -850,6 +860,29 @@ func SafeJoin(root, rel string) (string, error) {
return joined, nil
}
var envReferencePattern = regexp.MustCompile(`\$\{([A-Za-z_][A-Za-z0-9_]*)\}`)
// ExpandEnv expands ${NAME} references using the current process environment.
// It deliberately supports only shell-independent variable references; job
// files are never passed through a shell. An unset variable is an error so a
// typo cannot silently turn into a relative path.
func ExpandEnv(value string) (string, error) {
var missing string
expanded := envReferencePattern.ReplaceAllStringFunc(value, func(reference string) string {
name := reference[2 : len(reference)-1]
resolved, ok := os.LookupEnv(name)
if !ok {
missing = name
return reference
}
return resolved
})
if missing != "" {
return "", fmt.Errorf("environment variable %q referenced by ${%s} is not set", missing, missing)
}
return expanded, nil
}
// deepestExistingAncestor returns p itself if it exists, otherwise the nearest
// existing parent directory (falling back to the filesystem root).
func deepestExistingAncestor(p string) string {
+15
View File
@@ -135,6 +135,21 @@ func TestParseHumanSize(t *testing.T) {
}
}
func TestExpandEnv(t *testing.T) {
t.Setenv("RELSPEC_TEST_ROOT", "build/output")
got, err := ExpandEnv("${RELSPEC_TEST_ROOT}/schema-${RELSPEC_TEST_ROOT}")
if err != nil {
t.Fatal(err)
}
if want := "build/output/schema-build/output"; got != want {
t.Fatalf("ExpandEnv = %q, want %q", got, want)
}
if _, err := ExpandEnv("build/${RELSPEC_TEST_MISSING}"); err == nil {
t.Fatal("expected missing environment variable error")
}
}
func TestLoadRejectsDuplicateJobAcrossFiles(t *testing.T) {
dir := t.TempDir()
a := filepath.Join(dir, "relspec.yml")
+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)
}
}
}
+58
View File
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
} else {
// Schema doesn't exist, add it
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)
r.SchemasAdded++
}
@@ -91,6 +100,45 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
if !opts.SkipDomains {
r.mergeDomains(target, source)
}
mergeDatabaseMetadata(target, source)
}
// mergeDatabaseMetadata adds missing metadata keys and unions []string values,
// so per-file reader state (e.g. pending DBML commented refs) survives a merge.
func mergeDatabaseMetadata(target, source *models.Database) {
if len(source.Metadata) == 0 {
return
}
if target.Metadata == nil {
target.Metadata = make(map[string]any, len(source.Metadata))
}
for key, srcVal := range source.Metadata {
tgtVal, exists := target.Metadata[key]
if !exists {
if list, ok := srcVal.([]string); ok {
srcVal = append([]string(nil), list...)
}
target.Metadata[key] = srcVal
continue
}
tgtList, tgtOK := tgtVal.([]string)
srcList, srcOK := srcVal.([]string)
if !tgtOK || !srcOK {
continue
}
seen := make(map[string]bool, len(tgtList))
for _, v := range tgtList {
seen[v] = true
}
for _, v := range srcList {
if !seen[v] {
tgtList = append(tgtList, v)
seen[v] = true
}
}
target.Metadata[key] = tgtList
}
}
func (r *MergeResult) mergeSchemaContents(target, source *models.Schema, opts *MergeOptions) {
@@ -401,6 +449,8 @@ func cloneTable(table *models.Table) *models.Table {
Description: table.Description,
Schema: table.Schema,
Comment: table.Comment,
Tablespace: table.Tablespace,
GUID: table.GUID,
Sequence: table.Sequence,
UpdatedAt: table.UpdatedAt,
Columns: make(map[string]*models.Column),
@@ -430,6 +480,14 @@ func cloneTable(table *models.Table) *models.Table {
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
}
+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")
}
}
+29
View File
@@ -721,3 +721,32 @@ func TestComplexMerge(t *testing.T) {
t.Error("Expected ukey_users_guid constraint to exist")
}
}
func TestMergeDatabases_Metadata(t *testing.T) {
target := &models.Database{Metadata: map[string]any{
"refs": []string{"a", "b"},
"name": "target",
}}
source := &models.Database{Metadata: map[string]any{
"refs": []string{"b", "c"},
"name": "source",
"extra": []string{"x"},
}}
MergeDatabases(target, source, nil)
if got := target.Metadata["refs"].([]string); strings.Join(got, ",") != "a,b,c" {
t.Errorf("refs = %v, want [a b c]", got)
}
if got := target.Metadata["name"]; got != "target" {
t.Errorf("name = %v, want target (existing scalar keys are kept)", got)
}
extra := target.Metadata["extra"].([]string)
if strings.Join(extra, ",") != "x" {
t.Errorf("extra = %v, want [x]", extra)
}
source.Metadata["extra"].([]string)[0] = "changed"
if extra[0] != "x" {
t.Error("copied slice must not alias the source")
}
}
+1
View File
@@ -20,6 +20,7 @@ const (
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
MySQLDatabaseType DatabaseType = "mysql" // MySQL/MariaDB database
)
// 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")
}
}
+171
View File
@@ -0,0 +1,171 @@
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, 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)
}
})
}
}
+2
View File
@@ -64,6 +64,8 @@ The reader recognizes the following Bun struct tags:
- `autoincrement` - Auto-increment column
- `default` - Default value
- `unique` - Unique constraint
- `generated` - Generated column marker (`Generated`; expression is not in the tag)
- `identity` - Identity column marker (`Identity`, `ALWAYS`)
- `rel` - Relationship definition
## Example Bun Model
+48
View File
@@ -0,0 +1,48 @@
package bun
import (
"os"
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func TestReader_GeneratedAndIdentityMarkers(t *testing.T) {
src := `package models
import "github.com/uptrace/bun"
type Person struct {
bun.BaseModel ` + "`bun:\"table:people,alias:p\"`" + `
ID int64 ` + "`bun:\"id,pk,type:bigint,autoincrement,identity\"`" + `
Seq int64 ` + "`bun:\"seq,type:bigint,scanonly,identity,notnull\"`" + `
FullName string ` + "`bun:\"full_name,type:text,scanonly,generated,nullzero\"`" + `
First string ` + "`bun:\"first,type:text,nullzero\"`" + `
}
`
path := filepath.Join(t.TempDir(), "person.go")
if err := os.WriteFile(path, []byte(src), 0o644); err != nil {
t.Fatal(err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
cols := db.Schemas[0].Tables[0].Columns
if !cols["id"].Identity || cols["id"].IdentityGeneration != "ALWAYS" {
t.Errorf("id: identity = %v/%q, want true/ALWAYS", cols["id"].Identity, cols["id"].IdentityGeneration)
}
if !cols["seq"].Identity || cols["seq"].IdentityGeneration != "ALWAYS" {
t.Errorf("seq: identity = %v/%q, want true/ALWAYS", cols["seq"].Identity, cols["seq"].IdentityGeneration)
}
if !cols["full_name"].Generated {
t.Error("full_name should be generated")
}
if cols["first"].Generated || cols["first"].Identity {
t.Error("first must be an ordinary column")
}
}
+117
View File
@@ -0,0 +1,117 @@
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)
}
}
+7
View File
@@ -659,6 +659,13 @@ func (r *Reader) parseColumn(fieldName string, fieldType ast.Expr, tag string, s
hasExplicitNullableMarker = true
case "autoincrement":
column.AutoIncrement = true
case "generated":
// GENERATED ... STORED marker written by the Bun writer; the expression is not in the tag
column.Generated = true
case "identity":
// GENERATED ALWAYS AS IDENTITY marker written by the Bun writer
column.Identity = true
column.IdentityGeneration = "ALWAYS"
case "default":
// Default value from Bun tag (e.g., default:gen_random_uuid())
column.Default = value
+18
View File
@@ -90,11 +90,28 @@ Ref: posts.user_id > users.id [delete: cascade]
- Default values (`default`)
- Inline references (`ref`)
- Standalone `Ref` blocks
- Commented cross-file refs (`// Ref:` — see below)
- Indexes and composite indexes
- Table notes and column notes
- Enums
- Dialect directives (`@postgres:` / `@sqlite:` — see below)
## Commented cross-file refs
`// Ref:` / `// ref:` lines (ignored by dbdiagram) become FKs + relationships once both ends are loaded.
| Rule | Behaviour |
|---|---|
| When | Single file / directory: end of read. `--from-list`, `merge`, jobs: after all inputs are combined |
| Match | `schema.table.column` on both sides, case-insensitive |
| Operators | `>`, `<`, `-` (parsed like `Ref:`) |
| Duplicate of an FK on the same columns | Skipped silently |
| Target not loaded | Kept pending; skipped with a warning on the final pass |
| Column type mismatch | Warning, FK still created (`serial`≈`integer`, `bigserial`≈`bigint`) |
| Pending state | `Database.Metadata["dbml.commented_refs"]` (`[]string`) |
API: `dbml.ResolveCommentedRefs(db, final bool) []string` returns warnings.
## Dialect directives
Lines of the form `@<namespace>[(<column>)]: <args>` embed database-specific
@@ -140,6 +157,7 @@ grammar and the supported-directive matrix.
## Notes
- Column notes `GENERATED ALWAYS AS (expr) STORED` and `GENERATED ALWAYS|BY DEFAULT AS IDENTITY` set `Generated`/`GenerationExpression` and `Identity`/`IdentityGeneration`; they are not kept as comments
- DBML is designed for database documentation and diagramming
- Schema name defaults to `public`
- Relationship cardinality is preserved
+218
View File
@@ -0,0 +1,218 @@
package dbml
import (
"fmt"
"regexp"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
// CommentedRefsMetadataKey is the Database.Metadata key holding commented
// `// Ref:` lines not yet resolved against the model ([]string).
const CommentedRefsMetadataKey = "dbml.commented_refs"
// commentedRefRegex matches `// Ref: ...` and `// ref: ...`.
var commentedRefRegex = regexp.MustCompile(`^//\s*[Rr]ef\s*:\s*(.+)$`)
// commentedRef returns the ref body of a trimmed `// Ref:` comment line.
func commentedRef(line string) (string, bool) {
m := commentedRefRegex.FindStringSubmatch(line)
if m == nil {
return "", false
}
ref := strings.TrimSpace(m[1])
return ref, ref != ""
}
// PendingCommentedRefs returns the commented refs not yet resolved.
func PendingCommentedRefs(db *models.Database) []string {
if db == nil || db.Metadata == nil {
return nil
}
refs, _ := db.Metadata[CommentedRefsMetadataKey].([]string)
return refs
}
func setPendingCommentedRefs(db *models.Database, refs []string) {
if len(refs) == 0 {
if db.Metadata != nil {
delete(db.Metadata, CommentedRefsMetadataKey)
}
return
}
if db.Metadata == nil {
db.Metadata = make(map[string]any)
}
db.Metadata[CommentedRefsMetadataKey] = refs
}
// addPendingCommentedRef queues a ref, skipping exact repeats.
func addPendingCommentedRef(db *models.Database, ref string) {
refs := PendingCommentedRefs(db)
for _, existing := range refs {
if existing == ref {
return
}
}
setPendingCommentedRefs(db, append(refs, ref))
}
// ResolveCommentedRefs turns pending commented refs into foreign keys and
// relationships when both ends (schema.table.column) exist in db. A ref that
// matches an existing FK on the same columns is dropped as a duplicate.
// Unmatched refs stay pending unless final is set, in which case they are
// dropped with a warning. Returns human-readable warnings.
func ResolveCommentedRefs(db *models.Database, final bool) []string {
refs := PendingCommentedRefs(db)
if len(refs) == 0 {
return nil
}
var warnings []string
var pending []string
parser := &Reader{}
for _, ref := range refs {
fk := parser.parseRef(ref)
if fk == nil || len(fk.Columns) == 0 || len(fk.Columns) != len(fk.ReferencedColumns) {
warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: cannot parse", ref))
continue
}
srcTable, srcCols, srcMissing := lookupColumns(db, fk.Schema, fk.Table, fk.Columns)
dstTable, dstCols, dstMissing := lookupColumns(db, fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns)
if srcMissing != "" || dstMissing != "" {
if final {
missing := srcMissing
if missing == "" {
missing = dstMissing
}
warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: %s not found", ref, missing))
} else {
pending = append(pending, ref)
}
continue
}
fk.Schema, fk.Table, fk.Columns = srcTable.Schema, srcTable.Name, columnNames(srcCols)
fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns = dstTable.Schema, dstTable.Name, columnNames(dstCols)
if hasFKOnColumns(srcTable, fk.Columns) {
continue // already declared by an uncommented Ref or inline ref
}
if _, taken := srcTable.Constraints[fk.Name]; taken {
warnings = append(warnings, fmt.Sprintf("skipping commented ref %q: constraint %s already exists", ref, fk.Name))
continue
}
for i := range srcCols {
if !compatibleFKTypes(srcCols[i].Type, dstCols[i].Type) {
warnings = append(warnings, fmt.Sprintf("commented ref %q: type mismatch %s.%s.%s (%s) -> %s.%s.%s (%s)",
ref, srcTable.Schema, srcTable.Name, srcCols[i].Name, srcCols[i].Type,
dstTable.Schema, dstTable.Name, dstCols[i].Name, dstCols[i].Type))
}
}
if srcTable.Constraints == nil {
srcTable.Constraints = make(map[string]*models.Constraint)
}
srcTable.Constraints[fk.Name] = fk
addFKRelationship(srcTable, fk)
}
setPendingCommentedRefs(db, pending)
return warnings
}
// lookupColumns finds a table and its columns, case-insensitively. missing
// names the first object not found, or is empty.
func lookupColumns(db *models.Database, schemaName, tableName string, cols []string) (*models.Table, []*models.Column, string) {
qualified := schemaName + "." + tableName
var table *models.Table
for _, schema := range db.Schemas {
if !strings.EqualFold(schema.Name, schemaName) {
continue
}
for _, t := range schema.Tables {
if strings.EqualFold(t.Name, tableName) {
table = t
break
}
}
}
if table == nil {
return nil, nil, "table " + qualified
}
found := make([]*models.Column, 0, len(cols))
for _, name := range cols {
col := table.Columns[name]
if col == nil {
for _, c := range table.Columns {
if strings.EqualFold(c.Name, name) {
col = c
break
}
}
}
if col == nil {
return nil, nil, "column " + qualified + "." + name
}
found = append(found, col)
}
return table, found, ""
}
func columnNames(cols []*models.Column) []string {
names := make([]string, len(cols))
for i, c := range cols {
names[i] = c.Name
}
return names
}
// hasFKOnColumns reports whether table already has a foreign key over cols.
func hasFKOnColumns(table *models.Table, cols []string) bool {
for _, c := range table.Constraints {
if c.Type != models.ForeignKeyConstraint || len(c.Columns) != len(cols) {
continue
}
same := true
for i := range cols {
if !strings.EqualFold(c.Columns[i], cols[i]) {
same = false
break
}
}
if same {
return true
}
}
return false
}
// fkTypeAliases maps serial and alias spellings to their storage type.
var fkTypeAliases = map[string]string{
"smallserial": "smallint", "serial2": "smallint", "int2": "smallint",
"serial": "integer", "serial4": "integer", "int": "integer", "int4": "integer",
"bigserial": "bigint", "serial8": "bigint", "int8": "bigint",
}
// compatibleFKTypes compares column types ignoring case, length and serial
// vs. integer spelling. Unknown (empty) types are treated as compatible.
func compatibleFKTypes(a, b string) bool {
na, nb := normalizeFKType(a), normalizeFKType(b)
return na == "" || nb == "" || na == nb
}
func normalizeFKType(t string) string {
t = strings.ToLower(strings.TrimSpace(t))
if i := strings.Index(t, "("); i >= 0 {
t = strings.TrimSpace(t[:i])
}
if alias, ok := fkTypeAliases[t]; ok {
return alias
}
return t
}
+260
View File
@@ -0,0 +1,260 @@
package dbml
import (
"os"
"path/filepath"
"strings"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func readDBMLString(t *testing.T, content string) *models.Database {
t.Helper()
path := filepath.Join(t.TempDir(), "in.dbml")
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatalf("write fixture: %v", err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
return db
}
func findTable(db *models.Database, schema, table string) *models.Table {
for _, s := range db.Schemas {
if s.Name != schema {
continue
}
for _, t := range s.Tables {
if t.Name == table {
return t
}
}
}
return nil
}
func fkOn(table *models.Table, col string) *models.Constraint {
for _, c := range table.Constraints {
if c.Type == models.ForeignKeyConstraint && len(c.Columns) == 1 && c.Columns[0] == col {
return c
}
}
return nil
}
func relFor(table *models.Table, fkName string) *models.Relationship {
for _, r := range table.Relationships {
if r.ForeignKey == fkName {
return r
}
}
return nil
}
func TestCommentedRef(t *testing.T) {
tests := []struct {
line string
want string
ok bool
}{
{`// Ref: a.b.c > d.e.f`, `a.b.c > d.e.f`, true},
{`// ref: a.b.c - d.e.f`, `a.b.c - d.e.f`, true},
{`//Ref:a.b.c > d.e.f`, `a.b.c > d.e.f`, true},
{`// Reference notes`, "", false},
{`// see Ref: a.b.c > d.e.f`, "", false},
{`// Ref:`, "", false},
}
for _, tt := range tests {
got, ok := commentedRef(tt.line)
if ok != tt.ok || got != tt.want {
t.Errorf("commentedRef(%q) = (%q, %v), want (%q, %v)", tt.line, got, ok, tt.want, tt.ok)
}
}
}
// A commented ref whose tables are in the same file resolves on read.
func TestReader_CommentedRefSameFile(t *testing.T) {
db := readDBMLString(t, `Table "org"."department" {
"id_department" bigserial [pk]
}
Table "entity"."employee" {
"id_employee" bigserial [pk]
"rid_department" bigint
}
// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department" [delete: restrict, update: restrict]
`)
emp := findTable(db, "entity", "employee")
fk := fkOn(emp, "rid_department")
if fk == nil {
t.Fatal("expected FK on rid_department")
}
if fk.ReferencedSchema != "org" || fk.ReferencedTable != "department" || fk.ReferencedColumns[0] != "id_department" {
t.Errorf("FK target = %s.%s.%v", fk.ReferencedSchema, fk.ReferencedTable, fk.ReferencedColumns)
}
if fk.OnDelete != "restrict" || fk.OnUpdate != "restrict" {
t.Errorf("FK actions = %q/%q, want restrict/restrict", fk.OnDelete, fk.OnUpdate)
}
if relFor(emp, fk.Name) == nil {
t.Error("expected relationship for FK")
}
if refs := PendingCommentedRefs(db); len(refs) != 0 {
t.Errorf("pending = %v, want none", refs)
}
}
// A cross-file commented ref stays pending after a single-file read.
func TestReader_CommentedRefCrossFilePending(t *testing.T) {
db := readDBMLString(t, `Table "entity"."employee" {
"id_employee" bigserial [pk]
"rid_department" bigint
}
// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department"
`)
if fk := fkOn(findTable(db, "entity", "employee"), "rid_department"); fk != nil {
t.Fatal("FK must not resolve without the target table")
}
if refs := PendingCommentedRefs(db); len(refs) != 1 {
t.Fatalf("pending = %v, want 1", refs)
}
}
// Directory reads resolve commented refs after all files are merged.
func TestReader_CommentedRefDirectory(t *testing.T) {
db, err := NewReader(&readers.ReaderOptions{
FilePath: filepath.Join("..", "..", "..", "tests", "assets", "dbml", "multifile"),
}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
fk := fkOn(findTable(db, "public", "posts"), "user_id")
if fk == nil {
t.Fatal("expected FK posts.user_id from 9_refs.dbml commented ref")
}
if fk.ReferencedTable != "users" || fk.OnDelete != "CASCADE" {
t.Errorf("FK = %s ondelete %s, want users ondelete CASCADE", fk.ReferencedTable, fk.OnDelete)
}
}
func crossFileDB(t *testing.T, refs ...string) *models.Database {
t.Helper()
db := readDBMLString(t, `Table "org"."department" {
"id_department" bigserial [pk]
}
Table "entity"."employee" {
"id_employee" bigserial [pk]
"rid_department" bigint
"rid_manager" integer
"rid_team" bigint
}
`)
setPendingCommentedRefs(db, refs)
return db
}
func TestResolveCommentedRefs(t *testing.T) {
t.Run("lowercase ref and one-to-one", func(t *testing.T) {
db := crossFileDB(t,
`"entity"."employee"."rid_department" > "org"."department"."id_department"`,
`entity.employee.rid_team - org.department.id_department`,
)
if w := ResolveCommentedRefs(db, true); len(w) != 0 {
t.Errorf("warnings = %v, want none", w)
}
emp := findTable(db, "entity", "employee")
for _, col := range []string{"rid_department", "rid_team"} {
fk := fkOn(emp, col)
if fk == nil {
t.Fatalf("expected FK on %s", col)
}
if relFor(emp, fk.Name) == nil {
t.Errorf("expected relationship for %s", fk.Name)
}
}
if len(emp.Relationships) != 2 {
t.Errorf("relationships = %d, want 2 (same target must not overwrite)", len(emp.Relationships))
}
})
t.Run("type mismatch warns but resolves", func(t *testing.T) {
db := crossFileDB(t, `entity.employee.rid_manager > org.department.id_department`)
w := ResolveCommentedRefs(db, true)
if len(w) != 1 || !strings.Contains(w[0], "type mismatch") {
t.Errorf("warnings = %v, want one type mismatch", w)
}
if fkOn(findTable(db, "entity", "employee"), "rid_manager") == nil {
t.Error("expected FK on rid_manager")
}
})
t.Run("missing target stays pending until final", func(t *testing.T) {
db := crossFileDB(t, `entity.employee.rid_team > hr.team.id_team`)
if w := ResolveCommentedRefs(db, false); len(w) != 0 {
t.Errorf("non-final warnings = %v, want none", w)
}
if len(PendingCommentedRefs(db)) != 1 {
t.Fatal("ref should stay pending")
}
w := ResolveCommentedRefs(db, true)
if len(w) != 1 || !strings.Contains(w[0], "table hr.team not found") {
t.Errorf("final warnings = %v, want missing table", w)
}
if len(PendingCommentedRefs(db)) != 0 {
t.Error("final pass must clear pending refs")
}
})
t.Run("missing column", func(t *testing.T) {
db := crossFileDB(t, `entity.employee.rid_nope > org.department.id_department`)
w := ResolveCommentedRefs(db, true)
if len(w) != 1 || !strings.Contains(w[0], "column entity.employee.rid_nope not found") {
t.Errorf("warnings = %v, want missing column", w)
}
})
t.Run("deduplicates against uncommented ref", func(t *testing.T) {
db := readDBMLString(t, `Table "org"."department" {
"id_department" bigserial [pk]
}
Table "entity"."employee" {
"id_employee" bigserial [pk]
"rid_department" bigint
}
Ref: "entity"."employee"."rid_department" > "org"."department"."id_department"
// Ref: "entity"."employee"."rid_department" > "org"."department"."id_department"
`)
emp := findTable(db, "entity", "employee")
count := 0
for _, c := range emp.Constraints {
if c.Type == models.ForeignKeyConstraint {
count++
}
}
if count != 1 || len(emp.Relationships) != 1 {
t.Errorf("FKs = %d, relationships = %d, want 1 and 1", count, len(emp.Relationships))
}
})
}
func TestCompatibleFKTypes(t *testing.T) {
tests := []struct {
a, b string
want bool
}{
{"bigint", "bigserial", true},
{"integer", "serial", true},
{"INT4", "integer", true},
{"varchar(10)", "varchar(20)", true},
{"bigint", "serial", false},
{"uuid", "bigint", false},
{"", "bigint", true},
}
for _, tt := range tests {
if got := compatibleFKTypes(tt.a, tt.b); got != tt.want {
t.Errorf("compatibleFKTypes(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.want)
}
}
}
+38
View File
@@ -0,0 +1,38 @@
package dbml
import (
"regexp"
"strings"
"git.warky.dev/wdevs/relspecgo/pkg/models"
)
var (
generatedNoteRegex = regexp.MustCompile(`(?is)^GENERATED\s+ALWAYS\s+AS\s*\((.*)\)\s*STORED$`)
identityNoteRegex = regexp.MustCompile(`(?i)^GENERATED\s+(ALWAYS|BY\s+DEFAULT)\s+AS\s+IDENTITY$`)
)
// applyColumnMarkerNote interprets a column note written by the DBML writer for
// GENERATED ... STORED and GENERATED ... AS IDENTITY columns, setting the matching
// model fields. It returns true when the note was a marker, in which case it is not
// a user comment.
func applyColumnMarkerNote(column *models.Column, rawNote string) bool {
note := strings.TrimSpace(rawNote)
if len(note) >= 2 && (note[0] == '\'' || note[0] == '"') && note[len(note)-1] == note[0] {
note = note[1 : len(note)-1]
}
note = strings.ReplaceAll(strings.ReplaceAll(note, `\'`, `'`), `\\`, `\`)
note = strings.TrimSpace(note)
if m := generatedNoteRegex.FindStringSubmatch(note); m != nil {
column.Generated = true
column.GenerationExpression = strings.TrimSpace(m[1])
return true
}
if m := identityNoteRegex.FindStringSubmatch(note); m != nil {
column.Identity = true
column.IdentityGeneration = strings.Join(strings.Fields(strings.ToUpper(m[1])), " ")
return true
}
return false
}
+72 -15
View File
@@ -48,7 +48,22 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
return nil, fmt.Errorf("failed to read file: %w", err)
}
return r.parseDBML(string(content))
db, err := r.parseDBML(string(content))
if err != nil {
return nil, err
}
r.resolveCommentedRefs(db)
return db, nil
}
// resolveCommentedRefs resolves the commented refs whose tables are loaded.
// Unmatched refs stay pending for a later pass over a combined model.
func (r *Reader) resolveCommentedRefs(db *models.Database) {
for _, w := range ResolveCommentedRefs(db, false) {
if r.options.Progress != nil {
r.options.Progress("warning: " + w)
}
}
}
// ReadSchema reads and parses DBML input, returning a Schema model
@@ -125,6 +140,7 @@ func (r *Reader) readDirectoryDBML(dirPath string) (*models.Database, error) {
}
}
r.resolveCommentedRefs(db)
return db, nil
}
@@ -440,6 +456,10 @@ func mergeDatabase(baseDB, fileDB *models.Database) {
// Merge domains
baseDB.Domains = append(baseDB.Domains, fileDB.Domains...)
for _, ref := range PendingCommentedRefs(fileDB) {
addPendingCommentedRef(baseDB, ref)
}
// Use first non-empty description
if baseDB.Description == "" && fileDB.Description != "" {
baseDB.Description = fileDB.Description
@@ -493,8 +513,12 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
continue
}
// Skip empty lines and comments
// Skip empty lines and comments. A commented `// Ref:` is kept as a
// pending cross-file ref, resolved once every file is loaded.
if line == "" || strings.HasPrefix(line, "//") {
if ref, ok := commentedRef(line); ok {
addPendingCommentedRef(db, ref)
}
continue
}
@@ -581,7 +605,13 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
index := r.parseIndex(line, currentTable.Name, currentSchema)
if index != nil {
currentTable.Indexes[index.Name] = index
// Keep a reused name under a unique map key so the duplicate is
// not silently dropped; the inspector reports it.
key := index.Name
for n := 2; currentTable.Indexes[key] != nil; n++ {
key = fmt.Sprintf("%s#%d", index.Name, n)
}
currentTable.Indexes[key] = index
lastIndex = index
}
continue
@@ -659,20 +689,12 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
// for DBML refs so diffing equivalent schemas compares the same model.
for _, schema := range schemaMap {
for _, table := range schema.Tables {
for _, constraint := range table.Constraints {
for _, name := range sortedConstraintNames(table.Constraints) {
constraint := table.Constraints[name]
if constraint.Type != models.ForeignKeyConstraint {
continue
}
name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable)
relationship := models.InitRelationship(name, models.OneToMany)
relationship.FromTable = table.Name
relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), constraint.Columns...)
relationship.ToTable = constraint.ReferencedTable
relationship.ToSchema = constraint.ReferencedSchema
relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...)
relationship.ForeignKey = constraint.Name
table.Relationships[name] = relationship
addFKRelationship(table, constraint)
}
}
}
@@ -685,6 +707,39 @@ func (r *Reader) parseDBML(content string) (*models.Database, error) {
return db, nil
}
// sortedConstraintNames returns constraint keys in sorted order so derived
// relationship names do not depend on map iteration order.
func sortedConstraintNames(constraints map[string]*models.Constraint) []string {
names := make([]string, 0, len(constraints))
for name := range constraints {
names = append(names, name)
}
sort.Strings(names)
return names
}
// addFKRelationship derives the relationship for a foreign key, matching how
// the PostgreSQL reader models FKs.
func addFKRelationship(table *models.Table, constraint *models.Constraint) {
name := fmt.Sprintf("%s_to_%s", table.Name, constraint.ReferencedTable)
if existing, taken := table.Relationships[name]; taken && existing.ForeignKey != constraint.Name {
// A second FK to the same table must not overwrite the first.
name = fmt.Sprintf("%s_%s", name, strings.Join(constraint.Columns, "_"))
}
relationship := models.InitRelationship(name, models.OneToMany)
relationship.FromTable = table.Name
relationship.FromSchema = table.Schema
relationship.FromColumns = append([]string(nil), constraint.Columns...)
relationship.ToTable = constraint.ReferencedTable
relationship.ToSchema = constraint.ReferencedSchema
relationship.ToColumns = append([]string(nil), constraint.ReferencedColumns...)
relationship.ForeignKey = constraint.Name
if table.Relationships == nil {
table.Relationships = make(map[string]*models.Relationship)
}
table.Relationships[name] = relationship
}
// setTableNote preserves multiple table notes. The first maps to Description
// and the second to Comment, matching the model fields used by code writers.
func setTableNote(table *models.Table, note string) {
@@ -749,7 +804,9 @@ func (r *Reader) parseColumn(line, tableName, schemaName string) (*models.Column
} else if strings.HasPrefix(attr, "note:") {
// Parse column note/comment
note := strings.TrimSpace(strings.TrimPrefix(attr, "note:"))
column.Comment = strings.Trim(note, "'\"")
if !applyColumnMarkerNote(column, note) {
column.Comment = strings.Trim(note, "'\"")
}
} else if strings.HasPrefix(attr, "ref:") {
// Parse inline reference
// DBML semantics depend on context:
+36
View File
@@ -1085,3 +1085,39 @@ func TestReader_MultilineTableNote(t *testing.T) {
t.Errorf("column note = %q, want %q", got, want)
}
}
// A name reused inside one Indexes block must not silently drop the earlier
// index; both are kept so the inspector can report the duplicate.
func TestReader_DuplicateIndexNameInTableKept(t *testing.T) {
dbmlContent := `Table individual_actor {
id bigint [pk]
rid_actor bigint
kind text
Indexes {
(rid_actor) [name: 'idx_actor', unique]
(rid_actor, kind) [name: 'idx_actor']
}
}
`
dir := t.TempDir()
path := filepath.Join(dir, "dup_index.dbml")
if err := os.WriteFile(path, []byte(dbmlContent), 0o644); err != nil {
t.Fatalf("failed to write fixture: %v", err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
table := db.Schemas[0].Tables[0]
if len(table.Indexes) != 2 {
t.Fatalf("expected 2 indexes, got %d: %v", len(table.Indexes), table.Indexes)
}
for key, idx := range table.Indexes {
if idx.Name != "idx_actor" {
t.Errorf("index %q: Name = %q, want idx_actor", key, idx.Name)
}
}
}
+1
View File
@@ -85,6 +85,7 @@ export const postsRelations = relations(posts, ({ one }) => ({
## Notes
- `.generatedAlwaysAs(sql`expr`)` sets `Generated` + `GenerationExpression`; `.generatedAlwaysAsIdentity()` / `.generatedByDefaultAsIdentity()` set `Identity` + `IdentityGeneration`
- Supports both PostgreSQL and MySQL Drizzle schemas
- Extracts relationship information from `relations` definitions
- Schema defaults to `public` for PostgreSQL
+74 -1
View File
@@ -15,6 +15,9 @@ import (
// Reader implements the readers.Reader interface for Drizzle schema format
type Reader struct {
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
@@ -29,6 +32,7 @@ func (r *Reader) ReadDatabase() (*models.Database, error) {
if r.options.FilePath == "" {
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
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)
}
// 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
for _, file := range files {
content, err := os.ReadFile(file)
@@ -125,9 +136,22 @@ func (r *Reader) readDirectory(dirPath string) (*models.Database, error) {
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
func (r *Reader) parseDrizzle(content string) (*models.Database, error) {
db := models.InitDatabase("database")
r.collectEnumVars(content)
if r.options.Metadata != nil {
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
column.Type = r.drizzleTypeToSQL(drizzleType)
if enumName, ok := r.enumVars[drizzleType]; ok {
column.Type = enumName
}
// Default: columns are nullable unless specified
column.NotNull = false
@@ -495,9 +522,24 @@ func (r *Reader) parseColumnModifiers(line string, column *models.Column, table
}
}
// Check for .generatedAlwaysAsIdentity()
// Check for .generatedAlwaysAsIdentity() / .generatedByDefaultAsIdentity()
if strings.Contains(line, ".generatedAlwaysAsIdentity()") {
column.AutoIncrement = true
column.Identity = true
column.IdentityGeneration = "ALWAYS"
}
if strings.Contains(line, ".generatedByDefaultAsIdentity()") {
column.AutoIncrement = true
column.Identity = true
column.IdentityGeneration = "BY DEFAULT"
}
// Check for .generatedAlwaysAs(sql`expr`) (generated column)
if idx := strings.Index(line, ".generatedAlwaysAs("); idx != -1 {
if expr, ok := parseGeneratedAlwaysAs(line[idx+len(".generatedAlwaysAs("):]); ok {
column.Generated = true
column.GenerationExpression = expr
}
}
// Check for .references(() => otherTable.column)
@@ -615,3 +657,34 @@ func (r *Reader) varNameToTableName(varName string) string {
// For now, assume variable name matches table name
return varName
}
// parseGeneratedAlwaysAs extracts the SQL expression from the argument list of
// generatedAlwaysAs(...), which follows the opening parenthesis in rest. The expression is
// the contents of a sql`...` template (backticks inside it are escaped with a backslash)
// or, failing that, a quoted string.
func parseGeneratedAlwaysAs(rest string) (string, bool) {
rest = strings.TrimSpace(rest)
rest = strings.TrimPrefix(rest, "sql")
if rest == "" {
return "", false
}
quote := rest[0]
if quote != '`' && quote != '\'' && quote != '"' {
return "", false
}
var sb strings.Builder
for i := 1; i < len(rest); i++ {
ch := rest[i]
if ch == '\\' && i+1 < len(rest) {
i++
sb.WriteByte(rest[i])
continue
}
if ch == quote {
return strings.TrimSpace(sb.String()), true
}
sb.WriteByte(ch)
}
return "", 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")
}
}
+2
View File
@@ -78,6 +78,8 @@ The reader recognizes the following GORM struct tags:
- `not null` - NOT NULL constraint
- `autoIncrement` - Auto-increment column
- `default` - Default value
- `generated` - Generated column marker (`Generated`; expression is not in the tag)
- `identity` - Identity column marker (`Identity`, `ALWAYS`)
- `size` - Column size/length
- `index` - Create index
- `uniqueIndex` - Create unique index
+48
View File
@@ -0,0 +1,48 @@
package gorm
import (
"os"
"path/filepath"
"testing"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func TestReader_GeneratedAndIdentityMarkers(t *testing.T) {
src := `package models
type Person struct {
ID int64 ` + "`gorm:\"column:id;primaryKey;autoIncrement;type:bigint;identity\"`" + `
Seq int64 ` + "`gorm:\"column:seq;type:bigint;<-:false;identity;not null\"`" + `
FullName string ` + "`gorm:\"column:full_name;type:text;<-:false;generated\"`" + `
First string ` + "`gorm:\"column:first;type:text\"`" + `
}
func (Person) TableName() string {
return "people"
}
`
path := filepath.Join(t.TempDir(), "person.go")
if err := os.WriteFile(path, []byte(src), 0o644); err != nil {
t.Fatal(err)
}
db, err := NewReader(&readers.ReaderOptions{FilePath: path}).ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase() error = %v", err)
}
cols := db.Schemas[0].Tables[0].Columns
if !cols["id"].Identity || cols["id"].IdentityGeneration != "ALWAYS" {
t.Errorf("id: identity = %v/%q, want true/ALWAYS", cols["id"].Identity, cols["id"].IdentityGeneration)
}
if !cols["seq"].Identity {
t.Error("seq should be identity")
}
if !cols["full_name"].Generated {
t.Error("full_name should be generated")
}
if cols["first"].Generated || cols["first"].Identity {
t.Error("first must be an ordinary column")
}
}
+164
View File
@@ -0,0 +1,164 @@
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)
}
})
}
}
+9
View File
@@ -726,6 +726,15 @@ func (r *Reader) parseColumn(fieldName string, fieldType ast.Expr, tag string, s
if _, ok := parts["autoincrement"]; ok {
column.AutoIncrement = true
}
if _, ok := parts["generated"]; ok {
// GENERATED ... STORED marker written by the GORM writer; the expression is not in the tag
column.Generated = true
}
if _, ok := parts["identity"]; ok {
// GENERATED ALWAYS AS IDENTITY marker written by the GORM writer
column.Identity = true
column.IdentityGeneration = "ALWAYS"
}
if def, ok := parts["default"]; ok {
// Default value from GORM tag (e.g., default:gen_random_uuid())
column.Default = def
+5
View File
@@ -25,6 +25,11 @@ sqlserver://user:pass@192.168.1.100:1433/production
sqlserver://localhost/testdb?encrypt=disable
```
## Computed and Identity Columns
- `sys.computed_columns` -> `Generated` + `GenerationExpression` (outer parentheses removed; persisted or not)
- Identity columns -> `Identity` + `IdentityGeneration = ALWAYS`
## Supported Constraints
- Primary Keys
+37 -2
View File
@@ -104,8 +104,12 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
c.numeric_precision,
c.numeric_scale,
ISNULL(ep.value, '') as description,
COLUMNPROPERTY(OBJECT_ID(QUOTENAME(c.table_schema) + '.' + QUOTENAME(c.table_name)), c.column_name, 'IsIdentity') as is_identity
COLUMNPROPERTY(OBJECT_ID(QUOTENAME(c.table_schema) + '.' + QUOTENAME(c.table_name)), c.column_name, 'IsIdentity') as is_identity,
cc.definition as computed_definition
FROM information_schema.columns c
LEFT JOIN sys.computed_columns cc
ON cc.object_id = OBJECT_ID(QUOTENAME(c.table_schema) + '.' + QUOTENAME(c.table_name))
AND cc.name = c.column_name
LEFT JOIN sys.extended_properties ep
ON ep.major_id = OBJECT_ID(QUOTENAME(c.table_schema) + '.' + QUOTENAME(c.table_name))
AND ep.minor_id = COLUMNPROPERTY(OBJECT_ID(QUOTENAME(c.table_schema) + '.' + QUOTENAME(c.table_name)), c.column_name, 'ColumnId')
@@ -127,8 +131,9 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
var schema, tableName, columnName, isNullable, dataType, description string
var ordinalPosition int
var columnDefault, charMaxLength, numPrecision, numScale, isIdentity *int
var computedDefinition *string
if err := rows.Scan(&schema, &tableName, &columnName, &ordinalPosition, &columnDefault, &isNullable, &dataType, &charMaxLength, &numPrecision, &numScale, &description, &isIdentity); err != nil {
if err := rows.Scan(&schema, &tableName, &columnName, &ordinalPosition, &columnDefault, &isNullable, &dataType, &charMaxLength, &numPrecision, &numScale, &description, &isIdentity, &computedDefinition); err != nil {
return nil, err
}
@@ -144,6 +149,14 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
// Check if this is an identity column (auto-increment)
if isIdentity != nil && *isIdentity == 1 {
column.AutoIncrement = true
column.Identity = true
column.IdentityGeneration = "ALWAYS"
}
// Computed columns report their expression wrapped in an extra pair of parentheses
if computedDefinition != nil && strings.TrimSpace(*computedDefinition) != "" {
column.Generated = true
column.GenerationExpression = trimOuterParens(*computedDefinition)
}
if charMaxLength != nil && *charMaxLength > 0 {
@@ -414,3 +427,25 @@ func (r *Reader) queryIndexes(schemaName string) (map[string][]*models.Index, er
return indexes, rows.Err()
}
// trimOuterParens removes one pair of parentheses wrapping the whole expression, as
// SQL Server stores computed column definitions.
func trimOuterParens(expr string) string {
expr = strings.TrimSpace(expr)
if len(expr) < 2 || expr[0] != '(' || expr[len(expr)-1] != ')' {
return expr
}
depth := 0
for i := 0; i < len(expr); i++ {
switch expr[i] {
case '(':
depth++
case ')':
depth--
if depth == 0 && i != len(expr)-1 {
return expr
}
}
}
return strings.TrimSpace(expr[1 : len(expr)-1])
}
+15
View File
@@ -85,3 +85,18 @@ func TestConvertMSSQLToCanonical(t *testing.T) {
})
}
}
func TestTrimOuterParens(t *testing.T) {
tests := map[string]string{
"(([a])+([b]))": "([a])+([b])",
"([a]+[b])": "[a]+[b]",
"([a])+([b])": "([a])+([b])",
"[a]+[b]": "[a]+[b]",
" (upper([a])) ": "upper([a])",
}
for in, want := range tests {
if got := trimOuterParens(in); got != want {
t.Errorf("trimOuterParens(%q) = %q, want %q", in, got, want)
}
}
}
+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")
}
}
+27 -23
View File
@@ -422,31 +422,35 @@ func (r *Reader) queryPrimaryKeys(schemaName string) (map[string]*models.Constra
// queryForeignKeys retrieves all foreign key constraints for a schema
// Returns map[schema.table][]*Constraint
func (r *Reader) queryForeignKeys(schemaName string) (map[string][]*models.Constraint, error) {
// Columns are paired by position from pg_constraint so composite keys are not
// cross-joined (information_schema.constraint_column_usage has no ordering).
actionCase := func(col string) string {
return `CASE ` + col + ` WHEN 'a' THEN 'NO ACTION' WHEN 'r' THEN 'RESTRICT' WHEN 'c' THEN 'CASCADE' WHEN 'n' THEN 'SET NULL' WHEN 'd' THEN 'SET DEFAULT' END`
}
query := `
SELECT
tc.table_schema,
tc.table_name,
tc.constraint_name,
kcu.table_schema as foreign_table_schema,
kcu.table_name as foreign_table_name,
kcu.column_name as foreign_column,
ccu.table_schema as referenced_table_schema,
ccu.table_name as referenced_table_name,
ccu.column_name as referenced_column,
rc.update_rule,
rc.delete_rule
FROM information_schema.table_constraints tc
JOIN information_schema.key_column_usage kcu
ON tc.constraint_name = kcu.constraint_name
AND tc.table_schema = kcu.table_schema
JOIN information_schema.constraint_column_usage ccu
ON ccu.constraint_name = tc.constraint_name
JOIN information_schema.referential_constraints rc
ON rc.constraint_name = tc.constraint_name
AND rc.constraint_schema = tc.table_schema
WHERE tc.constraint_type = 'FOREIGN KEY'
AND tc.table_schema = $1
ORDER BY tc.table_schema, tc.table_name, tc.constraint_name, kcu.ordinal_position
ns.nspname AS table_schema,
cl.relname AS table_name,
con.conname AS constraint_name,
ns.nspname AS foreign_table_schema,
cl.relname AS foreign_table_name,
att.attname AS foreign_column,
fns.nspname AS referenced_table_schema,
fcl.relname AS referenced_table_name,
fatt.attname AS referenced_column,
` + actionCase("con.confupdtype") + ` AS update_rule,
` + actionCase("con.confdeltype") + ` AS delete_rule
FROM pg_catalog.pg_constraint con
JOIN pg_catalog.pg_class cl ON cl.oid = con.conrelid
JOIN pg_catalog.pg_namespace ns ON ns.oid = cl.relnamespace
JOIN pg_catalog.pg_class fcl ON fcl.oid = con.confrelid
JOIN pg_catalog.pg_namespace fns ON fns.oid = fcl.relnamespace
CROSS JOIN LATERAL unnest(con.conkey, con.confkey) WITH ORDINALITY AS k(attnum, fattnum, ord)
JOIN pg_catalog.pg_attribute att ON att.attrelid = con.conrelid AND att.attnum = k.attnum
JOIN pg_catalog.pg_attribute fatt ON fatt.attrelid = con.confrelid AND fatt.attnum = k.fattnum
WHERE con.contype = 'f'
AND ns.nspname = $1
ORDER BY ns.nspname, cl.relname, con.conname, k.ord
`
rows, err := r.conn.Query(r.ctx, query, schemaName)
+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)
}
}
}
+61
View File
@@ -1,10 +1,14 @@
package pgsql
import (
"context"
"os"
"reflect"
"strings"
"testing"
"github.com/jackc/pgx/v5"
"git.warky.dev/wdevs/relspecgo/pkg/models"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
@@ -499,3 +503,60 @@ func TestMapDataType_ExtensionTypesPreserveModifiers(t *testing.T) {
})
}
}
func TestReader_CompositeForeignKeyColumnsArePairedOnce(t *testing.T) {
connStr := getTestConnectionString(t)
ctx := context.Background()
conn, err := pgx.Connect(ctx, connStr)
if err != nil {
t.Fatalf("connect: %v", err)
}
defer conn.Close(ctx)
const schema = "relspec_fk_test"
setup := []string{
"DROP SCHEMA IF EXISTS " + schema + " CASCADE",
"CREATE SCHEMA " + schema,
"CREATE TABLE " + schema + ".parent (a int, b int, PRIMARY KEY (a, b))",
"CREATE TABLE " + schema + ".child (x int, y int, CONSTRAINT fk_child_parent FOREIGN KEY (x, y) REFERENCES " + schema + ".parent (a, b) ON DELETE CASCADE)",
}
for _, stmt := range setup {
if _, err := conn.Exec(ctx, stmt); err != nil {
t.Fatalf("setup %q: %v", stmt, err)
}
}
defer conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+schema+" CASCADE")
reader := NewReader(&readers.ReaderOptions{ConnectionString: connStr})
db, err := reader.ReadDatabase()
if err != nil {
t.Fatalf("ReadDatabase: %v", err)
}
for _, s := range db.Schemas {
if s.Name != schema {
continue
}
for _, tbl := range s.Tables {
if tbl.Name != "child" {
continue
}
fk := tbl.Constraints["fk_child_parent"]
if fk == nil {
t.Fatal("foreign key fk_child_parent not read")
}
if got := strings.Join(fk.Columns, ","); got != "x,y" {
t.Errorf("columns = %q, want x,y", got)
}
if got := strings.Join(fk.ReferencedColumns, ","); got != "a,b" {
t.Errorf("referenced columns = %q, want a,b", got)
}
if fk.OnDelete != "CASCADE" || fk.OnUpdate != "NO ACTION" {
t.Errorf("rules = %s/%s, want CASCADE/NO ACTION", fk.OnDelete, fk.OnUpdate)
}
return
}
}
t.Fatal("test schema/table not found in read result")
}
+18 -14
View File
@@ -13,7 +13,8 @@ import (
// Reader implements the readers.Reader interface for Prisma schema format
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
@@ -82,6 +83,8 @@ func (r *Reader) parsePrisma(content string) (*models.Database, error) {
schema := models.InitSchema("public")
schema.Enums = make([]*models.Enum, 0)
r.enumNames = collectEnumNames(content)
scanner := bufio.NewScanner(strings.NewReader(content))
// State tracking
@@ -600,20 +603,21 @@ func (r *Reader) isPrimitiveType(typeName string) bool {
return false
}
// isEnumType checks if a type name might be an enum
// Note: We can't definitively check against schema.Enums at parse time
// because enums might be defined after the model, so we just check
// if it starts with uppercase (Prisma convention for enums)
func (r *Reader) isEnumType(typeName string, table *models.Table) bool {
// 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' {
// Additional check: primitive types are already handled above
// So if it's uppercase and not primitive, it's likely an enum or model
// We'll assume it's an enum if it's a single word
return !strings.Contains(typeName, "_")
// isEnumType reports whether typeName is an enum declared in the schema.
// Enum names are collected up front because enums may be declared after the
// models that use them.
func (r *Reader) isEnumType(typeName string, _ *models.Table) bool {
return r.enumNames[typeName]
}
var enumDeclRegex = regexp.MustCompile(`(?m)^\s*enum\s+(\w+)\s*{`)
func collectEnumNames(content string) map[string]bool {
names := make(map[string]bool)
for _, m := range enumDeclRegex.FindAllStringSubmatch(content, -1) {
names[m[1]] = true
}
return false
return names
}
// createConstraintFromRelation creates a FK constraint from a @relation attribute
+351
View File
@@ -0,0 +1,351 @@
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)
}
}
}
+1
View File
@@ -40,6 +40,7 @@ options := &readers.ReaderOptions{
- Uses pure Go driver (modernc.org/sqlite) - no CGo required
- Supports both file path and connection string
- Auto-increment detection for INTEGER PRIMARY KEY columns
- Generated columns (virtual and stored) are read via `PRAGMA table_xinfo`; the expression is parsed from the `CREATE TABLE` SQL
- Foreign keys require `PRAGMA foreign_keys = ON` to be set
## Example Schema
+116
View File
@@ -0,0 +1,116 @@
package sqlite
import (
"regexp"
"strings"
)
var generatedAsRegex = regexp.MustCompile(`(?is)\bAS\s*\(`)
// parseGeneratedExpression extracts the expression of a generated column from a CREATE
// TABLE statement, e.g. `full TEXT GENERATED ALWAYS AS (a || b) STORED` yields `a || b`.
// It returns "" when the column or its expression cannot be found.
func parseGeneratedExpression(createSQL, columnName string) string {
open := strings.Index(createSQL, "(")
if open < 0 {
return ""
}
for _, def := range splitTopLevel(createSQL[open+1:]) {
if !strings.EqualFold(firstIdentifier(def), columnName) {
continue
}
loc := generatedAsRegex.FindStringIndex(def)
if loc == nil {
return ""
}
return balancedParens(def[loc[1]-1:])
}
return ""
}
// splitTopLevel splits a table body on commas that are outside parentheses and quotes,
// stopping at the parenthesis that closes the body.
func splitTopLevel(body string) []string {
var parts []string
depth := 0
var quote byte
start := 0
for i := 0; i < len(body); i++ {
ch := body[i]
if quote != 0 {
if ch == quote {
quote = 0
}
continue
}
switch ch {
case '\'', '"', '`':
quote = ch
case '[':
quote = ']'
case '(':
depth++
case ')':
if depth == 0 {
return append(parts, body[start:i])
}
depth--
case ',':
if depth == 0 {
parts = append(parts, body[start:i])
start = i + 1
}
}
}
return append(parts, body[start:])
}
// firstIdentifier returns the first (possibly quoted) identifier of a column definition.
func firstIdentifier(def string) string {
def = strings.TrimSpace(def)
if def == "" {
return ""
}
switch def[0] {
case '"', '\'', '`':
if end := strings.IndexByte(def[1:], def[0]); end >= 0 {
return def[1 : 1+end]
}
case '[':
if end := strings.IndexByte(def, ']'); end >= 0 {
return def[1:end]
}
}
if end := strings.IndexAny(def, " \t\r\n"); end >= 0 {
return def[:end]
}
return def
}
// balancedParens returns the text inside the parenthesis group that starts at s[0].
func balancedParens(s string) string {
depth := 0
var quote byte
for i := 0; i < len(s); i++ {
ch := s[i]
if quote != 0 {
if ch == quote {
quote = 0
}
continue
}
switch ch {
case '\'', '"':
quote = ch
case '(':
depth++
case ')':
depth--
if depth == 0 {
return strings.TrimSpace(s[1:i])
}
}
}
return ""
}
+64
View File
@@ -0,0 +1,64 @@
package sqlite
import (
"database/sql"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"git.warky.dev/wdevs/relspecgo/pkg/readers"
)
func TestParseGeneratedExpression(t *testing.T) {
createSQL := `CREATE TABLE people (
id INTEGER PRIMARY KEY,
"first" TEXT,
last TEXT,
full_name TEXT GENERATED ALWAYS AS (coalesce("first", '') || ' ' || last) STORED,
initials TEXT AS (substr("first", 1, 1) || substr(last, 1, 1)),
plain TEXT NOT NULL
)`
tests := []struct {
column string
want string
}{
{"full_name", `coalesce("first", '') || ' ' || last`},
{"initials", `substr("first", 1, 1) || substr(last, 1, 1)`},
{"plain", ""},
{"missing", ""},
}
for _, tt := range tests {
t.Run(tt.column, func(t *testing.T) {
assert.Equal(t, tt.want, parseGeneratedExpression(createSQL, tt.column))
})
}
}
func TestReader_GeneratedColumns(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "gen.db")
db, err := sql.Open("sqlite", dbPath)
require.NoError(t, err)
_, err = db.Exec(`CREATE TABLE people (
id INTEGER PRIMARY KEY,
first TEXT,
last TEXT,
full_name TEXT GENERATED ALWAYS AS (first || ' ' || last) STORED,
initials TEXT AS (substr(first, 1, 1) || substr(last, 1, 1))
)`)
require.NoError(t, err)
require.NoError(t, db.Close())
got, err := NewReader(&readers.ReaderOptions{FilePath: dbPath}).ReadDatabase()
require.NoError(t, err)
cols := got.Schemas[0].Tables[0].Columns
require.Contains(t, cols, "full_name", "generated columns must be read")
assert.True(t, cols["full_name"].Generated)
assert.Equal(t, "first || ' ' || last", cols["full_name"].GenerationExpression)
assert.True(t, cols["initials"].Generated)
assert.Equal(t, "substr(first, 1, 1) || substr(last, 1, 1)", cols["initials"].GenerationExpression)
assert.False(t, cols["first"].Generated)
}
+29 -4
View File
@@ -75,7 +75,8 @@ func (r *Reader) queryViews() ([]*models.View, error) {
// queryColumns retrieves all columns for a given table or view
func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, error) {
query := fmt.Sprintf("PRAGMA table_info(%s)", tableName)
// table_xinfo, unlike table_info, also lists generated columns (hidden = 2 virtual, 3 stored)
query := fmt.Sprintf("PRAGMA table_xinfo(%s)", tableName)
rows, err := r.db.QueryContext(r.ctx, query)
if err != nil {
@@ -84,24 +85,38 @@ func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, erro
defer rows.Close()
columns := make(map[string]*models.Column)
var tableSQL string
tableSQLLoaded := false
for rows.Next() {
var cid int
var name, dataType string
var notNull, pk int
var notNull, pk, hidden int
var defaultValue *string
if err := rows.Scan(&cid, &name, &dataType, &notNull, &defaultValue, &pk); err != nil {
if err := rows.Scan(&cid, &name, &dataType, &notNull, &defaultValue, &pk, &hidden); err != nil {
return nil, err
}
// Hidden virtual-table columns (hidden = 1) are not part of the schema
if hidden == 1 {
continue
}
column := models.InitColumn(name, tableName, "main")
column.Type = r.mapDataType(strings.ToUpper(dataType))
column.NotNull = (notNull == 1)
column.IsPrimaryKey = (pk > 0)
column.Sequence = uint(cid + 1)
if defaultValue != nil {
if hidden == 2 || hidden == 3 {
column.Generated = true
if !tableSQLLoaded {
tableSQL = r.tableSQL(tableName)
tableSQLLoaded = true
}
column.GenerationExpression = parseGeneratedExpression(tableSQL, name)
} else if defaultValue != nil {
column.Default = *defaultValue
}
@@ -116,6 +131,16 @@ func (r *Reader) queryColumns(tableName string) (map[string]*models.Column, erro
return columns, rows.Err()
}
// tableSQL returns the CREATE TABLE statement of a table, or "" when it cannot be read.
func (r *Reader) tableSQL(tableName string) string {
var sql string
err := r.db.QueryRowContext(r.ctx, `SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?`, tableName).Scan(&sql)
if err != nil {
return ""
}
return sql
}
// isAutoIncrement checks if a column is autoincrement
func (r *Reader) isAutoIncrement(tableName, columnName string) bool {
// Check sqlite_sequence table or parse CREATE TABLE statement
+2
View File
@@ -114,6 +114,8 @@ export class Post {
- `@JoinColumn()` - Foreign key column
- `@Index()` - Index definition
- `@Unique()` - Unique constraint
- `asExpression` / `generatedType` - Generated column (`Generated` + `GenerationExpression`)
- `@PrimaryGeneratedColumn('identity')`, `@Generated('identity')`, `generatedIdentity` - Identity column
## Notes
+94 -4
View File
@@ -128,7 +128,6 @@ func (r *Reader) extractEntities(content string) []entityInfo {
scanner := bufio.NewScanner(strings.NewReader(content))
entityRegex := regexp.MustCompile(`^export\s+class\s+(\w+)`)
decoratorRegex := regexp.MustCompile(`^\s*@(\w+)(\([^)]*\))?`)
fieldRegex := regexp.MustCompile(`^\s*(\w+):\s*([^;]+);`)
var currentEntity *entityInfo
@@ -145,8 +144,7 @@ func (r *Reader) extractEntities(content string) []entityInfo {
}
// Check for decorator
if matches := decoratorRegex.FindStringSubmatch(trimmed); matches != nil {
decorator := matches[0]
if decorator, ok := matchDecorator(trimmed); ok {
pendingDecorators = append(pendingDecorators, decorator)
continue
}
@@ -488,7 +486,11 @@ func (r *Reader) parseColumnDecorator(decorator string, column *models.Column, t
column.IsPrimaryKey = true
column.NotNull = true
if strings.Contains(decorator, "'uuid'") {
if strings.Contains(decorator, "'identity'") {
column.AutoIncrement = true
column.Identity = true
column.IdentityGeneration = parseGeneratedIdentity(decorator)
} else if strings.Contains(decorator, "'uuid'") {
column.Type = "uuid"
column.Default = "gen_random_uuid()"
} else if strings.Contains(decorator, "'increment'") || strings.Contains(decorator, "()") {
@@ -497,6 +499,17 @@ func (r *Reader) parseColumnDecorator(decorator string, column *models.Column, t
return
}
// @Generated('identity') on a non-key column; generatedIdentity is read from @Column
if strings.HasPrefix(decorator, "@Generated") {
if strings.Contains(decorator, "'identity'") {
column.Identity = true
if column.IdentityGeneration == "" {
column.IdentityGeneration = "BY DEFAULT"
}
}
return
}
// @Column
if strings.HasPrefix(decorator, "@Column") {
r.parseColumnOptions(decorator, column, table)
@@ -586,6 +599,14 @@ func (r *Reader) parseColumnOptions(decorator string, column *models.Column, tab
}
}
if matches := asExpressionRegex.FindStringSubmatch(content); matches != nil {
column.Generated = true
column.GenerationExpression = unescapeSingleQuoted(matches[1])
}
if strings.Contains(content, "generatedIdentity") {
column.IdentityGeneration = parseGeneratedIdentity(content)
}
if strings.Contains(content, "nullable: true") || strings.Contains(content, "nullable:true") {
column.NotNull = false
}
@@ -834,3 +855,72 @@ func (r *Reader) getPrimaryKeyColumn(table *models.Table) *models.Column {
return pk
}
var (
asExpressionRegex = regexp.MustCompile(`asExpression:\s*'((?:\\.|[^\\'])*)'`)
generatedIdentityRegexp = regexp.MustCompile(`generatedIdentity:\s*['"](ALWAYS|BY DEFAULT)['"]`)
)
// parseGeneratedIdentity returns the identity generation mode named in a decorator
// ("ALWAYS" or "BY DEFAULT"), defaulting to "BY DEFAULT" as TypeORM does.
func parseGeneratedIdentity(decorator string) string {
if matches := generatedIdentityRegexp.FindStringSubmatch(decorator); matches != nil {
return matches[1]
}
return "BY DEFAULT"
}
// unescapeSingleQuoted reverses escaping applied inside a quoted TypeScript string.
func unescapeSingleQuoted(s string) string {
var sb strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '\\' && i+1 < len(s) {
i++
}
sb.WriteByte(s[i])
}
return sb.String()
}
var decoratorNameRegex = regexp.MustCompile(`^@\w+`)
// matchDecorator returns the decorator at the start of line, including its argument list.
// Parentheses inside quoted strings (e.g. a generated column expression) do not end it.
func matchDecorator(line string) (string, bool) {
name := decoratorNameRegex.FindString(line)
if name == "" {
return "", false
}
rest := line[len(name):]
if !strings.HasPrefix(rest, "(") {
return name, true
}
depth := 0
var quote byte
for i := 0; i < len(rest); i++ {
ch := rest[i]
if quote != 0 {
switch ch {
case '\\':
i++
case quote:
quote = 0
}
continue
}
switch ch {
case '\'', '"', '`':
quote = ch
case '(':
depth++
case ')':
depth--
if depth == 0 {
return name + rest[:i+1], true
}
}
}
// Unterminated argument list: keep the whole line, as a best effort
return line, true
}
+381
View File
@@ -0,0 +1,381 @@
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, 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()
})
form.AddButton("Tables", func() {
se.showDomainTables(index)
})
form.AddButton("Delete", func() {
se.showDeleteDomainConfirm(index)
})

Some files were not shown because too many files have changed in this diff Show More