Compare commits
61
Commits
58e46e5b59
...
v1.0.86
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2796ebe381 | ||
|
|
9447d18323 | ||
|
|
9da69f4234 | ||
|
|
2c304dc363 | ||
|
|
02a394c99d | ||
|
|
b49859537b | ||
|
|
bed80b046b | ||
|
|
495a21b67b | ||
|
|
a32647ee16 | ||
|
|
08e1417393 | ||
|
|
70282fff73 | ||
|
|
43265dac0f | ||
|
|
66b90ca54b | ||
|
|
47108809aa | ||
|
|
720476fd6e | ||
|
|
572d03fe42 | ||
|
|
f1b9079b2d | ||
|
|
bb671c3680 | ||
|
|
bc8284db25 | ||
|
|
29e747393d | ||
|
|
df980a3434 | ||
|
|
778379538b | ||
|
|
7a9219b6e3 | ||
|
|
fd9c37cd25 | ||
|
|
0235a28add | ||
|
|
d961536186 | ||
|
|
d36806047b | ||
|
|
948419ffd3 | ||
|
|
b38f53c603 | ||
|
|
ccba53c494 | ||
|
|
53327b9a5a | ||
|
|
734b14d48d | ||
|
|
938f0ed51f | ||
|
|
6e2e7eb19e | ||
|
|
82e86cdf39 | ||
|
|
b7ae00ab52 | ||
|
|
8f8664357e | ||
|
|
7bd61a21b5 | ||
|
|
6f6b9834ca | ||
|
|
9cc10715e3 | ||
|
|
bbab5ce936 | ||
|
|
275424c605 | ||
|
|
b6c0cd3d1b | ||
|
|
d7d1d99ebc | ||
|
|
43849324ce | ||
|
|
e1df3b7cee | ||
|
|
f7b5d5f054 | ||
|
|
99fc4b0944 | ||
|
|
278d488363 | ||
|
|
2f69205aa0 | ||
|
|
b91985c493 | ||
|
|
80a3453233 | ||
|
|
ade71b75fa | ||
|
|
799b7feb53 | ||
|
|
eb40201b59 | ||
|
|
1461f8f69f | ||
|
|
ea70e19a46 | ||
|
|
52f4642d97 | ||
|
|
a84592374c | ||
|
|
ca226e83df | ||
|
|
19a2cc1fe3 |
@@ -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
@@ -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.
|
||||
|
||||
@@ -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,9 +20,11 @@ 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 'main.version=$(VERSION)' -X 'main.buildDate=$(BUILD_DATE)'
|
||||
LDFLAGS := -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.Version=$(VERSION)' -X 'git.warky.dev/wdevs/relspecgo/pkg/buildinfo.BuildDate=$(BUILD_DATE)'
|
||||
|
||||
# Auto-detect container runtime (Docker or Podman)
|
||||
CONTAINER_RUNTIME := $(shell \
|
||||
@@ -207,7 +209,7 @@ docker-test-integration: docker-up ## Start DB and run integration tests
|
||||
$(GOTEST) -v ./pkg/readers/pgsql/ -count=1 || (make docker-down && exit 1)
|
||||
@make docker-down
|
||||
|
||||
release: ## Create and push a new release tag (auto-increments patch version)
|
||||
release: lint fmt-check test build ## Run lint, format check, tests, build, then create and push a new release tag
|
||||
@echo "Creating new release..."
|
||||
@latest_tag=$$(git describe --tags --abbrev=0 2>/dev/null || echo ""); \
|
||||
if [ -z "$$latest_tag" ]; then \
|
||||
@@ -223,6 +225,15 @@ release: ## Create and push a new release tag (auto-increments patch version)
|
||||
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: ## Create and push a new release tag (auto-increments patch version)
|
||||
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
|
||||
|
||||
@@ -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
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var (
|
||||
batchSourceType string
|
||||
batchInputs []string
|
||||
batchTargetType string
|
||||
batchTargetDir string
|
||||
batchPackageName string
|
||||
batchSchemaFilter string
|
||||
batchFlattenSchema bool
|
||||
batchNullableTypes string
|
||||
batchNullableArrays string
|
||||
batchContinueOnError bool
|
||||
batchKeepGoing bool
|
||||
batchDryRun bool
|
||||
)
|
||||
|
||||
var batchCmd = &cobra.Command{
|
||||
Use: "batch",
|
||||
Short: "Convert many input files to a target format in one run",
|
||||
Long: `Convert each input file independently to the target format.
|
||||
|
||||
Unlike 'convert --from-list', which merges all inputs into one output, batch
|
||||
mode writes one output per input into --to-dir. The output is named after the
|
||||
input file (without its extension). Directory-style targets (gorm, bun,
|
||||
drizzle) get a sub-directory per input.
|
||||
|
||||
Inputs are given with --input, which accepts file paths and glob patterns and
|
||||
may be repeated or comma-separated. Inputs are processed in sorted order and
|
||||
duplicates are removed. The command exits non-zero if any input fails.
|
||||
|
||||
For named, multi-step workflows use 'relspec job run' instead.
|
||||
|
||||
Examples:
|
||||
# Convert every DBML file in a directory to JSON
|
||||
relspec batch --from dbml --input "schemas/*.dbml" --to json --to-dir out/
|
||||
|
||||
# Convert specific files to GORM models, one package directory per input
|
||||
relspec batch --from json --input a.json,b.json \
|
||||
--to gorm --to-dir models/ --package models
|
||||
|
||||
# Validate everything first, writing nothing
|
||||
relspec batch --from yaml --input "specs/*.yaml" --to pgsql --to-dir sql/ --dry-run
|
||||
|
||||
# Report all failures instead of stopping at the first
|
||||
relspec batch --from json --input "*.json" --to yaml --to-dir out/ --keep-going`,
|
||||
RunE: runBatch,
|
||||
}
|
||||
|
||||
func init() {
|
||||
batchCmd.Flags().StringVar(&batchSourceType, "from", "", "Source format for every input (dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, sqlite)")
|
||||
batchCmd.Flags().StringSliceVar(&batchInputs, "input", nil, "Input file path or glob pattern (repeatable, comma-separated)")
|
||||
batchCmd.Flags().StringVar(&batchTargetType, "to", "", "Target format")
|
||||
batchCmd.Flags().StringVar(&batchTargetDir, "to-dir", "", "Output directory; one output per input is written here")
|
||||
batchCmd.Flags().StringVar(&batchPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
||||
batchCmd.Flags().StringVar(&batchSchemaFilter, "schema", "", "Filter to a specific schema by name")
|
||||
batchCmd.Flags().BoolVar(&batchFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table")
|
||||
batchCmd.Flags().StringVar(&batchNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm)")
|
||||
batchCmd.Flags().StringVar(&batchNullableArrays, "array-nullable", "", "Nullable array representation for the Bun writer")
|
||||
batchCmd.Flags().BoolVar(&batchContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL (pgsql output only)")
|
||||
batchCmd.Flags().BoolVar(&batchKeepGoing, "keep-going", false, "Process remaining inputs after a failure; still exits non-zero")
|
||||
batchCmd.Flags().BoolVar(&batchDryRun, "dry-run", false, "Read and validate every input and print the plan without writing any output")
|
||||
|
||||
for _, f := range []string{"from", "input", "to", "to-dir"} {
|
||||
if err := batchCmd.MarkFlagRequired(f); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error marking %s flag as required: %v\n", f, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// batchDirTargets are writers that emit a directory of files rather than one file.
|
||||
var batchDirTargets = map[string]bool{"gorm": true, "bun": true, "drizzle": true}
|
||||
|
||||
// batchExtensions maps single-file target formats to their output extension.
|
||||
var batchExtensions = map[string]string{
|
||||
"dbml": ".dbml", "dctx": ".dctx", "drawdb": ".ddb", "json": ".json",
|
||||
"yaml": ".yaml", "yml": ".yaml", "pgsql": ".sql", "postgres": ".sql",
|
||||
"postgresql": ".sql", "sql": ".sql", "mssql": ".sql", "sqlserver": ".sql",
|
||||
"mssql2016": ".sql", "mssql2017": ".sql", "mssql2019": ".sql", "mssql2022": ".sql",
|
||||
"sqlite": ".sql", "sqlite3": ".sql", "prisma": ".prisma", "typeorm": ".ts",
|
||||
"graphql": ".graphql", "gql": ".graphql",
|
||||
}
|
||||
|
||||
// expandBatchInputs resolves paths and glob patterns into a sorted,
|
||||
// de-duplicated file list. A pattern that matches nothing is an error.
|
||||
func expandBatchInputs(patterns []string) ([]string, error) {
|
||||
seen := map[string]bool{}
|
||||
var files []string
|
||||
for _, p := range patterns {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
matches, err := filepath.Glob(p)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid pattern %q: %w", p, err)
|
||||
}
|
||||
if len(matches) == 0 {
|
||||
return nil, fmt.Errorf("no files match %q", p)
|
||||
}
|
||||
for _, m := range matches {
|
||||
if !seen[m] {
|
||||
seen[m] = true
|
||||
files = append(files, m)
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(files) == 0 {
|
||||
return nil, fmt.Errorf("no input files given")
|
||||
}
|
||||
sort.Strings(files)
|
||||
return files, nil
|
||||
}
|
||||
|
||||
// batchOutputPaths returns the output path for each input. It errors when two
|
||||
// inputs would collide on the same output name.
|
||||
func batchOutputPaths(files []string, targetType, dir string) ([]string, error) {
|
||||
key := strings.ToLower(targetType)
|
||||
ext := ""
|
||||
if !batchDirTargets[key] {
|
||||
var ok bool
|
||||
if ext, ok = batchExtensions[key]; !ok {
|
||||
return nil, fmt.Errorf("unsupported target format: %s", targetType)
|
||||
}
|
||||
}
|
||||
outs := make([]string, len(files))
|
||||
owner := map[string]string{}
|
||||
for i, f := range files {
|
||||
stem := strings.TrimSuffix(filepath.Base(f), filepath.Ext(f))
|
||||
out := filepath.Join(dir, stem+ext)
|
||||
if prev, dup := owner[out]; dup {
|
||||
return nil, fmt.Errorf("inputs %s and %s would both write %s", prev, f, out)
|
||||
}
|
||||
owner[out] = f
|
||||
outs[i] = out
|
||||
}
|
||||
return outs, nil
|
||||
}
|
||||
|
||||
func runBatch(cmd *cobra.Command, args []string) error {
|
||||
files, err := expandBatchInputs(batchInputs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outs, err := batchOutputPaths(files, batchTargetType, batchTargetDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\n=== RelSpec Batch Converter ===\n")
|
||||
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
|
||||
fmt.Fprintf(os.Stderr, "Inputs: %d file(s), %s -> %s\n\n", len(files), batchSourceType, batchTargetType)
|
||||
|
||||
out := outWriter(cmd)
|
||||
if batchDryRun {
|
||||
fmt.Fprintf(out, "RelSpec batch plan (dry run - nothing written):\n")
|
||||
}
|
||||
|
||||
var failed []string
|
||||
for i, f := range files {
|
||||
fmt.Fprintf(os.Stderr, "[%d/%d] %s -> %s\n", i+1, len(files), f, outs[i])
|
||||
if err := processBatchItem(cmd, f, outs[i]); err != nil {
|
||||
fmt.Fprintf(os.Stderr, " ✗ %v\n", err)
|
||||
failed = append(failed, fmt.Sprintf("%s: %v", f, err))
|
||||
if !batchKeepGoing {
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
fmt.Fprintf(os.Stderr, " ✓ done\n")
|
||||
}
|
||||
|
||||
fmt.Fprintf(os.Stderr, "\n=== Batch Complete: %d ok, %d failed ===\n", len(files)-len(failed), len(failed))
|
||||
if len(failed) > 0 {
|
||||
return fmt.Errorf("batch finished with %d failure(s):\n %s", len(failed), strings.Join(failed, "\n "))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func processBatchItem(cmd *cobra.Command, in, outPath string) error {
|
||||
db, err := readDatabaseForConvert(batchSourceType, in, "")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read source: %w", err)
|
||||
}
|
||||
finalizeCommentedRefs(db, stderrWarn)
|
||||
|
||||
if batchDryRun {
|
||||
if err := validateWriteTarget(db, batchTargetType, batchPackageName, batchSchemaFilter, ""); err != nil {
|
||||
return fmt.Errorf("dry run validation failed: %w", err)
|
||||
}
|
||||
w := outWriter(cmd)
|
||||
fmt.Fprintf(w, " %s -> %s (database '%s')\n", in, outPath, db.Name)
|
||||
printDryRunPlan(w, db)
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(batchTargetDir, 0o755); err != nil {
|
||||
return fmt.Errorf("failed to create output directory: %w", err)
|
||||
}
|
||||
if err := writeDatabase(db, batchTargetType, outPath, batchPackageName, batchSchemaFilter, batchFlattenSchema, batchNullableTypes, batchNullableArrays, batchContinueOnError, ""); err != nil {
|
||||
return fmt.Errorf("failed to write target: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func saveBatchState(t *testing.T) {
|
||||
t.Helper()
|
||||
a, b, c, d, e, f, g := batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun
|
||||
t.Cleanup(func() {
|
||||
batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun = a, b, c, d, e, f, g
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunBatch_ConvertsEachInput(t *testing.T) {
|
||||
saveBatchState(t)
|
||||
dir := t.TempDir()
|
||||
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||
writeTestJSON(t, filepath.Join(dir, "b.json"), []string{"posts"})
|
||||
outDir := filepath.Join(dir, "out")
|
||||
|
||||
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||
batchPackageName, batchKeepGoing, batchDryRun = "", false, false
|
||||
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||
|
||||
cmd, _ := newDryRunCmd()
|
||||
if err := runBatch(cmd, nil); err != nil {
|
||||
t.Fatalf("batch: %v", err)
|
||||
}
|
||||
for _, name := range []string{"a.yaml", "b.yaml"} {
|
||||
if _, err := os.Stat(filepath.Join(outDir, name)); err != nil {
|
||||
t.Errorf("expected %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunBatch_DryRunWritesNothing(t *testing.T) {
|
||||
saveBatchState(t)
|
||||
dir := t.TempDir()
|
||||
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||
outDir := filepath.Join(dir, "out")
|
||||
|
||||
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||
batchPackageName, batchKeepGoing, batchDryRun = "", false, true
|
||||
batchInputs = []string{filepath.Join(dir, "a.json")}
|
||||
|
||||
cmd, buf := newDryRunCmd()
|
||||
if err := runBatch(cmd, nil); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(outDir); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the output directory")
|
||||
}
|
||||
if !strings.Contains(buf.String(), "users") || !strings.Contains(buf.String(), "a.yaml") {
|
||||
t.Errorf("plan incomplete:\n%s", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunBatch_FailureHandling(t *testing.T) {
|
||||
saveBatchState(t)
|
||||
dir := t.TempDir()
|
||||
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||
if err := os.WriteFile(filepath.Join(dir, "b.json"), []byte("{not json"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeTestJSON(t, filepath.Join(dir, "c.json"), []string{"posts"})
|
||||
outDir := filepath.Join(dir, "out")
|
||||
|
||||
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||
batchPackageName, batchDryRun = "", false
|
||||
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||
cmd, _ := newDryRunCmd()
|
||||
|
||||
// Default: stop at first failure.
|
||||
batchKeepGoing = false
|
||||
err := runBatch(cmd, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "b.json") {
|
||||
t.Fatalf("expected failure naming b.json, got %v", err)
|
||||
}
|
||||
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); !os.IsNotExist(statErr) {
|
||||
t.Error("c.json should not be processed without --keep-going")
|
||||
}
|
||||
|
||||
// --keep-going: remaining inputs are processed, exit still fails.
|
||||
batchKeepGoing = true
|
||||
if err := runBatch(cmd, nil); err == nil {
|
||||
t.Fatal("expected non-zero result with --keep-going")
|
||||
}
|
||||
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); statErr != nil {
|
||||
t.Errorf("c.yaml should be written with --keep-going: %v", statErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpandBatchInputs(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
for _, n := range []string{"b.json", "a.json"} {
|
||||
if err := os.WriteFile(filepath.Join(dir, n), []byte("{}"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
got, err := expandBatchInputs([]string{filepath.Join(dir, "*.json"), filepath.Join(dir, "a.json")})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(got) != 2 || filepath.Base(got[0]) != "a.json" || filepath.Base(got[1]) != "b.json" {
|
||||
t.Errorf("want sorted deduped [a b], got %v", got)
|
||||
}
|
||||
if _, err := expandBatchInputs([]string{filepath.Join(dir, "*.nope")}); err == nil {
|
||||
t.Error("unmatched pattern should error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBatchOutputPaths(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
files []string
|
||||
target string
|
||||
want []string
|
||||
wantErr string
|
||||
}{
|
||||
{"file target", []string{"x/a.dbml"}, "json", []string{"out/a.json"}, ""},
|
||||
{"dir target", []string{"x/a.json"}, "gorm", []string{"out/a"}, ""},
|
||||
{"collision", []string{"x/a.json", "y/a.json"}, "yaml", nil, "both write"},
|
||||
{"unsupported", []string{"a.json"}, "nope", nil, "unsupported target"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := batchOutputPaths(tt.files, tt.target, "out")
|
||||
if tt.wantErr != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("want error %q, got %v", tt.wantErr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil || len(got) != len(tt.want) || got[0] != filepath.FromSlash(tt.want[0]) {
|
||||
t.Fatalf("got %v, %v; want %v", got, err, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+153
-15
@@ -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))
|
||||
|
||||
@@ -0,0 +1,312 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
const fixturesDir = "../../tests/assets"
|
||||
|
||||
// readableFormats maps each file-based reader format to an existing fixture.
|
||||
var readableFormats = []struct {
|
||||
format string
|
||||
path string
|
||||
}{
|
||||
{"dbml", "dbml/simple.dbml"},
|
||||
{"json", "json/database.json"},
|
||||
{"yaml", "yaml/database.yaml"},
|
||||
{"yml", "yaml/database.yaml"},
|
||||
{"drawdb", "drawdb/simple.json"},
|
||||
{"dctx", "dctx/p1.dctx"},
|
||||
{"graphql", "graphql/simple.graphql"},
|
||||
{"gql", "graphql/simple.graphql"},
|
||||
{"prisma", "prisma/example.prisma"},
|
||||
{"typeorm", "typeorm/example.ts"},
|
||||
{"drizzle", "drizzle/schema.ts"},
|
||||
{"gorm", "gorm/simple.go"},
|
||||
{"bun", "bun/simple.go"},
|
||||
}
|
||||
|
||||
func TestReadDatabaseForConvert_FileFormats(t *testing.T) {
|
||||
for _, tt := range readableFormats {
|
||||
t.Run(tt.format, func(t *testing.T) {
|
||||
db, err := readDatabaseForConvert(tt.format, filepath.Join(fixturesDir, tt.path), "")
|
||||
if err != nil {
|
||||
t.Fatalf("read: %v", err)
|
||||
}
|
||||
if db == nil || len(db.Schemas) == 0 {
|
||||
t.Fatalf("no schemas read: %+v", db)
|
||||
}
|
||||
// Uppercase format names are accepted.
|
||||
if _, err := readDatabaseForConvert(strings.ToUpper(tt.format), filepath.Join(fixturesDir, tt.path), ""); err != nil {
|
||||
t.Errorf("uppercase format: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadDatabaseForConvert_Errors(t *testing.T) {
|
||||
filePathFormats := []string{"dbml", "dctx", "drawdb", "json", "yaml", "gorm", "bun", "drizzle", "prisma", "typeorm", "graphql"}
|
||||
for _, f := range filePathFormats {
|
||||
t.Run("missing path "+f, func(t *testing.T) {
|
||||
_, err := readDatabaseForConvert(f, "", "")
|
||||
if err == nil || !strings.Contains(err.Error(), "file path is required") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
connFormats := []string{"pgsql", "postgres", "postgresql", "mssql", "sqlserver", "mysql", "mariadb"}
|
||||
for _, f := range connFormats {
|
||||
t.Run("missing conn "+f, func(t *testing.T) {
|
||||
_, err := readDatabaseForConvert(f, "", "")
|
||||
if err == nil || !strings.Contains(err.Error(), "connection string is required") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
if _, err := readDatabaseForConvert("sqlite", "", ""); err == nil || !strings.Contains(err.Error(), "required for SQLite") {
|
||||
t.Errorf("sqlite: %v", err)
|
||||
}
|
||||
if _, err := readDatabaseForConvert("nope", "x", ""); err == nil || !strings.Contains(err.Error(), "unsupported source format") {
|
||||
t.Errorf("unsupported: %v", err)
|
||||
}
|
||||
if _, err := readDatabaseForConvert("dbml", filepath.Join(t.TempDir(), "missing.dbml"), ""); err == nil || !strings.Contains(err.Error(), "failed to read database") {
|
||||
t.Errorf("missing file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadDatabase_DiffReader(t *testing.T) {
|
||||
for _, f := range []string{"dbml", "json", "yaml", "drawdb", "dctx"} {
|
||||
for _, tt := range readableFormats {
|
||||
if tt.format != f {
|
||||
continue
|
||||
}
|
||||
t.Run(f, func(t *testing.T) {
|
||||
db, err := readDatabase(f, filepath.Join(fixturesDir, tt.path), "", "source")
|
||||
if err != nil || db == nil || len(db.Schemas) == 0 {
|
||||
t.Fatalf("read: %v %+v", err, db)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
for _, f := range []string{"dbml", "dctx", "drawdb", "json", "yaml", "sqldir"} {
|
||||
if _, err := readDatabase(f, "", "", "src"); err == nil || !strings.Contains(err.Error(), "src: file path is required") {
|
||||
t.Errorf("%s missing path: %v", f, err)
|
||||
}
|
||||
}
|
||||
if _, err := readDatabase("pgsql", "", "", "src"); err == nil || !strings.Contains(err.Error(), "connection string is required") {
|
||||
t.Errorf("pgsql: %v", err)
|
||||
}
|
||||
if _, err := readDatabase("sqlite", "", "", "src"); err == nil {
|
||||
t.Error("sqlite without path must fail")
|
||||
}
|
||||
if _, err := readDatabase("nope", "x", "", "src"); err == nil || !strings.Contains(err.Error(), "unsupported database format") {
|
||||
t.Errorf("unsupported: %v", err)
|
||||
}
|
||||
if _, err := readDatabase("json", filepath.Join(t.TempDir(), "missing.json"), "", "src"); err == nil || !strings.Contains(err.Error(), "src: failed to read database") {
|
||||
t.Errorf("missing file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaskPassword(t *testing.T) {
|
||||
tests := []struct {
|
||||
in, want string
|
||||
}{
|
||||
{"", ""},
|
||||
{"postgres://user:secret@host:5432/db", "postgres://user:***@host:5432/db"},
|
||||
{"postgres://user@host:5432/db", "postgres://user@host:5432/db"},
|
||||
{"host=h user=u password=secret dbname=d", "host=h user=u password=*** dbname=d"},
|
||||
{"host=h user=u", "host=h user=u"},
|
||||
{"/tmp/file.db", "/tmp/file.db"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := maskPassword(tt.in); got != tt.want {
|
||||
t.Errorf("maskPassword(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
if got := maskPasswordInDiff(tt.in); got != tt.want {
|
||||
t.Errorf("maskPasswordInDiff(%q) = %q, want %q", tt.in, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSchemaNames(t *testing.T) {
|
||||
db := models.InitDatabase("d")
|
||||
if got := getSchemaNames(db); len(got) != 0 {
|
||||
t.Errorf("empty: %v", got)
|
||||
}
|
||||
db.Schemas = []*models.Schema{{Name: "a"}, {Name: "b"}}
|
||||
if got := strings.Join(getSchemaNames(db), ","); got != "a,b" {
|
||||
t.Errorf("got %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExtraFields(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
write := func(name, body string) string {
|
||||
p := filepath.Join(dir, name)
|
||||
if err := os.WriteFile(p, []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
valid := write("valid.json", `[{"name":"extra"}]`)
|
||||
|
||||
if got, err := loadExtraFields("bun", valid); err != nil || !strings.Contains(got, "extra") {
|
||||
t.Errorf("valid: %q %v", got, err)
|
||||
}
|
||||
if _, err := loadExtraFields("BUN", valid); err != nil {
|
||||
t.Errorf("case-insensitive format: %v", err)
|
||||
}
|
||||
if _, err := loadExtraFields("gorm", valid); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
|
||||
t.Errorf("non-bun: %v", err)
|
||||
}
|
||||
if _, err := loadExtraFields("bun", filepath.Join(dir, "missing.json")); err == nil || !strings.Contains(err.Error(), "failed to read") {
|
||||
t.Errorf("missing: %v", err)
|
||||
}
|
||||
if _, err := loadExtraFields("bun", write("bad.json", `{not json`)); err == nil || !strings.Contains(err.Error(), "invalid --extra-fields JSON") {
|
||||
t.Errorf("bad json: %v", err)
|
||||
}
|
||||
if _, err := loadExtraFields("bun", write("empty.json", `[]`)); err == nil || !strings.Contains(err.Error(), "at least one field") {
|
||||
t.Errorf("empty: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func multiSchemaDB() *models.Database {
|
||||
db := models.InitDatabase("multi")
|
||||
for _, n := range []string{"a", "b"} {
|
||||
s := models.InitSchema(n)
|
||||
tbl := models.InitTable("t_"+n, n)
|
||||
c := models.InitColumn("id", tbl.Name, n)
|
||||
c.Type = "integer"
|
||||
c.IsPrimaryKey = true
|
||||
tbl.Columns["id"] = c
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
func TestValidateWriteTarget(t *testing.T) {
|
||||
db := multiSchemaDB()
|
||||
single := models.InitDatabase("single")
|
||||
single.Schemas = []*models.Schema{models.InitSchema("only")}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
db *models.Database
|
||||
dbType, pkg, schemaFilter, extraFields, wantErrSubstr string
|
||||
}{
|
||||
{"json ok", db, "json", "", "", "", ""},
|
||||
{"pgsql alias ok", db, "sql", "", "", "", ""},
|
||||
{"gorm needs package", db, "gorm", "", "", "", "package name is required"},
|
||||
{"bun needs package", db, "bun", "", "", "", "package name is required"},
|
||||
{"gorm with package", db, "gorm", "models", "", "", ""},
|
||||
{"unsupported", db, "nope", "", "", "", "unsupported target format"},
|
||||
{"schema filter found", db, "json", "", "a", "", ""},
|
||||
{"schema filter missing", db, "json", "", "zzz", "", "not found in database"},
|
||||
{"dctx multi schema", db, "dctx", "", "", "", "multiple schemas found"},
|
||||
{"dctx multi schema with filter", db, "dctx", "", "a", "", ""},
|
||||
{"dctx single schema", single, "dctx", "", "", "", ""},
|
||||
{"dctx no schemas", models.InitDatabase("e"), "dctx", "", "", "", "no schemas found"},
|
||||
{"extra fields non-bun", db, "json", "", "", "x.json", "only supported for Bun"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := validateWriteTarget(tt.db, tt.dbType, tt.pkg, tt.schemaFilter, tt.extraFields)
|
||||
if tt.wantErrSubstr == "" {
|
||||
if err != nil {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErrSubstr) {
|
||||
t.Errorf("got %v, want substring %q", err, tt.wantErrSubstr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_Formats(t *testing.T) {
|
||||
db := multiSchemaDB()
|
||||
formats := []struct{ format, file string }{
|
||||
{"json", "out.json"},
|
||||
{"yaml", "out.yaml"},
|
||||
{"yml", "out.yml"},
|
||||
{"dbml", "out.dbml"},
|
||||
{"drawdb", "out.drawdb.json"},
|
||||
{"pgsql", "out.sql"},
|
||||
{"postgres", "out2.sql"},
|
||||
{"sql", "out3.sql"},
|
||||
{"mssql", "out_ms.sql"},
|
||||
{"mysql", "out_my.sql"},
|
||||
{"sqlite", "out_lite.sql"},
|
||||
{"graphql", "out.graphql"},
|
||||
{"gql", "out2.graphql"},
|
||||
{"prisma", "out.prisma"},
|
||||
{"typeorm", "out.ts"},
|
||||
{"drizzle", "out_drizzle.ts"},
|
||||
}
|
||||
for _, tt := range formats {
|
||||
t.Run(tt.format, func(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), tt.file)
|
||||
if err := writeDatabase(db, tt.format, out, "", "", false, "", "", false, ""); err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
info, err := os.Stat(out)
|
||||
if err != nil || info.Size() == 0 {
|
||||
t.Errorf("output missing or empty: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_GoFormatsWriteIntoDir(t *testing.T) {
|
||||
db := multiSchemaDB()
|
||||
for _, f := range []string{"gorm", "bun"} {
|
||||
t.Run(f, func(t *testing.T) {
|
||||
out := filepath.Join(t.TempDir(), "models.go")
|
||||
if err := writeDatabase(db, f, out, "models", "", false, "", "", false, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := os.Stat(out); err != nil {
|
||||
t.Errorf("no output: %v", err)
|
||||
}
|
||||
if err := writeDatabase(db, f, out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "package name is required") {
|
||||
t.Errorf("missing package: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteDatabase_SchemaFilterAndDCTX(t *testing.T) {
|
||||
db := multiSchemaDB()
|
||||
out := filepath.Join(t.TempDir(), "o.json")
|
||||
|
||||
if err := writeDatabase(db, "json", out, "", "a", false, "", "", false, ""); err != nil {
|
||||
t.Errorf("schema filter: %v", err)
|
||||
}
|
||||
if err := writeDatabase(db, "json", out, "", "zzz", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "not found in database") {
|
||||
t.Errorf("missing schema: %v", err)
|
||||
}
|
||||
if err := writeDatabase(db, "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "multiple schemas found") {
|
||||
t.Errorf("dctx multi: %v", err)
|
||||
}
|
||||
single := models.InitDatabase("s")
|
||||
single.Schemas = []*models.Schema{db.Schemas[0]}
|
||||
if err := writeDatabase(single, "dctx", filepath.Join(t.TempDir(), "o.dctx"), "", "", false, "", "", false, ""); err != nil {
|
||||
t.Errorf("dctx single: %v", err)
|
||||
}
|
||||
if err := writeDatabase(models.InitDatabase("e"), "dctx", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "no schemas found") {
|
||||
t.Errorf("dctx empty: %v", err)
|
||||
}
|
||||
if err := writeDatabase(db, "nope", out, "", "", false, "", "", false, ""); err == nil || !strings.Contains(err.Error(), "unsupported target format") {
|
||||
t.Errorf("unsupported: %v", err)
|
||||
}
|
||||
if err := writeDatabase(db, "json", out, "", "", false, "", "", false, filepath.Join(t.TempDir(), "x.json")); err == nil || !strings.Contains(err.Error(), "only supported for Bun") {
|
||||
t.Errorf("extra fields with json: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newDryRunCmd() (*cobra.Command, *bytes.Buffer) {
|
||||
var buf bytes.Buffer
|
||||
cmd := &cobra.Command{}
|
||||
cmd.SetOut(&buf)
|
||||
return cmd, &buf
|
||||
}
|
||||
|
||||
func TestRunConvert_DryRunWritesNothing(t *testing.T) {
|
||||
defer func(a, b, c, d string, e bool) {
|
||||
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||
|
||||
dir := t.TempDir()
|
||||
in := filepath.Join(dir, "in.json")
|
||||
out := filepath.Join(dir, "out.json")
|
||||
writeTestJSON(t, in, []string{"users", "posts"})
|
||||
|
||||
convertSourceType, convertSourcePath = "json", in
|
||||
convertTargetType, convertTargetPath = "json", out
|
||||
convertDryRun = true
|
||||
|
||||
cmd, buf := newDryRunCmd()
|
||||
if err := runConvert(cmd, nil); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the output file")
|
||||
}
|
||||
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||
if !strings.Contains(buf.String(), want) {
|
||||
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
// Normal behavior is unchanged.
|
||||
convertDryRun = false
|
||||
if err := runConvert(cmd, nil); err != nil {
|
||||
t.Fatalf("real run: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(out); err != nil {
|
||||
t.Fatalf("real run should write output: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunConvert_DryRunValidatesTarget(t *testing.T) {
|
||||
defer func(a, b, c, d string, e bool) {
|
||||
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||
|
||||
dir := t.TempDir()
|
||||
in := filepath.Join(dir, "in.json")
|
||||
out := filepath.Join(dir, "models")
|
||||
writeTestJSON(t, in, []string{"users"})
|
||||
|
||||
convertSourceType, convertSourcePath = "json", in
|
||||
convertDryRun = true
|
||||
|
||||
// gorm without --package must fail validation, as a real run would.
|
||||
convertTargetType, convertTargetPath = "gorm", out
|
||||
cmd, _ := newDryRunCmd()
|
||||
err := runConvert(cmd, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "package name is required") {
|
||||
t.Fatalf("expected package validation error, got %v", err)
|
||||
}
|
||||
|
||||
convertTargetType = "nope"
|
||||
err = runConvert(cmd, nil)
|
||||
if err == nil || !strings.Contains(err.Error(), "unsupported target format") {
|
||||
t.Fatalf("expected unsupported format error, got %v", err)
|
||||
}
|
||||
if _, statErr := os.Stat(out); !os.IsNotExist(statErr) {
|
||||
t.Fatal("dry run must not create the output path")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunSplit_DryRunWritesNothing(t *testing.T) {
|
||||
defer func(a, b, c, d, e string, f bool) {
|
||||
splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun = a, b, c, d, e, f
|
||||
}(splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun)
|
||||
|
||||
dir := t.TempDir()
|
||||
in := filepath.Join(dir, "in.json")
|
||||
out := filepath.Join(dir, "subset.json")
|
||||
writeTestJSON(t, in, []string{"users", "posts", "comments"})
|
||||
|
||||
splitSourceType, splitSourcePath = "json", in
|
||||
splitTargetType, splitTargetPath = "json", out
|
||||
splitTables = "users,posts"
|
||||
splitDryRun = true
|
||||
|
||||
cmd, buf := newDryRunCmd()
|
||||
if err := runSplit(cmd, nil); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the output file")
|
||||
}
|
||||
got := buf.String()
|
||||
if !strings.Contains(got, "2 table(s)") || strings.Contains(got, "comments") {
|
||||
t.Errorf("plan should show only the 2 selected tables:\n%s", got)
|
||||
}
|
||||
|
||||
// A selection that matches nothing fails validation in dry-run too.
|
||||
splitTables = "does_not_exist"
|
||||
if err := runSplit(cmd, nil); err == nil {
|
||||
t.Fatal("expected error for empty selection")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunMerge_DryRunWritesNothing(t *testing.T) {
|
||||
saved := saveMergeState()
|
||||
defer restoreMergeState(saved)
|
||||
defer func(v bool) { mergeDryRun = v }(mergeDryRun)
|
||||
|
||||
dir := t.TempDir()
|
||||
target := filepath.Join(dir, "target.json")
|
||||
source := filepath.Join(dir, "source.json")
|
||||
out := filepath.Join(dir, "merged.json")
|
||||
writeTestJSON(t, target, []string{"users"})
|
||||
writeTestJSON(t, source, []string{"posts"})
|
||||
|
||||
mergeTargetType, mergeTargetPath, mergeTargetConn = "json", target, ""
|
||||
mergeSourceType, mergeSourcePath, mergeSourceConn = "json", source, ""
|
||||
mergeFromList = nil
|
||||
mergeOutputType, mergeOutputPath, mergeOutputConn = "json", out, ""
|
||||
mergeSkipTables, mergeReportPath = "", ""
|
||||
mergeDryRun = true
|
||||
|
||||
cmd, buf := newDryRunCmd()
|
||||
if err := runMerge(cmd, nil); err != nil {
|
||||
t.Fatalf("dry run: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||
t.Fatal("dry run must not create the output file")
|
||||
}
|
||||
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||
if !strings.Contains(buf.String(), want) {
|
||||
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
mergeOutputType = "nope"
|
||||
if err := runMerge(cmd, nil); err == nil || !strings.Contains(err.Error(), "unsupported format") {
|
||||
t.Fatalf("expected unsupported output format error, got %v", err)
|
||||
}
|
||||
}
|
||||
+57
-19
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"os"
|
||||
)
|
||||
|
||||
// asciiLogo (see version.go) is printed by the `version` command.
|
||||
|
||||
func main() {
|
||||
args := os.Args[1:]
|
||||
isSilent := hasSilentFlag(args)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
+17
-35
@@ -2,52 +2,26 @@ package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
|
||||
"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.
|
||||
// The actual values are set there via ldflags (see Makefile).
|
||||
var (
|
||||
// Version information, set via ldflags during build
|
||||
version = "dev"
|
||||
buildDate = "unknown"
|
||||
version = buildinfo.Version
|
||||
buildDate = buildinfo.BuildDate
|
||||
prisma7 bool
|
||||
noVersion bool
|
||||
silent bool
|
||||
strictDirectives bool
|
||||
typeMapFlags []string
|
||||
typeMappings map[string]string
|
||||
)
|
||||
|
||||
func init() {
|
||||
// If version wasn't set via ldflags, try to get it from build info
|
||||
if version == "dev" {
|
||||
if info, ok := debug.ReadBuildInfo(); ok {
|
||||
// Try to get version from VCS
|
||||
var vcsRevision, vcsTime string
|
||||
for _, setting := range info.Settings {
|
||||
switch setting.Key {
|
||||
case "vcs.revision":
|
||||
if len(setting.Value) >= 7 {
|
||||
vcsRevision = setting.Value[:7]
|
||||
}
|
||||
case "vcs.time":
|
||||
vcsTime = setting.Value
|
||||
}
|
||||
}
|
||||
|
||||
if vcsRevision != "" {
|
||||
version = vcsRevision
|
||||
}
|
||||
|
||||
if vcsTime != "" {
|
||||
if t, err := time.Parse(time.RFC3339, vcsTime); err == nil {
|
||||
buildDate = t.UTC().Format("2006-01-02 15:04:05 UTC")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
Use: "relspec",
|
||||
Short: "RelSpec - Database schema conversion and analysis tool",
|
||||
@@ -57,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)
|
||||
@@ -71,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:, …)")
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRunDiff(t *testing.T) {
|
||||
oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO := sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath
|
||||
t.Cleanup(func() {
|
||||
sourceType, sourcePath, sourceConn, targetType, targetPath, targetConn, outputFormat, outputPath = oldS, oldSP, oldSC, oldT, oldTP, oldTC, oldF, oldO
|
||||
})
|
||||
|
||||
src := filepath.Join(fixturesDir, "dbml/simple.dbml")
|
||||
cmplx := filepath.Join(fixturesDir, "dbml/complex.dbml")
|
||||
|
||||
for _, format := range []string{"summary", "json", "html"} {
|
||||
t.Run(format, func(t *testing.T) {
|
||||
sourceType, sourcePath, sourceConn = "dbml", src, ""
|
||||
targetType, targetPath, targetConn = "dbml", cmplx, ""
|
||||
outputFormat = format
|
||||
outputPath = filepath.Join(t.TempDir(), "diff.out")
|
||||
if format == "summary" {
|
||||
outputPath = ""
|
||||
}
|
||||
if err := runDiff(nil, nil); err != nil {
|
||||
t.Fatalf("runDiff: %v", err)
|
||||
}
|
||||
if outputPath != "" {
|
||||
if b, err := os.ReadFile(outputPath); err != nil || len(b) == 0 {
|
||||
t.Errorf("empty output: %v", err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("bad source", func(t *testing.T) {
|
||||
sourceType, sourcePath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
|
||||
targetType, targetPath = "dbml", src
|
||||
outputFormat, outputPath = "summary", ""
|
||||
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source database") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
})
|
||||
t.Run("bad target", func(t *testing.T) {
|
||||
sourceType, sourcePath = "dbml", src
|
||||
targetType, targetPath = "dbml", filepath.Join(t.TempDir(), "missing.dbml")
|
||||
outputFormat, outputPath = "summary", ""
|
||||
if err := runDiff(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read target database") {
|
||||
t.Errorf("got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunInspect(t *testing.T) {
|
||||
oldT, oldP, oldC, oldR, oldF, oldO, oldS := inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter
|
||||
t.Cleanup(func() {
|
||||
inspectSourceType, inspectSourcePath, inspectSourceConn, inspectRulesPath, inspectOutputFormat, inspectOutputPath, inspectSchemaFilter = oldT, oldP, oldC, oldR, oldF, oldO, oldS
|
||||
})
|
||||
|
||||
inspectSourceType = "dbml"
|
||||
inspectSourcePath = filepath.Join(fixturesDir, "dbml/simple.dbml")
|
||||
inspectSourceConn = ""
|
||||
inspectRulesPath = filepath.Join(t.TempDir(), "no-rules.yaml") // missing: defaults used or error
|
||||
inspectSchemaFilter = ""
|
||||
|
||||
// Whatever the rules outcome, the run must not panic; formats are exercised.
|
||||
for _, format := range []string{"markdown", "json"} {
|
||||
inspectOutputFormat = format
|
||||
inspectOutputPath = filepath.Join(t.TempDir(), "report."+format)
|
||||
_ = runInspect(nil, nil)
|
||||
}
|
||||
|
||||
inspectOutputFormat = "bogus"
|
||||
inspectOutputPath = ""
|
||||
if err := runInspect(nil, nil); err == nil {
|
||||
t.Error("bogus output format must fail")
|
||||
}
|
||||
|
||||
inspectSourcePath = filepath.Join(t.TempDir(), "missing.dbml")
|
||||
if err := runInspect(nil, nil); err == nil || !strings.Contains(err.Error(), "failed to read source") {
|
||||
t.Errorf("missing source: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -24,6 +24,7 @@ var (
|
||||
splitExcludeTables string
|
||||
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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -4,13 +4,16 @@ import (
|
||||
"fmt"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||
)
|
||||
|
||||
var versionCmd = &cobra.Command{
|
||||
Use: "version",
|
||||
Short: "Print version information",
|
||||
Run: func(cmd *cobra.Command, args []string) {
|
||||
fmt.Printf("RelSpec %s\n", version)
|
||||
fmt.Printf("Built: %s\n", buildDate)
|
||||
fmt.Print(buildinfo.AsciiLogo)
|
||||
fmt.Printf("RelSpec %s\n", buildinfo.Version)
|
||||
fmt.Printf("Built: %s\n", buildinfo.BuildDate)
|
||||
},
|
||||
}
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
// watchSnapshot maps a file path to its modification time and size.
|
||||
type watchSnapshot map[string]string
|
||||
|
||||
// takeWatchSnapshot records the state of every file under the given paths.
|
||||
// Directories are walked recursively. Anything at or below the excluded path
|
||||
// (typically the output path) is skipped so regenerating output does not
|
||||
// retrigger the watcher. Missing paths are simply absent from the snapshot, so
|
||||
// creating them later counts as a change.
|
||||
func takeWatchSnapshot(paths []string, exclude string) watchSnapshot {
|
||||
snap := watchSnapshot{}
|
||||
exclude = absPathOrSelf(exclude)
|
||||
for _, root := range paths {
|
||||
_ = filepath.WalkDir(root, func(p string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if exclude != "" && isWithin(absPathOrSelf(p), exclude) {
|
||||
if d.IsDir() {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if d.IsDir() {
|
||||
return nil
|
||||
}
|
||||
info, err := d.Info()
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
snap[p] = fmt.Sprintf("%d-%d", info.ModTime().UnixNano(), info.Size())
|
||||
return nil
|
||||
})
|
||||
}
|
||||
return snap
|
||||
}
|
||||
|
||||
func (s watchSnapshot) equal(o watchSnapshot) bool {
|
||||
if len(s) != len(o) {
|
||||
return false
|
||||
}
|
||||
for k, v := range s {
|
||||
if o[k] != v {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func absPathOrSelf(p string) string {
|
||||
if p == "" {
|
||||
return ""
|
||||
}
|
||||
if abs, err := filepath.Abs(p); err == nil {
|
||||
return abs
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// isWithin reports whether path equals dir or is located below it.
|
||||
func isWithin(path, dir string) bool {
|
||||
if path == dir {
|
||||
return true
|
||||
}
|
||||
return strings.HasPrefix(path, dir+string(filepath.Separator))
|
||||
}
|
||||
|
||||
// watchLoop runs fn once immediately and again whenever the watched paths
|
||||
// change, until ctx is cancelled. Changes are debounced: fn runs only after
|
||||
// the snapshot has stayed unchanged for one poll interval. Errors from fn are
|
||||
// reported to w and do not stop the loop.
|
||||
func watchLoop(ctx context.Context, w io.Writer, paths []string, exclude string, interval time.Duration, fn func() error) {
|
||||
run := func() {
|
||||
if err := fn(); err != nil {
|
||||
fmt.Fprintf(w, "Error: %v\n", err)
|
||||
}
|
||||
fmt.Fprintf(w, "Watching for changes (Ctrl-C to stop)...\n")
|
||||
}
|
||||
|
||||
last := takeWatchSnapshot(paths, exclude)
|
||||
run()
|
||||
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
cur := takeWatchSnapshot(paths, exclude)
|
||||
if cur.equal(last) {
|
||||
continue
|
||||
}
|
||||
// Debounce: wait until writes settle.
|
||||
for settled := false; !settled; {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-time.After(interval):
|
||||
}
|
||||
next := takeWatchSnapshot(paths, exclude)
|
||||
settled = next.equal(cur)
|
||||
cur = next
|
||||
}
|
||||
fmt.Fprintf(w, "\nChange detected, regenerating...\n")
|
||||
last = cur
|
||||
run()
|
||||
}
|
||||
}
|
||||
|
||||
// runConvertWatch runs the conversion once and then again whenever the source
|
||||
// files change, until interrupted.
|
||||
func runConvertWatch(parent context.Context, w io.Writer, run func() error) error {
|
||||
var paths []string
|
||||
switch {
|
||||
case len(convertFromList) > 0:
|
||||
paths = convertFromList
|
||||
case convertSourcePath != "":
|
||||
paths = []string{convertSourcePath}
|
||||
default:
|
||||
return fmt.Errorf("--watch requires --from-path or --from-list (live database connections cannot be watched)")
|
||||
}
|
||||
if convertDryRun {
|
||||
return fmt.Errorf("--watch cannot be combined with --dry-run")
|
||||
}
|
||||
if convertWatchInterval <= 0 {
|
||||
return fmt.Errorf("--watch-interval must be positive")
|
||||
}
|
||||
if parent == nil {
|
||||
parent = context.Background()
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(parent, os.Interrupt, syscall.SIGTERM)
|
||||
defer stop()
|
||||
watchLoop(ctx, w, paths, convertTargetPath, convertWatchInterval, run)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestWatchSnapshotExcludesOutput(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
out := filepath.Join(dir, "out")
|
||||
if err := os.MkdirAll(out, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
src := filepath.Join(dir, "schema.dbml")
|
||||
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
before := takeWatchSnapshot([]string{dir}, out)
|
||||
if err := os.WriteFile(filepath.Join(out, "gen.go"), []byte("x"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if after := takeWatchSnapshot([]string{dir}, out); !before.equal(after) {
|
||||
t.Errorf("writing into the excluded output path changed the snapshot")
|
||||
}
|
||||
if err := os.WriteFile(src, []byte("changed"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if after := takeWatchSnapshot([]string{dir}, out); before.equal(after) {
|
||||
t.Errorf("modifying a source file did not change the snapshot")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWatchLoopRerunsOnChange(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
src := filepath.Join(dir, "schema.dbml")
|
||||
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var runs atomic.Int32
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
watchLoop(ctx, io.Discard, []string{src}, "", 10*time.Millisecond, func() error {
|
||||
runs.Add(1)
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
|
||||
waitFor(t, func() bool { return runs.Load() == 1 })
|
||||
if err := os.WriteFile(src, []byte("changed content"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
waitFor(t, func() bool { return runs.Load() == 2 })
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
|
||||
func TestRunConvertWatchValidation(t *testing.T) {
|
||||
oldPath, oldList, oldDry, oldInt := convertSourcePath, convertFromList, convertDryRun, convertWatchInterval
|
||||
defer func() {
|
||||
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = oldPath, oldList, oldDry, oldInt
|
||||
}()
|
||||
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = "", nil, false, time.Second
|
||||
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||
t.Error("expected error without --from-path/--from-list")
|
||||
}
|
||||
convertSourcePath, convertDryRun = "x.dbml", true
|
||||
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||
t.Error("expected error combining --watch with --dry-run")
|
||||
}
|
||||
}
|
||||
|
||||
func waitFor(t *testing.T, cond func() bool) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if cond() {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("condition not met in time")
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
# Format Usage Examples
|
||||
|
||||
Examples for `relspec convert` covering the file-based reader and writer
|
||||
formats. The "Writers" and "Readers" sections below were run against
|
||||
`examples/test_schema.dbml`. The cross-format and live-database examples were
|
||||
not run; they follow the flags shown in `relspec convert --help` and require
|
||||
matching input files or reachable databases.
|
||||
|
||||
Any reader can be combined with any writer: pick `--from`/`--from-path` for the
|
||||
source and `--to`/`--to-path` for the target. Add `--silent` to suppress progress
|
||||
output.
|
||||
|
||||
## Writers: DBML to every format
|
||||
|
||||
```bash
|
||||
S="--from dbml --from-path examples/test_schema.dbml"
|
||||
|
||||
relspec convert $S --to json --to-path schema.json
|
||||
relspec convert $S --to yaml --to-path schema.yaml
|
||||
relspec convert $S --to dctx --to-path schema.dctx
|
||||
relspec convert $S --to drawdb --to-path schema.drawdb.json
|
||||
relspec convert $S --to graphql --to-path schema.graphql
|
||||
relspec convert $S --to prisma --to-path schema.prisma
|
||||
relspec convert $S --to pgsql --to-path schema.pg.sql
|
||||
relspec convert $S --to mssql --to-path schema.mssql.sql
|
||||
relspec convert $S --to sqlite --to-path schema.sqlite.sql
|
||||
relspec convert $S --to drizzle --to-path schema.ts
|
||||
relspec convert $S --to typeorm --to-path entities.ts
|
||||
relspec convert $S --to gorm --to-path models.go --package models
|
||||
relspec convert $S --to bun --to-path models.go --package models
|
||||
```
|
||||
|
||||
Notes:
|
||||
|
||||
- Code-generation writers (`gorm`, `bun`) take `--package`. They also accept
|
||||
`--types baselib|stdlib|sqltypes` to choose the nullable type package.
|
||||
- When `--to-path` is a directory it must already exist.
|
||||
- `sqlite` output automatically flattens `schema.table` names. Use
|
||||
`--flatten-schema` for other formats if the target has no schema support.
|
||||
- `dctx` supports a single schema only; use `--schema <name>` to select one.
|
||||
|
||||
## Readers: file-based formats into DBML (or JSON where noted)
|
||||
|
||||
```bash
|
||||
relspec convert --from json --from-path schema.json --to dbml --to-path out.dbml
|
||||
relspec convert --from yaml --from-path schema.yaml --to dbml --to-path out.dbml
|
||||
relspec convert --from dctx --from-path schema.dctx --to dbml --to-path out.dbml
|
||||
relspec convert --from drawdb --from-path schema.drawdb.json --to dbml --to-path out.dbml
|
||||
relspec convert --from graphql --from-path schema.graphql --to dbml --to-path out.dbml
|
||||
relspec convert --from prisma --from-path schema.prisma --to dbml --to-path out.dbml
|
||||
relspec convert --from drizzle --from-path schema.ts --to dbml --to-path out.dbml
|
||||
relspec convert --from typeorm --from-path entities.ts --to dbml --to-path out.dbml
|
||||
relspec convert --from bun --from-path models.go --to dbml --to-path out.dbml
|
||||
relspec convert --from gorm --from-path models.go --to json --to-path out.json
|
||||
```
|
||||
|
||||
Code-first readers (`gorm`, `bun`, `drizzle`, `typeorm`) accept a single file or a
|
||||
directory of model files.
|
||||
|
||||
> Known issue: reading GORM models and writing DBML currently panics in the DBML
|
||||
> writer (`pkg/writers/dbml/writer.go`, `constraintToDBML`). Use another target
|
||||
> such as JSON until this is fixed.
|
||||
|
||||
## Cross-format combinations
|
||||
|
||||
```bash
|
||||
# ORM models to SQL DDL
|
||||
relspec convert --from gorm --from-path models.go --to pgsql --to-path schema.sql
|
||||
|
||||
# Prisma to Drizzle
|
||||
relspec convert --from prisma --from-path schema.prisma --to drizzle --to-path schema.ts
|
||||
|
||||
# DrawDB diagram to GraphQL
|
||||
relspec convert --from drawdb --from-path diagram.json --to graphql --to-path schema.graphql
|
||||
|
||||
# Merge several files while converting
|
||||
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
||||
```
|
||||
|
||||
## Live databases
|
||||
|
||||
These need a reachable database:
|
||||
|
||||
```bash
|
||||
# PostgreSQL
|
||||
relspec convert --from pgsql --from-conn "postgres://user:pass@localhost:5432/mydb" \
|
||||
--to dbml --to-path schema.dbml
|
||||
|
||||
# SQL Server
|
||||
relspec convert --from mssql --from-conn "<mssql connection string>" \
|
||||
--to json --to-path schema.json
|
||||
|
||||
# SQLite database file (--from-conn takes the file path)
|
||||
relspec convert --from sqlite --from-conn ./app.db --to dbml --to-path schema.dbml
|
||||
```
|
||||
|
||||
## Formats outside `convert`
|
||||
|
||||
- `sqldir` (SQL script directory reader) and `sqlexec` (SQL execution writer) are
|
||||
used by `relspec scripts` and `relspec job`, and `sqldir` by `relspec diff`.
|
||||
See [SCRIPTS_COMMAND.md](SCRIPTS_COMMAND.md) and [JOB_FILES.md](JOB_FILES.md).
|
||||
- The `template` writer is exposed through `relspec templ`. See
|
||||
[TEMPLATE_MODE.md](TEMPLATE_MODE.md).
|
||||
+13
-2
@@ -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
|
||||
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,6 @@
|
||||
// Code generated by RelSpec. DO NOT EDIT.
|
||||
export interface {{.Name}} {
|
||||
{{- range values .Table.Columns}}
|
||||
{{.Name}}: {{.Type}};
|
||||
{{- end}}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
// Package buildinfo exposes the RelSpec version and build date so that both the
|
||||
// CLI and the schema writers can stamp generated output with the same values.
|
||||
package buildinfo
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"runtime/debug"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Version and BuildDate are set via -ldflags at build time (see Makefile). When
|
||||
// built without ldflags they are backfilled from the Go module build info.
|
||||
var (
|
||||
Version = "dev"
|
||||
BuildDate = "unknown"
|
||||
)
|
||||
|
||||
func init() {
|
||||
if Version != "dev" {
|
||||
return
|
||||
}
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
var rev, vcsTime string
|
||||
for _, s := range info.Settings {
|
||||
switch s.Key {
|
||||
case "vcs.revision":
|
||||
if len(s.Value) >= 7 {
|
||||
rev = s.Value[:7]
|
||||
}
|
||||
case "vcs.time":
|
||||
vcsTime = s.Value
|
||||
}
|
||||
}
|
||||
if rev != "" {
|
||||
Version = rev
|
||||
}
|
||||
if t, err := time.Parse(time.RFC3339, vcsTime); err == nil {
|
||||
BuildDate = t.UTC().Format("2006-01-02 15:04:05 UTC")
|
||||
}
|
||||
}
|
||||
|
||||
// GeneratedComment returns the one-line provenance string embedded in generated
|
||||
// files, e.g. "RelSpec dev (built: unknown)".
|
||||
func GeneratedComment() string {
|
||||
return fmt.Sprintf("RelSpec %s (built: %s)", Version, BuildDate)
|
||||
}
|
||||
|
||||
const AsciiLogo = `
|
||||
██████╗ ███████╗██╗ ███████╗██████╗ ███████╗ ██████╗
|
||||
██╔══██╗██╔════╝██║ ██╔════╝██╔══██╗██╔════╝██╔════╝
|
||||
██████╔╝█████╗ ██║ ███████╗██████╔╝█████╗ ██║
|
||||
██╔══██╗██╔══╝ ██║ ╚════██║██╔═══╝ ██╔══╝ ██║
|
||||
██║ ██║███████╗███████╗███████║██║ ███████╗╚██████╗
|
||||
╚═╝ ╚═╝╚══════╝╚══════╝╚══════╝╚═╝ ╚══════╝ ╚═════╝
|
||||
[ IN ] ──▶ [ RELSPEC ] ──▶ [ OUT ]
|
||||
╔══════════════════════════════════════╗
|
||||
║ ║
|
||||
║ © WARKY DEVS ║
|
||||
║ Author: Hein (hein@warky.dev) ║
|
||||
║ ║
|
||||
╚══════════════════════════════════════╝
|
||||
`
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -37,6 +37,7 @@ func TestGetDefaultConfig(t *testing.T) {
|
||||
"missing_primary_key",
|
||||
"orphaned_foreign_key",
|
||||
"circular_dependency",
|
||||
"duplicate_index_name",
|
||||
}
|
||||
|
||||
for _, ruleName := range expectedRules {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -0,0 +1,262 @@
|
||||
package jobs
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// validateYAML loads one job file and returns the Validate error text ("" when valid).
|
||||
func validateYAML(t *testing.T, body string) string {
|
||||
t.Helper()
|
||||
set := loadOne(t, "version: 1\njobs:\n"+body)
|
||||
if err := set.Validate(); err != nil {
|
||||
return err.Error()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestValidateJobTable(t *testing.T) {
|
||||
in := " inputs:\n - path: a.dbml\n format: dbml\n"
|
||||
out := " output:\n format: json\n path: out.json\n"
|
||||
tests := []struct {
|
||||
name string
|
||||
job string
|
||||
want string // substring of the error, "" for valid
|
||||
}{
|
||||
{"missing command", " x:\n description: d\n", "missing command"},
|
||||
{"convert valid", " x:\n command: convert\n" + in + out, ""},
|
||||
{"convert script dirs", " x:\n command: convert\n script_dirs: [s]\n" + in + out, "script_dirs is not valid"},
|
||||
{"convert missing output", " x:\n command: convert\n" + in, "missing output"},
|
||||
{"convert output missing format", " x:\n command: convert\n" + in + " output:\n path: o\n", "output: missing format"},
|
||||
{"convert output unsupported format", " x:\n command: convert\n" + in + " output:\n format: nope\n path: o\n", "unsupported output format"},
|
||||
{"convert output missing path", " x:\n command: convert\n" + in + " output:\n format: json\n", "output: missing path"},
|
||||
{"convert output conn_env on non-exec format", " x:\n command: convert\n" + in + " output:\n format: json\n conn_env: DB\n", "not supported for format"},
|
||||
{"convert output path and conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n path: o.sql\n", "either path or conn_env"},
|
||||
{"convert output conn_env ok", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: DB\n", ""},
|
||||
{"output secret conn_env", " x:\n command: convert\n" + in + " output:\n format: pgsql\n conn_env: postgres://u:p@h/db\n", "environment variable name"},
|
||||
{"merge needs two inputs", " x:\n command: merge\n" + in + out, "at least 2 input"},
|
||||
{"input missing format", " x:\n command: convert\n inputs:\n - path: a\n" + out, "missing format"},
|
||||
{"input unsupported format", " x:\n command: convert\n inputs:\n - path: a\n format: nope\n" + out, "unsupported input format"},
|
||||
{"input file missing path", " x:\n command: convert\n inputs:\n - format: dbml\n" + out, "missing path"},
|
||||
{"input file with conn_env", " x:\n command: convert\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n" + out, "does not use conn_env"},
|
||||
{"input db missing conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n" + out, "requires conn_env"},
|
||||
{"input db with path", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n" + out, "takes conn_env, not path"},
|
||||
{"input db ok", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: DB\n" + out, ""},
|
||||
{"input secret conn_env", " x:\n command: convert\n inputs:\n - format: pgsql\n conn_env: \"host=h password=p\"\n" + out, "environment variable name"},
|
||||
{"bad log size", " x:\n command: convert\n log_max_size: lots\n" + in + out, "log_max_size"},
|
||||
{"absolute logfile", " x:\n command: convert\n logfile: /var/log/x.log\n" + in + out, "absolute paths"},
|
||||
{"home path", " x:\n command: convert\n template: ~/t\n" + in + out, "home-relative"},
|
||||
{"report path traversal", " x:\n command: inspect\n" + in + " report:\n format: json\n path: ../r.json\n", "escapes"},
|
||||
{"script_dir traversal", " x:\n command: scripts-list\n script_dirs: [../x]\n", "escapes"},
|
||||
|
||||
{"templ valid", " x:\n command: templ\n" + in + " template: t.tmpl\n mode: table\n output:\n format: text\n path: o\n", ""},
|
||||
{"templ pgsql input valid", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n template: t.tmpl\n", ""},
|
||||
{"templ no inputs", " x:\n command: templ\n template: t.tmpl\n", "at least 1 input"},
|
||||
{"templ no template", " x:\n command: templ\n" + in, "requires template"},
|
||||
{"templ bad mode", " x:\n command: templ\n" + in + " template: t\n mode: weird\n", "unsupported mode"},
|
||||
{"templ script dirs", " x:\n command: templ\n" + in + " template: t\n script_dirs: [s]\n", "script_dirs is not valid"},
|
||||
{"templ db output", " x:\n command: templ\n" + in + " template: t\n output:\n conn_env: DB\n", "does not support database output"},
|
||||
{"templ non-text output", " x:\n command: templ\n" + in + " template: t\n output:\n format: json\n path: o\n", "only output.format: text"},
|
||||
{"templ input missing format", " x:\n command: templ\n inputs:\n - path: a\n template: t\n", "missing format"},
|
||||
{"templ pgsql input without conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n template: t\n", "requires conn_env"},
|
||||
{"templ pgsql input with path", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: DB\n path: a\n template: t\n", "takes conn_env, not path"},
|
||||
{"templ file input without path", " x:\n command: templ\n inputs:\n - format: dbml\n template: t\n", "missing path"},
|
||||
{"templ file input with conn_env", " x:\n command: templ\n inputs:\n - path: a\n format: dbml\n conn_env: DB\n template: t\n", "does not use conn_env"},
|
||||
{"templ unsupported input format", " x:\n command: templ\n inputs:\n - path: a\n format: nope\n template: t\n", "unsupported templ input format"},
|
||||
{"templ secret conn_env", " x:\n command: templ\n inputs:\n - format: pgsql\n conn_env: a/b\n template: t\n", "environment variable name"},
|
||||
|
||||
{"split needs input", " x:\n command: split\n" + out, "at least 1 input"},
|
||||
{"split script dirs", " x:\n command: split\n" + in + " script_dirs: [s]\n" + out, "script_dirs is not valid"},
|
||||
{"split report", " x:\n command: split\n" + in + " report:\n format: json\n path: r\n" + out, "report is not valid"},
|
||||
{"split db output", " x:\n command: split\n" + in + " output:\n format: pgsql\n conn_env: DB\n", "writes a file"},
|
||||
|
||||
{"inspect script dirs", " x:\n command: inspect\n" + in + " script_dirs: [s]\n report:\n path: r\n", "script_dirs is not valid"},
|
||||
{"inspect output", " x:\n command: inspect\n" + in + out + " report:\n path: r\n", "output is not valid"},
|
||||
{"inspect bad report format", " x:\n command: inspect\n" + in + " report:\n format: html\n path: r\n", "not supported"},
|
||||
{"inspect report without path", " x:\n command: inspect\n" + in + " report:\n format: json\n", "requires report.path"},
|
||||
{"inspect default format ok", " x:\n command: inspect\n" + in + " report:\n path: r.md\n", ""},
|
||||
{"diff summary without path ok", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: summary\n", ""},
|
||||
{"diff json needs path", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n report:\n format: json\n", "requires report.path"},
|
||||
{"diff output", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n" + out + " report:\n format: summary\n", "output is not valid"},
|
||||
{"diff script dirs", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n script_dirs: [s]\n report:\n format: summary\n", "script_dirs is not valid"},
|
||||
{"diff no report", " x:\n command: diff\n" + in + " - path: b.dbml\n format: dbml\n", "requires a report block"},
|
||||
|
||||
{"scripts-list inputs", " x:\n command: scripts-list\n script_dirs: [s]\n" + in, "inputs is not valid"},
|
||||
{"scripts-list output", " x:\n command: scripts-list\n script_dirs: [s]\n" + out, "output is not valid"},
|
||||
{"scripts-exec inputs", " x:\n command: scripts-exec\n script_dirs: [s]\n" + in + " output:\n conn_env: DB\n", "inputs is not valid"},
|
||||
{"scripts-exec report", " x:\n command: scripts-exec\n script_dirs: [s]\n report:\n path: r\n output:\n conn_env: DB\n", "report is not valid"},
|
||||
{"scripts-exec output path", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n path: p\n", "output.path is not supported"},
|
||||
{"scripts-exec non-pgsql", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: mssql\n", "only supports pgsql"},
|
||||
{"scripts-exec secret conn_env", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: \"postgres://u@h/d\"\n", "environment variable name"},
|
||||
{"scripts-exec no script dirs", " x:\n command: scripts-exec\n output:\n conn_env: DB\n", "requires at least one script_dir"},
|
||||
{"scripts-exec pgsql format ok", " x:\n command: scripts-exec\n script_dirs: [s]\n output:\n conn_env: DB\n format: pgsql\n", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := validateYAML(t, tt.job)
|
||||
if tt.want == "" {
|
||||
if got != "" {
|
||||
t.Errorf("expected valid, got: %s", got)
|
||||
}
|
||||
return
|
||||
}
|
||||
if !strings.Contains(got, tt.want) {
|
||||
t.Errorf("error %q does not contain %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromJobInputShape(t *testing.T) {
|
||||
producer := " p:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: out.json\n"
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"path", " - from_job: p\n path: x\n", "takes no path"},
|
||||
{"format", " - from_job: p\n format: json\n", "drop format"},
|
||||
{"conn_env", " - from_job: p\n conn_env: DB\n", "takes no conn_env"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
for _, cmd := range []string{"convert", "templ"} {
|
||||
t.Run(cmd+"/"+tt.name, func(t *testing.T) {
|
||||
extra := " output:\n format: json\n path: o.json\n"
|
||||
if cmd == "templ" {
|
||||
extra = " template: t.tmpl\n"
|
||||
}
|
||||
got := validateYAML(t, producer+" c:\n command: "+cmd+"\n inputs:\n"+tt.input+extra)
|
||||
if !strings.Contains(got, tt.want) {
|
||||
t.Errorf("error %q does not contain %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvedLogPolicy(t *testing.T) {
|
||||
keep2 := 2
|
||||
keep0 := 0
|
||||
tests := []struct {
|
||||
name string
|
||||
job Job
|
||||
want LogPolicy
|
||||
}{
|
||||
{"built-in defaults", Job{}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||
{"file defaults", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}}, LogPolicy{MaxSizeBytes: 1 << 20, Keep: 7}},
|
||||
{"file defaults invalid size falls back", Job{fileDefaults: &Defaults{LogMaxSize: "junk", LogKeep: 0}}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||
{"job overrides file", Job{fileDefaults: &Defaults{LogMaxSize: "1MB", LogKeep: 7}, LogMaxSize: "2kb", LogKeep: &keep2}, LogPolicy{MaxSizeBytes: 2 << 10, Keep: 2}},
|
||||
{"job keep zero is honoured", Job{LogKeep: &keep0}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: 0}},
|
||||
{"job invalid size ignored", Job{LogMaxSize: "junk"}, LogPolicy{MaxSizeBytes: defaultLogMaxSizeBytes, Keep: defaultLogKeep}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := tt.job.ResolvedLogPolicy(); got != tt.want {
|
||||
t.Errorf("got %+v, want %+v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadAppliesFileDefaultsAndDir(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
p := filepath.Join(dir, "relspec.yml")
|
||||
write(t, p, "version: 1\ndefaults:\n log_max_size: 1MB\n log_keep: 9\n"+"jobs:\n a:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n")
|
||||
set, err := Load([]string{p})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
job := set.Jobs["a"]
|
||||
if job.Dir() != dir {
|
||||
t.Errorf("Dir = %q, want %q", job.Dir(), dir)
|
||||
}
|
||||
if pol := job.ResolvedLogPolicy(); pol.MaxSizeBytes != 1<<20 || pol.Keep != 9 {
|
||||
t.Errorf("policy %+v", pol)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetNamesSorted(t *testing.T) {
|
||||
set := &Set{Jobs: map[string]*Job{"b": {}, "a": {}, "c": {}}}
|
||||
if got := strings.Join(set.Names(), ","); got != "a,b,c" {
|
||||
t.Errorf("got %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlanErrors(t *testing.T) {
|
||||
set := loadOne(t, "version: 1\njobs:\n"+
|
||||
" a:\n command: convert\n depends_on: [ghost]\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o.json\n"+
|
||||
" b:\n command: convert\n inputs:\n - path: a.dbml\n format: dbml\n output:\n format: json\n path: o2.json\n")
|
||||
|
||||
if _, err := set.Plan("nope", true); err == nil || !strings.Contains(err.Error(), "unknown job") || !strings.Contains(err.Error(), "a, b") {
|
||||
t.Errorf("unknown job: %v", err)
|
||||
}
|
||||
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "unknown job \"ghost\"") {
|
||||
t.Errorf("unknown dependency: %v", err)
|
||||
}
|
||||
// Without dependencies the declared dependency is not walked.
|
||||
if got, err := set.Plan("a", false); err != nil || len(got) != 1 || got[0].Name != "a" {
|
||||
t.Errorf("no-deps plan: %v %v", got, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlanCycleAtRuntime(t *testing.T) {
|
||||
set := &Set{Jobs: map[string]*Job{
|
||||
"a": {Name: "a", DependsOn: []string{"b"}},
|
||||
"b": {Name: "b", DependsOn: []string{"a"}},
|
||||
}}
|
||||
if _, err := set.Plan("a", true); err == nil || !strings.Contains(err.Error(), "cycle") {
|
||||
t.Errorf("want cycle error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverErrorsAndFiltering(t *testing.T) {
|
||||
if _, err := Discover(filepath.Join(t.TempDir(), "missing")); err == nil {
|
||||
t.Error("missing dir must fail")
|
||||
}
|
||||
dir := t.TempDir()
|
||||
for _, f := range []string{"relspec.yaml", "relspec.b.yml", "relspec.a.yaml", "relspec.txt", "other.yml", "relspec"} {
|
||||
write(t, filepath.Join(dir, f), "")
|
||||
}
|
||||
if err := os.Mkdir(filepath.Join(dir, "relspec.dir.yml"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := Discover(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var names []string
|
||||
for _, p := range got {
|
||||
names = append(names, filepath.Base(p))
|
||||
}
|
||||
if strings.Join(names, ",") != "relspec.yaml,relspec.a.yaml,relspec.b.yml" {
|
||||
t.Errorf("got %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeJoinCases(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if got, err := SafeJoin(root, "sub/file.sql"); err != nil || !strings.HasSuffix(got, filepath.Join("sub", "file.sql")) {
|
||||
t.Errorf("nested: %q %v", got, err)
|
||||
}
|
||||
for _, bad := range []string{"", "/etc/passwd", "~/x", "..", "../x", "a/../../x"} {
|
||||
if _, err := SafeJoin(root, bad); err == nil {
|
||||
t.Errorf("SafeJoin(%q) must fail", bad)
|
||||
}
|
||||
}
|
||||
if _, err := SafeJoin(filepath.Join(root, "does", "not", "exist"), "x"); err == nil {
|
||||
t.Error("unresolvable root must fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLooksLikeSecret(t *testing.T) {
|
||||
for in, want := range map[string]bool{
|
||||
"": false, "DB_URL": false, "MY_DB": false,
|
||||
"postgres://u:p@h/db": true, "host=h": true, "a b": true, "a/b": true, "u@h": true, "k:v": true,
|
||||
} {
|
||||
if got := looksLikeSecret(in); got != want {
|
||||
t.Errorf("looksLikeSecret(%q) = %v, want %v", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -82,6 +82,15 @@ func (r *MergeResult) merge(target, source *models.Database, opts *MergeOptions)
|
||||
} else {
|
||||
// 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
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
package merge
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func TestMergeSequences(t *testing.T) {
|
||||
target := models.InitSchema("public")
|
||||
target.Sequences = []*models.Sequence{{Name: "Existing", StartValue: 1, IncrementBy: 1}}
|
||||
|
||||
source := models.InitSchema("public")
|
||||
source.Sequences = []*models.Sequence{
|
||||
{Name: "existing", StartValue: 100, IncrementBy: 10}, // conflicting: must not overwrite
|
||||
{Name: "fresh", StartValue: 5, IncrementBy: 2, MinValue: 1, MaxValue: 99, CacheSize: 3, Cycle: true, OwnedByTable: "t", OwnedByColumn: "id", Comment: "c", Description: "d"},
|
||||
}
|
||||
|
||||
res := &MergeResult{}
|
||||
res.mergeSequences(target, source)
|
||||
|
||||
if res.SequencesAdded != 1 || len(target.Sequences) != 2 {
|
||||
t.Fatalf("added=%d len=%d", res.SequencesAdded, len(target.Sequences))
|
||||
}
|
||||
if target.Sequences[0].StartValue != 1 || target.Sequences[0].IncrementBy != 1 {
|
||||
t.Errorf("existing sequence was modified: %+v", target.Sequences[0])
|
||||
}
|
||||
added := target.Sequences[1]
|
||||
if added.Name != "fresh" || added.StartValue != 5 || added.IncrementBy != 2 || added.MinValue != 1 || added.MaxValue != 99 ||
|
||||
added.CacheSize != 3 || !added.Cycle || added.OwnedByTable != "t" || added.OwnedByColumn != "id" || added.Comment != "c" || added.Description != "d" {
|
||||
t.Errorf("clone lost fields: %+v", added)
|
||||
}
|
||||
if added == source.Sequences[1] {
|
||||
t.Error("sequence must be cloned, not shared")
|
||||
}
|
||||
source.Sequences[1].StartValue = 777
|
||||
if added.StartValue != 5 {
|
||||
t.Error("clone must be independent of source")
|
||||
}
|
||||
if cloneSequence(nil) != nil {
|
||||
t.Error("cloneSequence(nil) must be nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloneSchemaIsIndependent(t *testing.T) {
|
||||
src := models.InitSchema("public")
|
||||
src.Description, src.Owner, src.Comment, src.Sequence = "d", "o", "c", 4
|
||||
src.Permissions["r"] = "all"
|
||||
src.Metadata["k"] = "v"
|
||||
src.Scripts = []*models.Script{{Name: "s"}}
|
||||
|
||||
tbl := models.InitTable("t", "public")
|
||||
col := models.InitColumn("id", "t", "public")
|
||||
col.Type = "integer"
|
||||
tbl.Columns["id"] = col
|
||||
tbl.Constraints["pk"] = &models.Constraint{Name: "pk", Type: models.PrimaryKeyConstraint, Columns: []string{"id"}}
|
||||
tbl.Indexes["i"] = &models.Index{Name: "i", Columns: []string{"id"}, Include: []string{"x"}}
|
||||
tbl.Metadata["tm"] = 1
|
||||
src.Tables = []*models.Table{tbl}
|
||||
|
||||
v := models.InitView("v", "public")
|
||||
v.Definition = "select 1"
|
||||
v.Columns["c"] = &models.Column{Name: "c"}
|
||||
v.Metadata["vm"] = 1
|
||||
src.Views = []*models.View{v}
|
||||
src.Sequences = []*models.Sequence{{Name: "sq", StartValue: 3}}
|
||||
src.Enums = []*models.Enum{{Name: "e", Values: []string{"a", "b"}}}
|
||||
src.Relations = []*models.Relationship{{Name: "r", FromColumns: []string{"a"}, ToColumns: []string{"b"}, Properties: map[string]string{"p": "q"}}}
|
||||
|
||||
got := cloneSchema(src)
|
||||
if got == src || got.Name != "public" || got.Description != "d" || got.Owner != "o" || got.Comment != "c" || got.Sequence != 4 {
|
||||
t.Fatalf("scalar fields: %+v", got)
|
||||
}
|
||||
if got.Permissions["r"] != "all" || got.Metadata["k"] != "v" || len(got.Scripts) != 1 {
|
||||
t.Errorf("maps/scripts: %+v", got)
|
||||
}
|
||||
if len(got.Tables) != 1 || got.Tables[0] == tbl || got.Tables[0].Columns["id"] == col || got.Tables[0].Columns["id"].Type != "integer" {
|
||||
t.Errorf("tables not deep cloned: %+v", got.Tables)
|
||||
}
|
||||
if len(got.Views) != 1 || got.Views[0] == v || got.Views[0].Definition != "select 1" || got.Views[0].Columns["c"] == v.Columns["c"] || got.Views[0].Metadata["vm"] != 1 {
|
||||
t.Errorf("views not deep cloned: %+v", got.Views)
|
||||
}
|
||||
if len(got.Sequences) != 1 || got.Sequences[0] == src.Sequences[0] || got.Sequences[0].StartValue != 3 {
|
||||
t.Errorf("sequences: %+v", got.Sequences)
|
||||
}
|
||||
if len(got.Enums) != 1 || got.Enums[0] == src.Enums[0] || strings.Join(got.Enums[0].Values, ",") != "a,b" {
|
||||
t.Errorf("enums: %+v", got.Enums)
|
||||
}
|
||||
if len(got.Relations) != 1 || got.Relations[0] == src.Relations[0] || got.Relations[0].Properties["p"] != "q" {
|
||||
t.Errorf("relations: %+v", got.Relations)
|
||||
}
|
||||
|
||||
// Mutating the clone must not touch the source.
|
||||
got.Permissions["r"] = "none"
|
||||
got.Metadata["k"] = "changed"
|
||||
got.Tables[0].Columns["id"].Type = "text"
|
||||
got.Tables[0].Constraints["pk"].Columns[0] = "zzz"
|
||||
got.Tables[0].Indexes["i"].Columns[0] = "zzz"
|
||||
got.Tables[0].Metadata["tm"] = 2
|
||||
got.Enums[0].Values[0] = "zzz"
|
||||
got.Relations[0].FromColumns[0] = "zzz"
|
||||
got.Relations[0].Properties["p"] = "zzz"
|
||||
got.Views[0].Columns["c"].Name = "zzz"
|
||||
if src.Permissions["r"] != "all" || src.Metadata["k"] != "v" || col.Type != "integer" ||
|
||||
tbl.Constraints["pk"].Columns[0] != "id" || tbl.Indexes["i"].Columns[0] != "id" || tbl.Metadata["tm"] != 1 ||
|
||||
src.Enums[0].Values[0] != "a" || src.Relations[0].FromColumns[0] != "a" || src.Relations[0].Properties["p"] != "q" ||
|
||||
v.Columns["c"].Name != "c" {
|
||||
t.Error("clone shares state with the source")
|
||||
}
|
||||
|
||||
if cloneSchema(nil) != nil {
|
||||
t.Error("cloneSchema(nil) must be nil")
|
||||
}
|
||||
bare := cloneSchema(&models.Schema{Name: "bare"})
|
||||
if bare.Permissions != nil || bare.Metadata != nil {
|
||||
t.Errorf("nil maps must stay nil: %+v", bare)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloneNilInputs(t *testing.T) {
|
||||
if cloneTable(nil) != nil || cloneColumn(nil) != nil || cloneConstraint(nil) != nil || cloneIndex(nil) != nil ||
|
||||
cloneView(nil) != nil || cloneEnum(nil) != nil || cloneRelation(nil) != nil || cloneDomain(nil) != nil {
|
||||
t.Error("clone of nil must be nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloneDomainAndRelation(t *testing.T) {
|
||||
d := &models.Domain{Name: "d", Description: "x", Comment: "c", Sequence: 2, Metadata: map[string]any{"k": 1}, Tables: []*models.DomainTable{{TableName: "t", SchemaName: "s"}}}
|
||||
cd := cloneDomain(d)
|
||||
if cd == d || cd.Name != "d" || cd.Description != "x" || cd.Comment != "c" || cd.Sequence != 2 || cd.Metadata["k"] != 1 || len(cd.Tables) != 1 {
|
||||
t.Errorf("domain clone: %+v", cd)
|
||||
}
|
||||
cd.Metadata["k"] = 2
|
||||
if d.Metadata["k"] != 1 {
|
||||
t.Error("domain metadata shared")
|
||||
}
|
||||
|
||||
r := &models.Relationship{Name: "r", Type: "one_to_many", FromTable: "a", FromSchema: "s", ToTable: "b", ToSchema: "s", ForeignKey: "fk", ThroughTable: "l", ThroughSchema: "s", Description: "d", Sequence: 3}
|
||||
cr := cloneRelation(r)
|
||||
if cr == r || cr.Name != "r" || cr.Type != "one_to_many" || cr.FromTable != "a" || cr.ToTable != "b" || cr.ForeignKey != "fk" || cr.ThroughTable != "l" || cr.Description != "d" || cr.Sequence != 3 {
|
||||
t.Errorf("relation clone: %+v", cr)
|
||||
}
|
||||
if cr.Properties != nil {
|
||||
t.Errorf("nil properties must stay nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractTypeParts(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
col models.Column
|
||||
wantType string
|
||||
wantLen, wantPrec, wantScale int
|
||||
}{
|
||||
{"plain", models.Column{Type: "TEXT"}, "text", 0, 0, 0},
|
||||
{"trim and lower", models.Column{Type: " Integer "}, "integer", 0, 0, 0},
|
||||
{"embedded length", models.Column{Type: "varchar(50)"}, "varchar", 50, 0, 0},
|
||||
{"embedded precision and scale", models.Column{Type: "numeric(10,2)"}, "numeric", 0, 10, 2},
|
||||
{"embedded with spaces", models.Column{Type: "numeric( 10 , 2 )"}, "numeric", 0, 10, 2},
|
||||
{"fields win over embedded precision", models.Column{Type: "numeric(10,2)", Precision: 12, Scale: 4}, "numeric", 0, 12, 4},
|
||||
{"fields win over embedded length", models.Column{Type: "varchar(50)", Length: 80}, "varchar", 80, 0, 0},
|
||||
{"precision field blocks embedded length", models.Column{Type: "varchar(50)", Precision: 5}, "varchar", 0, 5, 0},
|
||||
{"non-numeric modifier", models.Column{Type: "varchar(max)"}, "varchar", 0, 0, 0},
|
||||
{"zero modifier", models.Column{Type: "char(0)"}, "char", 0, 0, 0},
|
||||
{"serial sugar", models.Column{Type: "bigserial"}, "bigint", 0, 0, 0},
|
||||
{"smallserial sugar", models.Column{Type: "smallserial"}, "smallint", 0, 0, 0},
|
||||
{"three modifiers ignored", models.Column{Type: "x(1,2,3)"}, "x", 0, 0, 0},
|
||||
{"empty", models.Column{}, "", 0, 0, 0},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
col := tt.col
|
||||
gt, gl, gp, gs := extractTypeParts(&col)
|
||||
if gt != tt.wantType || gl != tt.wantLen || gp != tt.wantPrec || gs != tt.wantScale {
|
||||
t.Errorf("got (%q,%d,%d,%d), want (%q,%d,%d,%d)", gt, gl, gp, gs, tt.wantType, tt.wantLen, tt.wantPrec, tt.wantScale)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnTypeConflict(t *testing.T) {
|
||||
c := func(typ string, l, p, s int) *models.Column {
|
||||
return &models.Column{Type: typ, Length: l, Precision: p, Scale: s}
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b *models.Column
|
||||
want bool
|
||||
}{
|
||||
{"nil target", nil, c("text", 0, 0, 0), false},
|
||||
{"nil source", c("text", 0, 0, 0), nil, false},
|
||||
{"same", c("text", 0, 0, 0), c("TEXT", 0, 0, 0), false},
|
||||
{"different base", c("text", 0, 0, 0), c("integer", 0, 0, 0), true},
|
||||
{"embedded equals field", c("varchar(50)", 0, 0, 0), c("varchar", 50, 0, 0), false},
|
||||
{"different length", c("varchar", 50, 0, 0), c("varchar", 80, 0, 0), true},
|
||||
{"different scale", c("numeric", 0, 10, 2), c("numeric", 0, 10, 3), true},
|
||||
{"serial vs int", c("bigserial", 0, 0, 0), c("bigint", 0, 0, 0), false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := columnTypeConflict(tt.a, tt.b); got != tt.want {
|
||||
t.Errorf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDescribeColumnType(t *testing.T) {
|
||||
tests := []struct {
|
||||
col *models.Column
|
||||
want string
|
||||
}{
|
||||
{nil, ""},
|
||||
{&models.Column{}, ""},
|
||||
{&models.Column{Type: " "}, ""},
|
||||
{&models.Column{Type: "text"}, "text"},
|
||||
{&models.Column{Type: " numeric ", Precision: 10, Scale: 2}, "numeric(10,2)"},
|
||||
{&models.Column{Type: "numeric", Precision: 10}, "numeric(10)"},
|
||||
{&models.Column{Type: "varchar", Length: 50}, "varchar(50)"},
|
||||
{&models.Column{Type: "varchar", Length: 50, Precision: 7}, "varchar(7)"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := describeColumnType(tt.col); got != tt.want {
|
||||
t.Errorf("describeColumnType(%+v) = %q, want %q", tt.col, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFirstNonEmpty(t *testing.T) {
|
||||
if got := firstNonEmpty("", " ", "x", "y"); got != "x" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
if got := firstNonEmpty(); got != "" {
|
||||
t.Errorf("none: %q", got)
|
||||
}
|
||||
if got := firstNonEmpty("", " "); got != "" {
|
||||
t.Errorf("all blank: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetColumnTypeConflictSummary(t *testing.T) {
|
||||
conflicts := []ColumnTypeConflict{
|
||||
{Schema: "s", Table: "t", Column: "a", TargetType: "text", SourceType: "integer"},
|
||||
{Schema: "s", Table: "t", Column: "b", TargetType: "int", SourceType: "text"},
|
||||
{Schema: "s", Table: "u", Column: "c", TargetType: "x", SourceType: "y"},
|
||||
}
|
||||
res := &MergeResult{TypeConflicts: conflicts}
|
||||
|
||||
if GetColumnTypeConflictSummary(nil, 5) != "" || GetColumnTypeConflictSummary(&MergeResult{}, 5) != "" {
|
||||
t.Error("no conflicts must yield empty summary")
|
||||
}
|
||||
|
||||
all := GetColumnTypeConflictSummary(res, 0)
|
||||
if !strings.Contains(all, "column type conflicts detected:") || !strings.Contains(all, "s.t.a: target=text source=integer") ||
|
||||
!strings.Contains(all, "s.u.c: target=x source=y") || strings.Contains(all, "more") {
|
||||
t.Errorf("unlimited summary:\n%s", all)
|
||||
}
|
||||
if neg := GetColumnTypeConflictSummary(res, -1); neg != all {
|
||||
t.Error("negative limit must behave as unlimited")
|
||||
}
|
||||
|
||||
limited := GetColumnTypeConflictSummary(res, 2)
|
||||
if !strings.Contains(limited, "s.t.b") || strings.Contains(limited, "s.u.c") || !strings.HasSuffix(limited, "... and 1 more") {
|
||||
t.Errorf("limited summary:\n%s", limited)
|
||||
}
|
||||
exact := GetColumnTypeConflictSummary(res, 3)
|
||||
if strings.Contains(exact, "more") {
|
||||
t.Errorf("limit == len must not truncate:\n%s", exact)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMinHelper(t *testing.T) {
|
||||
if min(1, 2) != 1 || min(2, 1) != 1 || min(3, 3) != 3 {
|
||||
t.Error("min")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package merge
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func sourceWithRelationship() *models.Database {
|
||||
db := models.InitDatabase("src")
|
||||
s := models.InitSchema("sales")
|
||||
orders := models.InitTable("orders", "sales")
|
||||
orders.Tablespace = "fast"
|
||||
orders.GUID = "guid-1"
|
||||
orders.Relationships["fk_cust"] = &models.Relationship{
|
||||
Name: "fk_cust", FromTable: "orders", ToTable: "customers",
|
||||
FromColumns: []string{"cust_id"}, ToColumns: []string{"id"},
|
||||
}
|
||||
s.Tables = append(s.Tables, orders, models.InitTable("Audit", "sales"))
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func TestCloneTable_CopiesRelationshipsTablespaceGUID(t *testing.T) {
|
||||
src := sourceWithRelationship()
|
||||
target := models.InitDatabase("tgt")
|
||||
MergeDatabases(target, src, nil)
|
||||
|
||||
got := target.Schemas[0].Tables[0]
|
||||
if got.Tablespace != "fast" || got.GUID != "guid-1" {
|
||||
t.Errorf("tablespace/guid lost: %+v", got)
|
||||
}
|
||||
rel := got.Relationships["fk_cust"]
|
||||
if rel == nil || rel.ToTable != "customers" {
|
||||
t.Fatalf("relationship lost: %+v", got.Relationships)
|
||||
}
|
||||
if rel == src.Schemas[0].Tables[0].Relationships["fk_cust"] {
|
||||
t.Error("relationship must be deep-copied")
|
||||
}
|
||||
rel.FromColumns[0] = "changed"
|
||||
if src.Schemas[0].Tables[0].Relationships["fk_cust"].FromColumns[0] != "cust_id" {
|
||||
t.Error("relationship columns shared with source")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMerge_SkipTablesAppliesToNewSchemas(t *testing.T) {
|
||||
src := sourceWithRelationship()
|
||||
target := models.InitDatabase("tgt")
|
||||
MergeDatabases(target, src, &MergeOptions{SkipTableNames: map[string]bool{"audit": true}})
|
||||
|
||||
tables := target.Schemas[0].Tables
|
||||
if len(tables) != 1 || tables[0].Name != "orders" {
|
||||
t.Errorf("skipped table copied into new schema: %+v", tables)
|
||||
}
|
||||
if len(src.Schemas[0].Tables) != 2 {
|
||||
t.Error("source must not be modified")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
+22
-17
@@ -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
|
||||
@@ -227,23 +228,27 @@ func (d *Sequence) SQLName() string {
|
||||
|
||||
// Column represents a table column
|
||||
type Column struct {
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table" yaml:"table" xml:"table"`
|
||||
Schema string `json:"schema" yaml:"schema" xml:"schema"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"`
|
||||
Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"`
|
||||
Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"`
|
||||
Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"`
|
||||
NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"`
|
||||
Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"`
|
||||
AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"`
|
||||
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Name string `json:"name" yaml:"name" xml:"name"`
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty" xml:"description,omitempty"`
|
||||
Table string `json:"table" yaml:"table" xml:"table"`
|
||||
Schema string `json:"schema" yaml:"schema" xml:"schema"`
|
||||
Type string `json:"type" yaml:"type" xml:"type"`
|
||||
Length int `json:"length,omitempty" yaml:"length,omitempty" xml:"length,omitempty"`
|
||||
Precision int `json:"precision,omitempty" yaml:"precision,omitempty" xml:"precision,omitempty"`
|
||||
Scale int `json:"scale,omitempty" yaml:"scale,omitempty" xml:"scale,omitempty"`
|
||||
NotNull bool `json:"not_null" yaml:"not_null" xml:"not_null"`
|
||||
Default any `json:"default,omitempty" yaml:"default,omitempty" xml:"default,omitempty"`
|
||||
AutoIncrement bool `json:"auto_increment" yaml:"auto_increment" xml:"auto_increment"`
|
||||
IsPrimaryKey bool `json:"is_primary_key" yaml:"is_primary_key" xml:"is_primary_key"`
|
||||
Comment string `json:"comment,omitempty" yaml:"comment,omitempty" xml:"comment,omitempty"`
|
||||
Collation string `json:"collation,omitempty" yaml:"collation,omitempty" xml:"collation,omitempty"`
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty" xml:"-"`
|
||||
Sequence uint `json:"sequence,omitempty" yaml:"sequence,omitempty" xml:"sequence,omitempty"`
|
||||
GUID string `json:"guid" yaml:"guid" xml:"guid"`
|
||||
Generated bool `json:"generated,omitempty" yaml:"generated,omitempty" xml:"generated,omitempty"`
|
||||
GenerationExpression string `json:"generation_expression,omitempty" yaml:"generation_expression,omitempty" xml:"generation_expression,omitempty"`
|
||||
Identity bool `json:"identity,omitempty" yaml:"identity,omitempty" xml:"identity,omitempty"`
|
||||
IdentityGeneration string `json:"identity_generation,omitempty" yaml:"identity_generation,omitempty" xml:"identity_generation,omitempty"` // "ALWAYS" or "BY DEFAULT"
|
||||
}
|
||||
|
||||
// SQLName returns the column name in lowercase for SQL compatibility.
|
||||
|
||||
@@ -0,0 +1,232 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSQLNameLowercases(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
got string
|
||||
}{
|
||||
{"database", (&Database{Name: "MyDB"}).SQLName()},
|
||||
{"domain", (&Domain{Name: "MyDomain"}).SQLName()},
|
||||
{"schema", (&Schema{Name: "MySchema"}).SQLName()},
|
||||
{"table", (&Table{Name: "MyTable"}).SQLName()},
|
||||
{"view", (&View{Name: "MyView"}).SQLName()},
|
||||
{"sequence", (&Sequence{Name: "MySeq"}).SQLName()},
|
||||
{"column", (&Column{Name: "MyCol"}).SQLName()},
|
||||
{"index", (&Index{Name: "MyIdx"}).SQLName()},
|
||||
{"relationship", (&Relationship{Name: "MyRel"}).SQLName()},
|
||||
{"constraint", (&Constraint{Name: "MyCon"}).SQLName()},
|
||||
{"enum", (&Enum{Name: "MyEnum"}).SQLName()},
|
||||
{"script", (&Script{Name: "MyScript"}).SQLName()},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if tt.got == "" || tt.got != lower(tt.got) {
|
||||
t.Errorf("SQLName not lowercase: %q", tt.got)
|
||||
}
|
||||
})
|
||||
}
|
||||
if got := (&Table{}).SQLName(); got != "" {
|
||||
t.Errorf("empty name: %q", got)
|
||||
}
|
||||
if got := (&Table{Name: "MyTable"}).SQLName(); got != "mytable" {
|
||||
t.Errorf("got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func lower(s string) string {
|
||||
b := []byte(s)
|
||||
for i, c := range b {
|
||||
if c >= 'A' && c <= 'Z' {
|
||||
b[i] = c + 32
|
||||
}
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
func TestUpdateDatePropagates(t *testing.T) {
|
||||
db := InitDatabase("d")
|
||||
schema := InitSchema("s")
|
||||
schema.RefDatabase = db
|
||||
table := InitTable("t", "s")
|
||||
table.RefSchema = schema
|
||||
|
||||
table.UpdateDate()
|
||||
for name, v := range map[string]string{"table": table.UpdatedAt, "schema": schema.UpdatedAt, "database": db.UpdatedAt} {
|
||||
ts, err := time.Parse(time.RFC3339, v)
|
||||
if err != nil {
|
||||
t.Fatalf("%s UpdatedAt %q: %v", name, v, err)
|
||||
}
|
||||
if time.Since(ts) > time.Minute {
|
||||
t.Errorf("%s UpdatedAt too old: %v", name, ts)
|
||||
}
|
||||
}
|
||||
|
||||
// Without references only the receiver is updated.
|
||||
lone := InitTable("lone", "s")
|
||||
lone.UpdateDate()
|
||||
if lone.UpdatedAt == "" {
|
||||
t.Error("lone table not updated")
|
||||
}
|
||||
loneSchema := InitSchema("x")
|
||||
loneSchema.UpdateDate()
|
||||
if loneSchema.UpdatedAt == "" {
|
||||
t.Error("lone schema not updated")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPrimaryKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cols []*Column
|
||||
want string
|
||||
}{
|
||||
{"none", []*Column{{Name: "a"}}, ""},
|
||||
{"single", []*Column{{Name: "a"}, {Name: "id", IsPrimaryKey: true}}, "id"},
|
||||
{"composite ordered by sequence", []*Column{
|
||||
{Name: "a", IsPrimaryKey: true, Sequence: 2},
|
||||
{Name: "b", IsPrimaryKey: true, Sequence: 1},
|
||||
}, "b"},
|
||||
{"composite without sequence falls back to name", []*Column{
|
||||
{Name: "z", IsPrimaryKey: true},
|
||||
{Name: "m", IsPrimaryKey: true},
|
||||
}, "m"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tbl := InitTable("t", "s")
|
||||
for _, c := range tt.cols {
|
||||
tbl.Columns[c.Name] = c
|
||||
}
|
||||
got := tbl.GetPrimaryKey()
|
||||
if tt.want == "" {
|
||||
if got != nil {
|
||||
t.Errorf("expected nil, got %s", got.Name)
|
||||
}
|
||||
return
|
||||
}
|
||||
if got == nil || got.Name != tt.want {
|
||||
t.Errorf("got %v, want %s", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
if InitTable("empty", "s").GetPrimaryKey() != nil {
|
||||
t.Error("empty table must have no PK")
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnLess(t *testing.T) {
|
||||
tests := []struct {
|
||||
a, b *Column
|
||||
want bool
|
||||
}{
|
||||
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b", Sequence: 2}, true},
|
||||
{&Column{Name: "a", Sequence: 2}, &Column{Name: "b", Sequence: 1}, false},
|
||||
{&Column{Name: "a"}, &Column{Name: "b"}, true},
|
||||
{&Column{Name: "b"}, &Column{Name: "a"}, false},
|
||||
{&Column{Name: "b", Sequence: 1}, &Column{Name: "a"}, false}, // one side unsequenced: by name
|
||||
{&Column{Name: "a", Sequence: 1}, &Column{Name: "b"}, true},
|
||||
}
|
||||
for i, tt := range tests {
|
||||
if got := columnLess(tt.a, tt.b); got != tt.want {
|
||||
t.Errorf("case %d: got %v, want %v", i, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetForeignKeys(t *testing.T) {
|
||||
tbl := InitTable("t", "s")
|
||||
add := func(name string, typ ConstraintType, seq uint) {
|
||||
c := InitConstraint(name, typ)
|
||||
c.Sequence = seq
|
||||
tbl.Constraints[name] = c
|
||||
}
|
||||
add("pk", PrimaryKeyConstraint, 0)
|
||||
add("fk_b", ForeignKeyConstraint, 0)
|
||||
add("fk_a", ForeignKeyConstraint, 0)
|
||||
add("uq", UniqueConstraint, 0)
|
||||
|
||||
got := tbl.GetForeignKeys()
|
||||
if len(got) != 2 || got[0].Name != "fk_a" || got[1].Name != "fk_b" {
|
||||
t.Errorf("by name: %v", got)
|
||||
}
|
||||
|
||||
tbl.Constraints["fk_a"].Sequence = 5
|
||||
tbl.Constraints["fk_b"].Sequence = 2
|
||||
got = tbl.GetForeignKeys()
|
||||
if got[0].Name != "fk_b" || got[1].Name != "fk_a" {
|
||||
t.Errorf("by sequence: %v", got)
|
||||
}
|
||||
|
||||
if got := InitTable("e", "s").GetForeignKeys(); got == nil || len(got) != 0 {
|
||||
t.Errorf("empty table must give non-nil empty slice, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitConstructors(t *testing.T) {
|
||||
db := InitDatabase("db")
|
||||
if db.Name != "db" || db.Schemas == nil || db.Domains == nil || db.Metadata == nil || db.GUID == "" {
|
||||
t.Errorf("InitDatabase: %+v", db)
|
||||
}
|
||||
s := InitSchema("s")
|
||||
if s.Name != "s" || s.Tables == nil || s.Views == nil || s.Sequences == nil || s.Permissions == nil || s.Metadata == nil || s.Scripts == nil || s.GUID == "" {
|
||||
t.Errorf("InitSchema: %+v", s)
|
||||
}
|
||||
tb := InitTable("t", "s")
|
||||
if tb.Name != "t" || tb.Schema != "s" || tb.Columns == nil || tb.Constraints == nil || tb.Indexes == nil || tb.Relationships == nil || tb.Metadata == nil || tb.GUID == "" {
|
||||
t.Errorf("InitTable: %+v", tb)
|
||||
}
|
||||
c := InitColumn("c", "t", "s")
|
||||
if c.Name != "c" || c.Table != "t" || c.Schema != "s" || c.Metadata == nil || c.GUID == "" {
|
||||
t.Errorf("InitColumn: %+v", c)
|
||||
}
|
||||
ix := InitIndex("i", "t", "s")
|
||||
if ix.Name != "i" || ix.Table != "t" || ix.Schema != "s" || ix.Columns == nil || ix.Include == nil || ix.Metadata == nil || ix.GUID == "" {
|
||||
t.Errorf("InitIndex: %+v", ix)
|
||||
}
|
||||
r := InitRelation("r", "s")
|
||||
if r.Name != "r" || r.FromSchema != "s" || r.ToSchema != "s" || r.Properties == nil || r.FromColumns == nil || r.ToColumns == nil || r.GUID == "" {
|
||||
t.Errorf("InitRelation: %+v", r)
|
||||
}
|
||||
rel := InitRelationship("rel", RelationType("one_to_many"))
|
||||
if rel.Name != "rel" || rel.Type != "one_to_many" || rel.Properties == nil || rel.GUID == "" {
|
||||
t.Errorf("InitRelationship: %+v", rel)
|
||||
}
|
||||
con := InitConstraint("k", UniqueConstraint)
|
||||
if con.Name != "k" || con.Type != UniqueConstraint || con.Columns == nil || con.ReferencedColumns == nil || con.GUID == "" {
|
||||
t.Errorf("InitConstraint: %+v", con)
|
||||
}
|
||||
sc := InitScript("sc")
|
||||
if sc.Name != "sc" || sc.RunAfter == nil || sc.Metadata == nil || sc.GUID == "" {
|
||||
t.Errorf("InitScript: %+v", sc)
|
||||
}
|
||||
v := InitView("v", "s")
|
||||
if v.Name != "v" || v.Schema != "s" || v.Columns == nil || v.Metadata == nil || v.GUID == "" {
|
||||
t.Errorf("InitView: %+v", v)
|
||||
}
|
||||
sq := InitSequence("sq", "s")
|
||||
if sq.Name != "sq" || sq.Schema != "s" || sq.IncrementBy != 1 || sq.StartValue != 1 || sq.GUID == "" {
|
||||
t.Errorf("InitSequence: %+v", sq)
|
||||
}
|
||||
d := InitDomain("d")
|
||||
if d.Name != "d" || d.Tables == nil || d.Metadata == nil || d.GUID == "" {
|
||||
t.Errorf("InitDomain: %+v", d)
|
||||
}
|
||||
dt := InitDomainTable("t", "s")
|
||||
if dt.TableName != "t" || dt.SchemaName != "s" || dt.GUID == "" {
|
||||
t.Errorf("InitDomainTable: %+v", dt)
|
||||
}
|
||||
e := InitEnum("e", "s")
|
||||
if e.Name != "e" || e.Schema != "s" || e.Values == nil || e.GUID == "" {
|
||||
t.Errorf("InitEnum: %+v", e)
|
||||
}
|
||||
|
||||
// GUIDs are unique per call.
|
||||
if InitTable("t", "s").GUID == InitTable("t", "s").GUID {
|
||||
t.Error("GUIDs must be unique")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+105
-18
@@ -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
|
||||
}
|
||||
|
||||
@@ -305,9 +321,39 @@ func sortDBMLFiles(files []string) []string {
|
||||
// Merges: Columns (map), Constraints (map), Indexes (map), Relationships (map)
|
||||
// Uses first non-empty Description
|
||||
func mergeTable(baseTable, fileTable *models.Table) {
|
||||
// Merge columns (map naturally merges - later keys overwrite)
|
||||
for key, col := range fileTable.Columns {
|
||||
baseTable.Columns[key] = col
|
||||
// Merge columns. Each file numbers its own columns from 1, so a table split
|
||||
// across files would otherwise end up with colliding Column.Sequence values
|
||||
// and writers would fall back to alphabetical order. Re-base the incoming
|
||||
// file's new columns after the highest sequence already present, preserving
|
||||
// their in-file order. Columns that overwrite an existing key keep the
|
||||
// original position.
|
||||
var maxSeq uint
|
||||
for _, col := range baseTable.Columns {
|
||||
if col.Sequence > maxSeq {
|
||||
maxSeq = col.Sequence
|
||||
}
|
||||
}
|
||||
|
||||
incoming := make([]*models.Column, 0, len(fileTable.Columns))
|
||||
for _, col := range fileTable.Columns {
|
||||
incoming = append(incoming, col)
|
||||
}
|
||||
sort.Slice(incoming, func(i, j int) bool {
|
||||
if incoming[i].Sequence != incoming[j].Sequence {
|
||||
return incoming[i].Sequence < incoming[j].Sequence
|
||||
}
|
||||
return incoming[i].Name < incoming[j].Name
|
||||
})
|
||||
|
||||
var added uint
|
||||
for _, col := range incoming {
|
||||
if existing, ok := baseTable.Columns[col.Name]; ok {
|
||||
col.Sequence = existing.Sequence
|
||||
} else {
|
||||
added++
|
||||
col.Sequence = maxSeq + added
|
||||
}
|
||||
baseTable.Columns[col.Name] = col
|
||||
}
|
||||
|
||||
// Merge constraints
|
||||
@@ -410,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
|
||||
@@ -463,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
|
||||
}
|
||||
|
||||
@@ -551,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
|
||||
@@ -629,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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -655,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) {
|
||||
@@ -719,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:
|
||||
|
||||
@@ -652,6 +652,22 @@ func TestReadDirectory_TableMerging(t *testing.T) {
|
||||
if emailCol.Type != "varchar(255)" {
|
||||
t.Errorf("Expected email type 'varchar(255)', got '%s'", emailCol.Type)
|
||||
}
|
||||
|
||||
// Merged columns must keep declaration order (file 1: id, email; file 3:
|
||||
// name, created_at) via strictly increasing, non-colliding Sequence values
|
||||
// so downstream writers do not fall back to alphabetical order.
|
||||
order := []string{"id", "email", "name", "created_at"}
|
||||
var prev uint
|
||||
for i, name := range order {
|
||||
col := usersTable.Columns[name]
|
||||
if col.Sequence == 0 {
|
||||
t.Fatalf("column %q has zero Sequence after merge", name)
|
||||
}
|
||||
if i > 0 && col.Sequence <= prev {
|
||||
t.Errorf("column %q Sequence %d not greater than previous %d (order not preserved across files)", name, col.Sequence, prev)
|
||||
}
|
||||
prev = col.Sequence
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadDirectory_CommentedRefsLast(t *testing.T) {
|
||||
@@ -1069,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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
package mysql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
_ "github.com/go-sql-driver/mysql"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
)
|
||||
|
||||
type Reader struct {
|
||||
options *readers.ReaderOptions
|
||||
db *sql.DB
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func NewReader(options *readers.ReaderOptions) *Reader {
|
||||
return &Reader{options: options, ctx: context.Background()}
|
||||
}
|
||||
|
||||
func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||
if r.options == nil || r.options.ConnectionString == "" {
|
||||
return nil, fmt.Errorf("connection string is required")
|
||||
}
|
||||
if err := r.connect(); err != nil {
|
||||
return nil, fmt.Errorf("failed to connect: %w", err)
|
||||
}
|
||||
defer r.close()
|
||||
var name, version string
|
||||
if err := r.db.QueryRowContext(r.ctx, "SELECT DATABASE()").Scan(&name); err != nil {
|
||||
return nil, fmt.Errorf("failed to get database name: %w", err)
|
||||
}
|
||||
_ = r.db.QueryRowContext(r.ctx, "SELECT VERSION()").Scan(&version)
|
||||
db := models.InitDatabase(name)
|
||||
db.DatabaseType, db.SourceFormat, db.DatabaseVersion = models.MySQLDatabaseType, "mysql", version
|
||||
schemas, err := r.querySchemas(name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query schemas: %w", err)
|
||||
}
|
||||
for _, schema := range schemas {
|
||||
tables, err := r.queryTables(schema.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
schema.Tables = tables
|
||||
for _, table := range tables {
|
||||
table.Columns, err = r.queryColumns(schema.Name, table.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
table.Constraints, err = r.queryConstraints(schema.Name, table.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
table.Indexes, err = r.queryIndexes(schema.Name, table.Name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
table.RefSchema = schema
|
||||
for _, c := range table.Constraints {
|
||||
if c.Type == models.ForeignKeyConstraint {
|
||||
r.deriveRelationship(table, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
schema.RefDatabase = db
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
|
||||
func (r *Reader) ReadSchema() (*models.Schema, error) {
|
||||
db, err := r.ReadDatabase()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(db.Schemas) == 0 {
|
||||
return nil, fmt.Errorf("no schemas found in database")
|
||||
}
|
||||
return db.Schemas[0], nil
|
||||
}
|
||||
|
||||
func (r *Reader) ReadTable() (*models.Table, error) {
|
||||
s, err := r.ReadSchema()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(s.Tables) == 0 {
|
||||
return nil, fmt.Errorf("no tables found in schema")
|
||||
}
|
||||
return s.Tables[0], nil
|
||||
}
|
||||
|
||||
func (r *Reader) connect() error {
|
||||
db, err := sql.Open("mysql", r.options.ConnectionString)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = db.PingContext(r.ctx); err != nil {
|
||||
db.Close()
|
||||
return err
|
||||
}
|
||||
r.db = db
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Reader) close() {
|
||||
if r.db != nil {
|
||||
_ = r.db.Close()
|
||||
}
|
||||
}
|
||||
func (r *Reader) mapDataType(t string) string { return mariadb.ConvertMariaDBToCanonical(t) }
|
||||
|
||||
func (r *Reader) querySchemas(current string) ([]*models.Schema, error) {
|
||||
rows, err := r.db.QueryContext(r.ctx, "SELECT SCHEMA_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME = ?", current)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Schema
|
||||
for rows.Next() {
|
||||
var n string
|
||||
if err := rows.Scan(&n); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, models.InitSchema(n))
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (r *Reader) queryTables(schema string) ([]*models.Table, error) {
|
||||
rows, err := r.db.QueryContext(r.ctx, "SELECT TABLE_NAME FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? AND TABLE_TYPE = 'BASE TABLE' ORDER BY TABLE_NAME", schema)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*models.Table
|
||||
for rows.Next() {
|
||||
var n string
|
||||
if err := rows.Scan(&n); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, models.InitTable(n, schema))
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (r *Reader) queryColumns(schema, table string) (map[string]*models.Column, error) {
|
||||
rows, err := r.db.QueryContext(r.ctx, `SELECT COLUMN_NAME, COLUMN_TYPE, IS_NULLABLE, COLUMN_DEFAULT, ORDINAL_POSITION, EXTRA, COLUMN_COMMENT FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? ORDER BY ORDINAL_POSITION`, schema, table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := map[string]*models.Column{}
|
||||
for rows.Next() {
|
||||
var name, typ, nullable, extra, comment string
|
||||
var def sql.NullString
|
||||
var pos int
|
||||
if err := rows.Scan(&name, &typ, &nullable, &def, &pos, &extra, &comment); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c := models.InitColumn(name, table, schema)
|
||||
c.Type = r.mapDataType(typ)
|
||||
c.NotNull = strings.EqualFold(nullable, "NO")
|
||||
c.Sequence = uint(pos)
|
||||
c.Comment = comment
|
||||
if def.Valid {
|
||||
c.Default = def.String
|
||||
}
|
||||
c.AutoIncrement = strings.Contains(strings.ToLower(extra), "auto_increment")
|
||||
out[name] = c
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (r *Reader) queryConstraints(schema, table string) (map[string]*models.Constraint, error) {
|
||||
rows, err := r.db.QueryContext(r.ctx, `SELECT CONSTRAINT_NAME, CONSTRAINT_TYPE, COLUMN_NAME, REFERENCED_TABLE_SCHEMA, REFERENCED_TABLE_NAME, REFERENCED_COLUMN_NAME, ORDINAL_POSITION FROM information_schema.KEY_COLUMN_USAGE k JOIN information_schema.TABLE_CONSTRAINTS t USING (CONSTRAINT_SCHEMA, TABLE_NAME, CONSTRAINT_NAME) WHERE k.TABLE_SCHEMA=? AND k.TABLE_NAME=? ORDER BY CONSTRAINT_NAME, ORDINAL_POSITION`, schema, table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := map[string]*models.Constraint{}
|
||||
for rows.Next() {
|
||||
var name, typ, col, rs, rt, rc string
|
||||
var pos int
|
||||
if err := rows.Scan(&name, &typ, &col, &rs, &rt, &rc, &pos); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c := out[name]
|
||||
if c == nil {
|
||||
ct := models.UniqueConstraint
|
||||
if typ == "PRIMARY KEY" {
|
||||
ct = models.PrimaryKeyConstraint
|
||||
}
|
||||
if typ == "FOREIGN KEY" {
|
||||
ct = models.ForeignKeyConstraint
|
||||
}
|
||||
c = models.InitConstraint(name, ct)
|
||||
c.Schema = schema
|
||||
c.Table = table
|
||||
c.ReferencedSchema = rs
|
||||
c.ReferencedTable = rt
|
||||
out[name] = c
|
||||
}
|
||||
c.Columns = append(c.Columns, col)
|
||||
if rc != "" {
|
||||
c.ReferencedColumns = append(c.ReferencedColumns, rc)
|
||||
}
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (r *Reader) queryIndexes(schema, table string) (map[string]*models.Index, error) {
|
||||
rows, err := r.db.QueryContext(r.ctx, `SELECT INDEX_NAME, NON_UNIQUE, COLUMN_NAME, SEQ_IN_INDEX, INDEX_TYPE FROM information_schema.STATISTICS WHERE TABLE_SCHEMA=? AND TABLE_NAME=? AND INDEX_NAME <> 'PRIMARY' ORDER BY INDEX_NAME, SEQ_IN_INDEX`, schema, table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := map[string]*models.Index{}
|
||||
for rows.Next() {
|
||||
var name, col, typ string
|
||||
var non, seq int
|
||||
if err := rows.Scan(&name, &non, &col, &seq, &typ); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
i := out[name]
|
||||
if i == nil {
|
||||
i = models.InitIndex(name, table, schema)
|
||||
i.Unique = non == 0
|
||||
i.Type = strings.ToLower(typ)
|
||||
out[name] = i
|
||||
}
|
||||
i.Columns = append(i.Columns, col)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (r *Reader) deriveRelationship(t *models.Table, c *models.Constraint) {
|
||||
n := fmt.Sprintf("%s_to_%s", t.Name, c.ReferencedTable)
|
||||
rel := models.InitRelationship(n, models.OneToMany)
|
||||
rel.FromTable = t.Name
|
||||
rel.FromSchema = t.Schema
|
||||
rel.FromColumns = append([]string(nil), c.Columns...)
|
||||
rel.ToTable = c.ReferencedTable
|
||||
rel.ToSchema = c.ReferencedSchema
|
||||
rel.ToColumns = append([]string(nil), c.ReferencedColumns...)
|
||||
rel.ForeignKey = c.Name
|
||||
t.Relationships[n] = rel
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package mysql
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||
)
|
||||
|
||||
func TestReaderMapDataType(t *testing.T) {
|
||||
r := NewReader(&readers.ReaderOptions{})
|
||||
for _, tc := range []struct{ input, want string }{{"varchar(64)", "string"}, {"bigint unsigned", "int64"}, {"datetime", "timestamp"}, {"json", "json"}} {
|
||||
if got := r.mapDataType(tc.input); got != tc.want {
|
||||
t.Errorf("mapDataType(%q) = %q, want %q", tc.input, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReaderRequiresConnectionString(t *testing.T) {
|
||||
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
|
||||
t.Fatal("expected missing connection string error")
|
||||
}
|
||||
}
|
||||
@@ -243,7 +243,11 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
|
||||
c.numeric_scale,
|
||||
c.udt_name,
|
||||
pg_catalog.format_type(a.atttypid, a.atttypmod) as formatted_data_type,
|
||||
col_description((c.table_schema||'.'||c.table_name)::regclass, c.ordinal_position) as description
|
||||
col_description((c.table_schema||'.'||c.table_name)::regclass, c.ordinal_position) as description,
|
||||
c.is_generated,
|
||||
c.generation_expression,
|
||||
c.is_identity,
|
||||
c.identity_generation
|
||||
FROM information_schema.columns c
|
||||
JOIN pg_catalog.pg_namespace n
|
||||
ON n.nspname = c.table_schema
|
||||
@@ -268,17 +272,36 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
|
||||
columnsMap := make(map[string]map[string]*models.Column)
|
||||
|
||||
for rows.Next() {
|
||||
var schema, tableName, columnName, isNullable, dataType, udtName, formattedDataType string
|
||||
var schema, tableName, columnName, isNullable, dataType, udtName, formattedDataType, isGenerated, isIdentity string
|
||||
var ordinalPosition int
|
||||
var columnDefault, description *string
|
||||
var columnDefault, description, generationExpression, identityGeneration *string
|
||||
var charMaxLength, numPrecision, numScale *int
|
||||
|
||||
if err := rows.Scan(&schema, &tableName, &columnName, &ordinalPosition, &columnDefault, &isNullable, &dataType, &charMaxLength, &numPrecision, &numScale, &udtName, &formattedDataType, &description); err != nil {
|
||||
if err := rows.Scan(&schema, &tableName, &columnName, &ordinalPosition, &columnDefault, &isNullable, &dataType, &charMaxLength, &numPrecision, &numScale, &udtName, &formattedDataType, &description, &isGenerated, &generationExpression, &isIdentity, &identityGeneration); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
column := models.InitColumn(columnName, tableName, schema)
|
||||
|
||||
// GENERATED ALWAYS ... STORED columns are computed by Postgres and cannot
|
||||
// have their default altered/dropped like a regular column.
|
||||
column.Generated = isGenerated == "ALWAYS"
|
||||
if generationExpression != nil {
|
||||
column.GenerationExpression = *generationExpression
|
||||
}
|
||||
|
||||
// GENERATED { ALWAYS | BY DEFAULT } AS IDENTITY columns are driven by an
|
||||
// internal sequence rather than a literal default (unlike serial columns,
|
||||
// they carry no pg_attrdef row at all), so they need the same DB-side
|
||||
// handling as generated columns even though the underlying mechanism differs.
|
||||
column.Identity = isIdentity == "YES"
|
||||
if identityGeneration != nil {
|
||||
column.IdentityGeneration = strings.ToUpper(strings.TrimSpace(*identityGeneration))
|
||||
}
|
||||
if column.Identity {
|
||||
column.AutoIncrement = true
|
||||
}
|
||||
|
||||
// Check if this is a serial type (has nextval default)
|
||||
hasNextval := false
|
||||
if columnDefault != nil {
|
||||
@@ -301,6 +324,24 @@ func (r *Reader) queryColumns(schemaName string) (map[string]map[string]*models.
|
||||
column.Description = *description
|
||||
}
|
||||
|
||||
if column.Generated {
|
||||
note := "GENERATED ALWAYS AS (" + column.GenerationExpression + ") STORED"
|
||||
if column.Description != "" {
|
||||
column.Description = column.Description + " " + note
|
||||
} else {
|
||||
column.Description = note
|
||||
}
|
||||
}
|
||||
|
||||
if column.Identity {
|
||||
note := "GENERATED " + column.IdentityGeneration + " AS IDENTITY"
|
||||
if column.Description != "" {
|
||||
column.Description = column.Description + " " + note
|
||||
} else {
|
||||
column.Description = note
|
||||
}
|
||||
}
|
||||
|
||||
if charMaxLength != nil {
|
||||
column.Length = *charMaxLength
|
||||
}
|
||||
@@ -381,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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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, ¬Null, &defaultValue, &pk); err != nil {
|
||||
if err := rows.Scan(&cid, &name, &dataType, ¬Null, &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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package sqltypes
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"math"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSqlNull_ValueScalarCases(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input SqlNull[any]
|
||||
want driver.Value
|
||||
}{
|
||||
{name: "invalid", input: SqlNull[any]{}, want: nil},
|
||||
{name: "integer", input: Null[any](int64(42), true), want: int64(42)},
|
||||
{name: "string", input: Null[any]("hello", true), want: "hello"},
|
||||
{name: "boolean", input: Null[any](true, true), want: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := tt.input.Value()
|
||||
if err != nil {
|
||||
t.Fatalf("Value returned error: %v", err)
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("Value() = %v (%T), want %v (%T)", got, got, tt.want, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlNull_Int64Conversions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input SqlNull[any]
|
||||
want int64
|
||||
}{
|
||||
{name: "invalid", input: SqlNull[any]{}, want: 0},
|
||||
{name: "signed integer", input: Null[any](int32(-12), true), want: -12},
|
||||
{name: "unsigned integer", input: Null[any](uint16(12), true), want: 12},
|
||||
{name: "float truncates", input: Null[any](float64(12.9), true), want: 12},
|
||||
{name: "numeric string", input: Null[any]("123", true), want: 123},
|
||||
{name: "invalid string", input: Null[any]("not a number", true), want: 0},
|
||||
{name: "true", input: Null[any](true, true), want: 1},
|
||||
{name: "false", input: Null[any](false, true), want: 0},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := tt.input.Int64(); got != tt.want {
|
||||
t.Errorf("Int64() = %d, want %d", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlNull_Float64Conversions(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input SqlNull[any]
|
||||
want float64
|
||||
}{
|
||||
{name: "invalid", input: SqlNull[any]{}, want: 0},
|
||||
{name: "float", input: Null[any](float32(1.25), true), want: 1.25},
|
||||
{name: "signed integer", input: Null[any](int64(-12), true), want: -12},
|
||||
{name: "unsigned integer", input: Null[any](uint16(12), true), want: 12},
|
||||
{name: "numeric string", input: Null[any]("12.5", true), want: 12.5},
|
||||
{name: "invalid string", input: Null[any]("not a number", true), want: 0},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := tt.input.Float64(); got != tt.want {
|
||||
t.Errorf("Float64() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlDate_JSONNullAndInvalid(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
json string
|
||||
valid bool
|
||||
}{
|
||||
{name: "null", json: "null", valid: false},
|
||||
{name: "invalid date", json: `"not-a-date"`, valid: false},
|
||||
{name: "valid date", json: `"2024-01-15"`, valid: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var got SqlDate
|
||||
if err := json.Unmarshal([]byte(tt.json), &got); err != nil {
|
||||
t.Fatalf("UnmarshalJSON returned error: %v", err)
|
||||
}
|
||||
if got.Valid != tt.valid {
|
||||
t.Errorf("Valid = %v, want %v", got.Valid, tt.valid)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
if data, err := json.Marshal(SqlDate{}); err != nil {
|
||||
t.Fatalf("MarshalJSON returned error: %v", err)
|
||||
} else if string(data) != "null" {
|
||||
t.Errorf("MarshalJSON() = %s, want null", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlTypeNowConstructors(t *testing.T) {
|
||||
before := time.Now()
|
||||
timestamp := SqlTimeStampNow()
|
||||
date := SqlDateNow()
|
||||
tm := SqlTimeNow()
|
||||
after := time.Now()
|
||||
|
||||
for name, got := range map[string]time.Time{
|
||||
"timestamp": timestamp.Time(),
|
||||
"date": date.Time(),
|
||||
"time": tm.Time(),
|
||||
} {
|
||||
if !got.After(before) && !got.Equal(before) || got.After(after) {
|
||||
t.Errorf("%s constructor returned %v outside [%v, %v]", name, got, before, after)
|
||||
}
|
||||
}
|
||||
if !timestamp.Valid || !date.Valid || !tm.Valid {
|
||||
t.Fatal("Now constructors must return valid values")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewSqlAndToJSONDT(t *testing.T) {
|
||||
if got := NewSql[int64]("42"); !got.Valid || got.Val != 42 {
|
||||
t.Errorf("NewSql[int64](\"42\") = %#v, want valid 42", got)
|
||||
}
|
||||
if got := NewSql[int64](nil); got.Valid {
|
||||
t.Errorf("NewSql[int64](nil) = %#v, want invalid", got)
|
||||
}
|
||||
if got := NewSqlFloat32(1.5); !got.Valid || got.Val != 1.5 {
|
||||
t.Errorf("NewSqlFloat32(1.5) = %#v, want valid 1.5", got)
|
||||
}
|
||||
|
||||
when := time.Date(2024, 1, 15, 10, 30, 45, 0, time.UTC)
|
||||
if got := ToJSONDT(when); got != "2024-01-15T10:30:45Z" {
|
||||
t.Errorf("ToJSONDT() = %q, want RFC3339 timestamp", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlNull_Float64PreservesInfinity(t *testing.T) {
|
||||
got := Null[float64](math.Inf(1), true).Float64()
|
||||
if !math.IsInf(got, 1) {
|
||||
t.Errorf("Float64() = %v, want +Inf", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package transform
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// Validation and normalization are currently pass-through stubs; these tests
|
||||
// pin that contract (no error, input returned unchanged).
|
||||
func TestTransformerStubs(t *testing.T) {
|
||||
tr := NewTransformer()
|
||||
if tr == nil {
|
||||
t.Fatal("nil transformer")
|
||||
}
|
||||
db := models.InitDatabase("d")
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("t", "public")
|
||||
|
||||
if err := tr.ValidateDatabase(db); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
if err := tr.ValidateSchema(schema); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
if err := tr.ValidateTable(table); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
if got, err := tr.NormalizeDatabase(db); err != nil || got != db {
|
||||
t.Errorf("NormalizeDatabase = %v, %v", got, err)
|
||||
}
|
||||
if got, err := tr.NormalizeSchema(schema); err != nil || got != schema {
|
||||
t.Errorf("NormalizeSchema = %v, %v", got, err)
|
||||
}
|
||||
if got, err := tr.NormalizeTable(table); err != nil || got != table {
|
||||
t.Errorf("NormalizeTable = %v, %v", got, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ConnKind identifies the database type a connection string targets.
|
||||
type ConnKind string
|
||||
|
||||
const (
|
||||
ConnPostgres ConnKind = "postgres"
|
||||
ConnMSSQL ConnKind = "mssql"
|
||||
ConnSQLite ConnKind = "sqlite"
|
||||
)
|
||||
|
||||
// connKinds lists the kinds offered by the builder dialog, in display order.
|
||||
var connKinds = []ConnKind{ConnPostgres, ConnMSSQL, ConnSQLite}
|
||||
|
||||
// maskedPassword is substituted for the password in previews.
|
||||
const maskedPassword = "****"
|
||||
|
||||
// ConnFields holds the editable parts of a connection string.
|
||||
type ConnFields struct {
|
||||
Kind ConnKind
|
||||
Host string
|
||||
Port string
|
||||
Database string
|
||||
User string
|
||||
Password string
|
||||
SSLMode string
|
||||
FilePath string // SQLite only
|
||||
|
||||
// Extra keeps query parameters the builder has no field for, so that
|
||||
// parsing and rebuilding an existing string does not drop them.
|
||||
Extra url.Values
|
||||
}
|
||||
|
||||
// DefaultConnFields returns sensible defaults for the given kind.
|
||||
func DefaultConnFields(kind ConnKind) ConnFields {
|
||||
f := ConnFields{Kind: kind}
|
||||
switch kind {
|
||||
case ConnPostgres:
|
||||
f.Host, f.Port, f.User, f.SSLMode = "localhost", "5432", "postgres", "disable"
|
||||
case ConnMSSQL:
|
||||
f.Host, f.Port, f.User, f.SSLMode = "localhost", "1433", "sa", "disable"
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
// SSLModes returns the valid SSL/encryption options for a kind.
|
||||
func SSLModes(kind ConnKind) []string {
|
||||
switch kind {
|
||||
case ConnPostgres:
|
||||
return []string{"disable", "allow", "prefer", "require", "verify-ca", "verify-full"}
|
||||
case ConnMSSQL:
|
||||
return []string{"disable", "false", "true"}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f ConnFields) sslParam() string {
|
||||
if f.Kind == ConnMSSQL {
|
||||
return "encrypt"
|
||||
}
|
||||
return "sslmode"
|
||||
}
|
||||
|
||||
// BuildConnString renders the fields as a connection string. With mask set,
|
||||
// a non-empty password is replaced by asterisks (for previews).
|
||||
func BuildConnString(f ConnFields, mask bool) string {
|
||||
if f.Kind == ConnSQLite {
|
||||
return f.FilePath
|
||||
}
|
||||
|
||||
u := &url.URL{Scheme: "postgres"}
|
||||
if f.Kind == ConnMSSQL {
|
||||
u.Scheme = "sqlserver"
|
||||
}
|
||||
|
||||
if f.Port != "" {
|
||||
u.Host = net.JoinHostPort(f.Host, f.Port)
|
||||
} else {
|
||||
u.Host = f.Host
|
||||
}
|
||||
|
||||
if f.User != "" {
|
||||
if f.Password != "" {
|
||||
pw := f.Password
|
||||
if mask {
|
||||
pw = maskedPassword
|
||||
}
|
||||
u.User = url.UserPassword(f.User, pw)
|
||||
} else {
|
||||
u.User = url.User(f.User)
|
||||
}
|
||||
}
|
||||
|
||||
query := url.Values{}
|
||||
for k, v := range f.Extra {
|
||||
query[k] = v
|
||||
}
|
||||
if f.Kind == ConnMSSQL {
|
||||
if f.Database != "" {
|
||||
query.Set("database", f.Database)
|
||||
}
|
||||
} else if f.Database != "" {
|
||||
u.Path = "/" + f.Database
|
||||
}
|
||||
if f.SSLMode != "" {
|
||||
query.Set(f.sslParam(), f.SSLMode)
|
||||
}
|
||||
u.RawQuery = query.Encode()
|
||||
|
||||
out := u.String()
|
||||
if mask {
|
||||
// url escapes '*' in the userinfo; keep the preview readable.
|
||||
out = strings.Replace(out, url.QueryEscape(maskedPassword), maskedPassword, 1)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// DetectConnKind guesses the kind from a connection string's scheme. Anything
|
||||
// that is not a recognised URL is treated as a SQLite file path.
|
||||
func DetectConnKind(s string) ConnKind {
|
||||
lower := strings.ToLower(strings.TrimSpace(s))
|
||||
switch {
|
||||
case strings.HasPrefix(lower, "postgres://"), strings.HasPrefix(lower, "postgresql://"):
|
||||
return ConnPostgres
|
||||
case strings.HasPrefix(lower, "sqlserver://"), strings.HasPrefix(lower, "mssql://"):
|
||||
return ConnMSSQL
|
||||
}
|
||||
return ConnSQLite
|
||||
}
|
||||
|
||||
// ParseConnString splits a connection string into fields. An empty string
|
||||
// yields the defaults for hint. Missing ports fall back to the kind default.
|
||||
func ParseConnString(s string, hint ConnKind) (ConnFields, error) {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return DefaultConnFields(hint), nil
|
||||
}
|
||||
|
||||
kind := DetectConnKind(s)
|
||||
if kind == ConnSQLite {
|
||||
path := s
|
||||
for _, prefix := range []string{"sqlite://", "sqlite3://"} {
|
||||
path = strings.TrimPrefix(path, prefix)
|
||||
}
|
||||
return ConnFields{Kind: ConnSQLite, FilePath: path}, nil
|
||||
}
|
||||
|
||||
u, err := url.Parse(s)
|
||||
if err != nil {
|
||||
return DefaultConnFields(kind), fmt.Errorf("invalid connection string: %w", err)
|
||||
}
|
||||
|
||||
f := ConnFields{
|
||||
Kind: kind,
|
||||
Host: u.Hostname(),
|
||||
Port: u.Port(),
|
||||
}
|
||||
if f.Port == "" {
|
||||
f.Port = DefaultConnFields(kind).Port
|
||||
}
|
||||
if u.User != nil {
|
||||
f.User = u.User.Username()
|
||||
f.Password, _ = u.User.Password()
|
||||
}
|
||||
|
||||
query := u.Query()
|
||||
if kind == ConnMSSQL {
|
||||
f.Database = query.Get("database")
|
||||
query.Del("database")
|
||||
} else {
|
||||
f.Database = strings.TrimPrefix(u.Path, "/")
|
||||
}
|
||||
f.SSLMode = query.Get(f.sslParam())
|
||||
query.Del(f.sslParam())
|
||||
if len(query) > 0 {
|
||||
f.Extra = query
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
_ "github.com/microsoft/go-mssqldb"
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
// connTestTimeout bounds how long "Test connection" may block.
|
||||
const connTestTimeout = 5 * time.Second
|
||||
|
||||
// TestConnection opens and pings the database described by f. Any occurrence
|
||||
// of the password in the returned error is masked.
|
||||
func TestConnection(f ConnFields) error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), connTestTimeout)
|
||||
defer cancel()
|
||||
|
||||
err := testConnection(ctx, f)
|
||||
if err != nil && f.Password != "" {
|
||||
err = fmt.Errorf("%s", strings.ReplaceAll(err.Error(), f.Password, maskedPassword))
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func testConnection(ctx context.Context, f ConnFields) error {
|
||||
switch f.Kind {
|
||||
case ConnPostgres:
|
||||
conn, err := pgx.Connect(ctx, BuildConnString(f, false))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return conn.Close(ctx)
|
||||
case ConnMSSQL:
|
||||
return pingSQL(ctx, "sqlserver", BuildConnString(f, false))
|
||||
case ConnSQLite:
|
||||
if f.FilePath == "" {
|
||||
return fmt.Errorf("file path is required")
|
||||
}
|
||||
// Opening a missing SQLite file would silently create it.
|
||||
if _, err := os.Stat(f.FilePath); err != nil {
|
||||
return err
|
||||
}
|
||||
return pingSQL(ctx, "sqlite", f.FilePath)
|
||||
}
|
||||
return fmt.Errorf("unsupported connection type %q", f.Kind)
|
||||
}
|
||||
|
||||
func pingSQL(ctx context.Context, driver, dsn string) error {
|
||||
db, err := sql.Open(driver, dsn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer db.Close()
|
||||
return db.PingContext(ctx)
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/gdamore/tcell/v2"
|
||||
"github.com/rivo/tview"
|
||||
)
|
||||
|
||||
// connBuilderPage is the page name of the connection string builder dialog.
|
||||
const connBuilderPage = "conn-builder"
|
||||
|
||||
// showConnStringBuilder opens the connection string builder, pre-filled by
|
||||
// parsing current. Save calls onDone with the built string; Esc/Back leaves
|
||||
// the caller's input untouched.
|
||||
func (se *SchemaEditor) showConnStringBuilder(current string, hint ConnKind, returnPage string, onDone func(connString string)) {
|
||||
fields, err := ParseConnString(current, hint)
|
||||
if err != nil {
|
||||
se.showErrorDialog("Error", err.Error()+"\nStarting from defaults.")
|
||||
}
|
||||
|
||||
title := tview.NewTextView().
|
||||
SetText("[::b]Connection String Builder").
|
||||
SetTextAlign(tview.AlignCenter).
|
||||
SetDynamicColors(true)
|
||||
|
||||
preview := tview.NewTextView()
|
||||
preview.SetBorder(true).SetTitle(" Preview (password masked) ").SetTitleAlign(tview.AlignLeft)
|
||||
|
||||
form := tview.NewForm()
|
||||
form.SetBorder(true).SetTitle(" Connection ").SetTitleAlign(tview.AlignLeft)
|
||||
|
||||
updatePreview := func() {
|
||||
preview.SetText(tview.Escape(BuildConnString(fields, true)))
|
||||
}
|
||||
|
||||
closeBuilder := func() {
|
||||
se.pages.RemovePage(connBuilderPage)
|
||||
se.pages.SwitchToPage(returnPage)
|
||||
}
|
||||
|
||||
var render func(focus int)
|
||||
render = func(focus int) {
|
||||
form.Clear(false)
|
||||
|
||||
kindIndex := 0
|
||||
kindLabels := make([]string, len(connKinds))
|
||||
for i, k := range connKinds {
|
||||
kindLabels[i] = string(k)
|
||||
if k == fields.Kind {
|
||||
kindIndex = i
|
||||
}
|
||||
}
|
||||
form.AddDropDown("Type", kindLabels, kindIndex, func(_ string, index int) {
|
||||
if connKinds[index] == fields.Kind {
|
||||
return
|
||||
}
|
||||
fields = DefaultConnFields(connKinds[index])
|
||||
render(0)
|
||||
})
|
||||
|
||||
if fields.Kind == ConnSQLite {
|
||||
form.AddInputField("File Path", fields.FilePath, 50, nil, func(v string) {
|
||||
fields.FilePath = v
|
||||
updatePreview()
|
||||
})
|
||||
if item, ok := form.GetFormItemByLabel("File Path").(*tview.InputField); ok {
|
||||
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() != tcell.KeyEnter {
|
||||
return event
|
||||
}
|
||||
se.showFileBrowser(FileBrowserConfig{
|
||||
Mode: FileBrowserLoad,
|
||||
StartPath: fields.FilePath,
|
||||
Extensions: FormatExtensions("sqlite"),
|
||||
ReturnPage: connBuilderPage,
|
||||
OnSelect: func(path string) { item.SetText(path) },
|
||||
})
|
||||
return nil
|
||||
})
|
||||
}
|
||||
} else {
|
||||
form.AddInputField("Host", fields.Host, 50, nil, func(v string) { fields.Host = v; updatePreview() })
|
||||
form.AddInputField("Port", fields.Port, 10, tview.InputFieldInteger, func(v string) { fields.Port = v; updatePreview() })
|
||||
form.AddInputField("Database", fields.Database, 50, nil, func(v string) { fields.Database = v; updatePreview() })
|
||||
form.AddInputField("User", fields.User, 50, nil, func(v string) { fields.User = v; updatePreview() })
|
||||
form.AddPasswordField("Password", fields.Password, 50, '*', func(v string) { fields.Password = v; updatePreview() })
|
||||
|
||||
label := "SSL Mode"
|
||||
if fields.Kind == ConnMSSQL {
|
||||
label = "Encrypt"
|
||||
}
|
||||
modes := SSLModes(fields.Kind)
|
||||
modeIndex := -1
|
||||
for i, m := range modes {
|
||||
if m == fields.SSLMode {
|
||||
modeIndex = i
|
||||
}
|
||||
}
|
||||
if modeIndex < 0 {
|
||||
// Keep a value parsed from an existing string even if it is not a listed option.
|
||||
modes = append([]string{fields.SSLMode}, modes...)
|
||||
modeIndex = 0
|
||||
}
|
||||
form.AddDropDown(label, modes, modeIndex, func(option string, _ int) {
|
||||
fields.SSLMode = option
|
||||
updatePreview()
|
||||
})
|
||||
}
|
||||
|
||||
form.AddButton("Save [F2]", connBuilderSave(se, &fields, closeBuilder, onDone))
|
||||
form.AddButton("Test [F3]", func() { se.testConnectionDialog(fields) })
|
||||
form.AddButton("Back [Esc]", closeBuilder)
|
||||
|
||||
updatePreview()
|
||||
form.SetFocus(focus)
|
||||
se.app.SetFocus(form)
|
||||
}
|
||||
|
||||
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
switch event.Key() {
|
||||
case tcell.KeyEscape:
|
||||
closeBuilder()
|
||||
return nil
|
||||
case tcell.KeyF2:
|
||||
connBuilderSave(se, &fields, closeBuilder, onDone)()
|
||||
return nil
|
||||
case tcell.KeyF3:
|
||||
se.testConnectionDialog(fields)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
|
||||
render(0)
|
||||
|
||||
flex := tview.NewFlex().SetDirection(tview.FlexRow).
|
||||
AddItem(title, 1, 0, false).
|
||||
AddItem(form, 0, 1, true).
|
||||
AddItem(preview, 4, 0, false)
|
||||
|
||||
se.pages.AddAndSwitchToPage(connBuilderPage, flex, true)
|
||||
se.app.SetFocus(form)
|
||||
}
|
||||
|
||||
// connBuilderSave returns the Save action: validate, write back, close.
|
||||
func connBuilderSave(se *SchemaEditor, fields *ConnFields, closeBuilder func(), onDone func(string)) func() {
|
||||
return func() {
|
||||
if msg := validateConnFields(*fields); msg != "" {
|
||||
se.showErrorDialog("Error", msg)
|
||||
return
|
||||
}
|
||||
result := BuildConnString(*fields, false)
|
||||
closeBuilder()
|
||||
onDone(result)
|
||||
}
|
||||
}
|
||||
|
||||
// validateConnFields returns a message describing the first missing required field, or "".
|
||||
func validateConnFields(f ConnFields) string {
|
||||
if f.Kind == ConnSQLite {
|
||||
if strings.TrimSpace(f.FilePath) == "" {
|
||||
return "File path is required"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
if strings.TrimSpace(f.Host) == "" {
|
||||
return "Host is required"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// testConnectionDialog runs TestConnection in the background and reports the result.
|
||||
func (se *SchemaEditor) testConnectionDialog(fields ConnFields) {
|
||||
if msg := validateConnFields(fields); msg != "" {
|
||||
se.showErrorDialog("Error", msg)
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
err := TestConnection(fields)
|
||||
se.app.QueueUpdateDraw(func() {
|
||||
if err != nil {
|
||||
se.showErrorDialog("Connection Failed", fmt.Sprintf("Connection failed:\n%v", err))
|
||||
return
|
||||
}
|
||||
se.showSuccessDialog("Connection OK", "Connection successful", nil)
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
// attachConnStringBuilder makes Enter on the named input open the builder.
|
||||
func (se *SchemaEditor) attachConnStringBuilder(form *tview.Form, label, returnPage string, format func() string) {
|
||||
item, ok := form.GetFormItemByLabel(label).(*tview.InputField)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() != tcell.KeyEnter {
|
||||
return event
|
||||
}
|
||||
hint := ConnPostgres
|
||||
if format != nil && format() == "sqlite" {
|
||||
hint = ConnSQLite
|
||||
}
|
||||
se.showConnStringBuilder(item.GetText(), hint, returnPage, func(s string) { item.SetText(s) })
|
||||
return nil
|
||||
})
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user