Compare commits
48
Commits
278d488363
...
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 |
@@ -78,6 +78,13 @@ jobs:
|
||||
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 }}"
|
||||
|
||||
+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
|
||||
@@ -243,7 +243,7 @@ release: lint fmt-check test build ## Run lint, format check, tests, build, then
|
||||
git push origin "$$version"; \
|
||||
echo "Tag $$version created and pushed to remote repository."
|
||||
|
||||
release-version: ## Auto-increment patch version, update VERSION/package files, commit, tag, and push
|
||||
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); \
|
||||
@@ -260,5 +260,18 @@ release-version: ## Auto-increment patch version, update VERSION/package files,
|
||||
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
|
||||
@@ -172,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+
|
||||
@@ -182,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
|
||||
@@ -196,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)
|
||||
```
|
||||
|
||||
@@ -216,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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+140
-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())
|
||||
|
||||
@@ -240,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)
|
||||
@@ -380,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
|
||||
@@ -408,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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -465,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)
|
||||
|
||||
@@ -529,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)
|
||||
}
|
||||
}
|
||||
+1
-1
@@ -844,7 +844,7 @@ func writeJobOutput(rj *resolvedJob, db *models.Database, lg *jobLogger) error {
|
||||
}
|
||||
lg.logf("writing output to database env:%s", rj.outputConnEnv)
|
||||
writerOpts := newWriterOptions("", o.Package, o.FlattenSchema, o.Types, o.ArrayNullable, o.ContinueOnError)
|
||||
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn}
|
||||
writerOpts.Metadata = map[string]interface{}{"connection_string": rj.outputConn, "full_ddl": o.FullDDL}
|
||||
return wpgsql.NewWriter(writerOpts).WriteDatabase(db)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)")
|
||||
}
|
||||
|
||||
@@ -262,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)
|
||||
@@ -283,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
|
||||
|
||||
@@ -443,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,
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
||||
@@ -17,6 +18,8 @@ var (
|
||||
noVersion bool
|
||||
silent bool
|
||||
strictDirectives bool
|
||||
typeMapFlags []string
|
||||
typeMappings map[string]string
|
||||
)
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
@@ -28,10 +31,16 @@ bidirectional conversion between various database schema formats.
|
||||
It reads database schemas from multiple sources (live databases, DBML,
|
||||
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
||||
JSON, YAML, SQL, etc.).`,
|
||||
PersistentPreRunE: func(cmd *cobra.Command, args []string) error {
|
||||
var err error
|
||||
typeMappings, err = writers.ParseTypeMappings(typeMapFlags)
|
||||
return err
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
rootCmd.AddCommand(convertCmd)
|
||||
rootCmd.AddCommand(batchCmd)
|
||||
rootCmd.AddCommand(diffCmd)
|
||||
rootCmd.AddCommand(inspectCmd)
|
||||
rootCmd.AddCommand(scriptsCmd)
|
||||
@@ -42,8 +51,10 @@ func init() {
|
||||
rootCmd.AddCommand(mergeCmd)
|
||||
rootCmd.AddCommand(splitCmd)
|
||||
rootCmd.AddCommand(versionCmd)
|
||||
rootCmd.AddCommand(updateCmd)
|
||||
rootCmd.AddCommand(reportCmd)
|
||||
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||
rootCmd.PersistentFlags().StringArrayVar(&typeMapFlags, "type-map", nil, "Override a SQL-to-Go type mapping for bun/gorm output as sqltype=gotype (repeatable), e.g. --type-map uuid=uuid.UUID --type-map numeric=decimal.Decimal")
|
||||
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
||||
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
||||
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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).
|
||||
@@ -187,6 +187,8 @@ 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
|
||||
@@ -195,6 +197,7 @@ jobs:
|
||||
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
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
# TUI mouse support plan
|
||||
|
||||
Issue: #46
|
||||
Status: design only; this document does not implement mouse input.
|
||||
|
||||
## 1. Current implementation and scope
|
||||
|
||||
The editor is created in `cmd/relspec/edit.go` by
|
||||
`ui.NewSchemaEditorWithConfigs(...).Run()`. `pkg/ui/editor.go` owns the
|
||||
`tview.Application`, `tview.Pages`, and application lifecycle. The current
|
||||
code never calls `Application.EnableMouse`, so tcell mouse reporting is off.
|
||||
The module uses tview v0.42.0 and tcell/v2 v2.13.9.
|
||||
|
||||
The first implementation should add `--no-mouse` to the `edit` Cobra command
|
||||
only. The flag is a local boolean, defaulting to false, and should be passed
|
||||
explicitly into the editor (prefer an options/config field rather than a
|
||||
package-global or environment variable). It must not affect convert, inspect,
|
||||
merge, or other commands. There is no environment-variable or persistent
|
||||
configuration setting in this issue: a command-line opt-out is predictable,
|
||||
visible in `edit --help`, and avoids adding configuration precedence rules.
|
||||
|
||||
At startup, the editor should call `app.EnableMouse(!noMouse)` before
|
||||
`Run()`. `--no-mouse` must mean that the application does not enable terminal
|
||||
mouse reporting and that no custom mouse handlers are relied upon. Keyboard
|
||||
behavior must remain identical in both modes.
|
||||
|
||||
Likely implementation files are `cmd/relspec/edit.go`,
|
||||
`pkg/ui/editor.go`, focused TUI mouse helpers/tests under `pkg/ui`, and a
|
||||
short user-facing note in the command help or TUI documentation. Do not
|
||||
refactor unrelated screens or data operations.
|
||||
|
||||
## 2. Widget and screen coverage
|
||||
|
||||
The application composes `Pages`, `Flex`, `TextView`, `List`, `Table`, `Form`,
|
||||
`Button`, `InputField`, `DropDown`, `TextArea`, `CheckBox`, and `Modal`.
|
||||
Vendored tview confirms mouse handlers exist for all of those relevant
|
||||
primitives, including focus on left-down, list/table selection, button clicks,
|
||||
form child dispatch, dropdown opening/drag selection, text-area cursor and
|
||||
scrolling, and modal button dispatch. `Pages`, `Flex`, and `Form` forward events
|
||||
to their children.
|
||||
|
||||
The coverage plan is:
|
||||
|
||||
* Main menu (`pkg/ui/main_menu.go`): left click focuses/selects a list entry;
|
||||
second activation opens it; buttons and exit confirmation remain reachable.
|
||||
* Schema, table, domain, object, relation, and database screens: click a row
|
||||
to select it; double-click the row to perform the same action as the
|
||||
keyboard Enter/selected callback where opening is meaningful; scroll lists
|
||||
and tables; click each action button.
|
||||
* Tables (`schema_screens.go`, `table_screens.go`, and object/relation tables):
|
||||
tview's table handler provides selection and scrolling, but it does not
|
||||
provide application-specific double-click activation. Add a small reusable
|
||||
wrapper/helper for the table instances that need it. It must preserve the
|
||||
existing selected row/column behavior and invoke the same callback as Enter,
|
||||
not duplicate mutation logic.
|
||||
* Forms (`load_save_screens.go` and the form-building screen files): click an
|
||||
input to focus it, click buttons to activate them, click a dropdown to open
|
||||
it and choose an option, scroll multiline help/text areas, and retain all
|
||||
existing keyboard Tab/Shift-Tab, shortcut, Enter, and Escape behavior.
|
||||
* Dialogs (`pkg/ui/dialogs.go` plus confirmation/error/success modals): modal
|
||||
buttons are clickable and the modal keeps focus above the underlying page.
|
||||
Clicking outside a modal must not activate the hidden page or dismiss a
|
||||
destructive confirmation. Escape and the existing button-key behavior stay
|
||||
authoritative.
|
||||
* The planned file browser and connection-string builder from issue #44 must
|
||||
use the same contracts: clickable entries/buttons and scrolling, with
|
||||
keyboard navigation and explicit cancel/accept paths. #46 should not
|
||||
implement #44's widgets; it should define the integration point and test
|
||||
them when #44 lands.
|
||||
|
||||
Do not promise drag semantics for every widget. Drag is appropriate for text
|
||||
selection/cursor movement and dropdown selection where tview already supports
|
||||
it. For ordinary list/table navigation, a click selects and the wheel scrolls;
|
||||
row dragging should not mutate data.
|
||||
|
||||
## 3. Exact mouse action contract
|
||||
|
||||
| Widget/type | Left down/click | Double click | Wheel/drag | Keyboard fallback |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| Main/list menu | focus and select row | invoke row selected callback | scroll list | arrows, Enter, shortcuts |
|
||||
| Data table | focus and select cell/row | invoke the screen's existing open/edit action for the selected row | vertical/horizontal scroll as supported by tview | arrows, PageUp/PageDown, Enter, existing shortcuts |
|
||||
| Button | focus | same as one activation, never duplicate the callback | none | Tab/Shift-Tab, Enter/Space and existing shortcut |
|
||||
| Input field | focus; place cursor if supported | no destructive action | text-area behavior if provided by tview | typing, arrows, Home/End, Tab, Escape |
|
||||
| Text area/help | focus and position cursor | select word only where tview supports it; no application action | scroll; drag text selection if supported | arrows, PageUp/PageDown, standard editing keys |
|
||||
| Dropdown | focus/open and choose the hit option | same as click; no duplicate selection | drag through options only while open | arrows, Enter, Escape, Tab |
|
||||
| Checkbox | toggle on click | no second toggle | none | Space and existing form navigation |
|
||||
| Modal | focus/click visible button | same button action once | no underlying-page scrolling | Tab/Shift-Tab, Enter, Escape, existing button keys |
|
||||
| Blank/border/title area | focus containing primitive where useful | none | no mutation | current screen shortcuts |
|
||||
|
||||
Right and middle clicks should have no application action in the first
|
||||
release. Wheel events should be consumed only by the scrollable primitive
|
||||
under the pointer. Double-click timing/translation should come from tview/
|
||||
tcell; custom code must not fire the action once for both the click and the
|
||||
double-click. Any custom table wrapper needs a small state machine or tview
|
||||
mouse action handling that is tested for this property.
|
||||
|
||||
## 4. tview gaps and implementation boundaries
|
||||
|
||||
Enabling mouse support is not sufficient for the desired behavior. tview's
|
||||
built-in Table handler selects cells and scrolls but has no repository-specific
|
||||
row-open callback on double click. Existing screen code also wires keyboard
|
||||
input captures directly on individual widgets, so mouse actions must call the
|
||||
same screen callbacks rather than route through synthetic key events.
|
||||
|
||||
Use tview's `MouseHandler`/`WrapMouseHandler` contracts and `setFocus` rather
|
||||
than reading terminal coordinates in each screen. A reusable table adapter
|
||||
may embed `*tview.Table`, delegate ordinary actions to the original table
|
||||
handler, and add the screen's double-click callback. Keep the adapter in
|
||||
`pkg/ui` and use it only where a row-opening action exists. Do not modify the
|
||||
vendored tview copy.
|
||||
|
||||
The `Pages`/`Modal` dispatch order must be verified: a visible modal consumes
|
||||
its click before the page below it. Page transitions should happen only in the
|
||||
existing callbacks, so a stale hidden page cannot receive a click.
|
||||
|
||||
## 5. Keyboard, terminal, and copy/paste behavior
|
||||
|
||||
Mouse is an enhancement, never a requirement. Every acceptance path must be
|
||||
reachable with the existing keyboard controls, including load/save, navigation,
|
||||
editing, confirmations, cancel, and exit. `--no-mouse` is the regression mode
|
||||
for proving this contract.
|
||||
|
||||
Mouse reporting is terminal capability dependent. On local terminals it is
|
||||
negotiated by tcell; tmux and SSH can suppress, translate, or fail to pass
|
||||
mouse reporting depending on their configuration. The application must still
|
||||
start and remain keyboard usable if mouse reporting is unavailable or broken.
|
||||
Documentation should state that terminal/tmux configuration may be required,
|
||||
and that SSH behavior depends on the remote terminal path. Windows Terminal and
|
||||
other Windows console hosts should be treated as supported only insofar as the
|
||||
selected tcell backend reports mouse events; the CLI must not assume POSIX
|
||||
escape sequences or add platform-specific code in this issue.
|
||||
|
||||
Enabling mouse capture normally prevents terminal-native selection/copy from
|
||||
seeing ordinary button-drag events. Document the standard workaround: hold the
|
||||
terminal's bypass modifier (commonly Shift, terminal-dependent) for selection,
|
||||
or use `--no-mouse` when native copy/paste is the priority. Do not implement a
|
||||
second clipboard protocol. Input-field/text-area copy/paste must continue to
|
||||
use tview/tcell paste handling and keyboard shortcuts; verify that enabling
|
||||
mouse does not intercept paste events.
|
||||
|
||||
## 6. Test strategy using tcell simulation
|
||||
|
||||
Add focused tests rather than attempting a full interactive end-to-end test.
|
||||
Use `tcell.NewSimulationScreen("")`, `screen.Init()`, construct the editor or
|
||||
an isolated primitive tree, and inject events with the actual API:
|
||||
`SimulationScreen.InjectMouse(x, y, buttons, mod)` and `InjectKey(...)`.
|
||||
Coordinates must be derived from the primitive's drawn rectangle or fixed by a
|
||||
small deterministic test layout; do not use arbitrary coordinates without
|
||||
checking the rendered screen.
|
||||
|
||||
Minimum cases:
|
||||
|
||||
1. Default editor configuration enables mouse; the explicit disabled option
|
||||
leaves it disabled. If the Application API is not observable directly,
|
||||
test through the simulation screen's event path plus a constructor-level
|
||||
option assertion.
|
||||
2. A list click changes focus/selection, and double-click invokes the existing
|
||||
selected action exactly once.
|
||||
3. A table click selects the expected row/cell, wheel events change the visible
|
||||
offset, and double-click invokes the row action exactly once.
|
||||
4. Form button, input field, checkbox, and dropdown clicks match their
|
||||
keyboard callbacks.
|
||||
5. A modal button click acts on the modal and cannot activate the underlying
|
||||
page; Escape still cancels.
|
||||
6. `--no-mouse` leaves keyboard selection/activation unchanged and mouse
|
||||
injection has no application effect.
|
||||
7. Existing dialogs and screen transitions do not leave a stale mouse capture
|
||||
after a page is removed.
|
||||
|
||||
Prefer callback counters and selected-index assertions over screen-text-only
|
||||
assertions. Run the relevant `pkg/ui` tests with `go test -race ./pkg/ui` and
|
||||
run the full repository test suite if time/resources permit.
|
||||
|
||||
## 7. Rollout and acceptance criteria
|
||||
|
||||
Implementation is ready for review when:
|
||||
|
||||
* `relspec edit --help` documents `--no-mouse` and mouse is enabled by default.
|
||||
* Only the edit TUI is affected; non-TUI commands have no changed behavior.
|
||||
* Main screens, tables, lists, forms, dropdowns, buttons, text areas, and
|
||||
visible dialogs support the action contract above.
|
||||
* Keyboard-only operation is complete and verified with `--no-mouse`.
|
||||
* Modal clicks cannot fall through to an underlying page.
|
||||
* Table double-click behavior is explicit, tested, and does not duplicate
|
||||
activation.
|
||||
* tcell simulation tests cover default-on, opt-out, selection, scrolling,
|
||||
activation, dialog focus, and keyboard fallback.
|
||||
* `go test -race ./pkg/ui`, appropriate command tests, `go test ./...`,
|
||||
formatting, and `git diff --check` pass (or any limitation is recorded).
|
||||
* User-facing docs explain tmux/SSH/Windows variability and the terminal
|
||||
modifier workaround for native copy/paste.
|
||||
|
||||
Roll out in two implementation slices if needed: first application option,
|
||||
standard tview handlers, tests, and documentation; second only the reusable
|
||||
table double-click adapter and screen wiring. Do not block the first slice on
|
||||
issue #44, but do not claim #44's future widgets are covered until they use the
|
||||
same contract.
|
||||
|
||||
## 8. Open decisions and dependencies
|
||||
|
||||
* Confirm whether the project wants a public editor options type or a small
|
||||
`SetMouseEnabled`/constructor parameter; avoid a global flag.
|
||||
* Confirm the preferred terminal copy modifier in project documentation, since
|
||||
tmux, SSH clients, and Windows Terminal differ.
|
||||
* Decide whether horizontal wheel events should be supported where tview/table
|
||||
exposes them; vertical scrolling is mandatory, horizontal is optional.
|
||||
* Decide whether double-click opens every data table or only tables with an
|
||||
unambiguous row action. The plan recommends the latter.
|
||||
* Confirm #44's file-browser and connection-builder primitive choices before
|
||||
wiring their mouse tests.
|
||||
* Confirm CI has a stable non-terminal environment for simulation-screen tests;
|
||||
no real terminal, tmux session, database, or network should be required.
|
||||
@@ -4,6 +4,7 @@ go 1.25.13
|
||||
|
||||
require (
|
||||
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=
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
# Maintainer: Hein (Warky Devs) <hein@warky.dev>
|
||||
pkgname=relspec
|
||||
pkgver=1.0.82
|
||||
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')
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
Name: relspec
|
||||
Version: 1.0.82
|
||||
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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
+8
-5
@@ -285,11 +285,14 @@ type Options struct {
|
||||
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:
|
||||
|
||||
@@ -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++
|
||||
}
|
||||
@@ -440,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),
|
||||
@@ -469,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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -157,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,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
|
||||
}
|
||||
@@ -804,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:
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -422,31 +422,35 @@ func (r *Reader) queryPrimaryKeys(schemaName string) (map[string]*models.Constra
|
||||
// queryForeignKeys retrieves all foreign key constraints for a schema
|
||||
// Returns map[schema.table][]*Constraint
|
||||
func (r *Reader) queryForeignKeys(schemaName string) (map[string][]*models.Constraint, error) {
|
||||
// Columns are paired by position from pg_constraint so composite keys are not
|
||||
// cross-joined (information_schema.constraint_column_usage has no ordering).
|
||||
actionCase := func(col string) string {
|
||||
return `CASE ` + col + ` WHEN 'a' THEN 'NO ACTION' WHEN 'r' THEN 'RESTRICT' WHEN 'c' THEN 'CASCADE' WHEN 'n' THEN 'SET NULL' WHEN 'd' THEN 'SET DEFAULT' END`
|
||||
}
|
||||
query := `
|
||||
SELECT
|
||||
tc.table_schema,
|
||||
tc.table_name,
|
||||
tc.constraint_name,
|
||||
kcu.table_schema as foreign_table_schema,
|
||||
kcu.table_name as foreign_table_name,
|
||||
kcu.column_name as foreign_column,
|
||||
ccu.table_schema as referenced_table_schema,
|
||||
ccu.table_name as referenced_table_name,
|
||||
ccu.column_name as referenced_column,
|
||||
rc.update_rule,
|
||||
rc.delete_rule
|
||||
FROM information_schema.table_constraints tc
|
||||
JOIN information_schema.key_column_usage kcu
|
||||
ON tc.constraint_name = kcu.constraint_name
|
||||
AND tc.table_schema = kcu.table_schema
|
||||
JOIN information_schema.constraint_column_usage ccu
|
||||
ON ccu.constraint_name = tc.constraint_name
|
||||
JOIN information_schema.referential_constraints rc
|
||||
ON rc.constraint_name = tc.constraint_name
|
||||
AND rc.constraint_schema = tc.table_schema
|
||||
WHERE tc.constraint_type = 'FOREIGN KEY'
|
||||
AND tc.table_schema = $1
|
||||
ORDER BY tc.table_schema, tc.table_name, tc.constraint_name, kcu.ordinal_position
|
||||
ns.nspname AS table_schema,
|
||||
cl.relname AS table_name,
|
||||
con.conname AS constraint_name,
|
||||
ns.nspname AS foreign_table_schema,
|
||||
cl.relname AS foreign_table_name,
|
||||
att.attname AS foreign_column,
|
||||
fns.nspname AS referenced_table_schema,
|
||||
fcl.relname AS referenced_table_name,
|
||||
fatt.attname AS referenced_column,
|
||||
` + actionCase("con.confupdtype") + ` AS update_rule,
|
||||
` + actionCase("con.confdeltype") + ` AS delete_rule
|
||||
FROM pg_catalog.pg_constraint con
|
||||
JOIN pg_catalog.pg_class cl ON cl.oid = con.conrelid
|
||||
JOIN pg_catalog.pg_namespace ns ON ns.oid = cl.relnamespace
|
||||
JOIN pg_catalog.pg_class fcl ON fcl.oid = con.confrelid
|
||||
JOIN pg_catalog.pg_namespace fns ON fns.oid = fcl.relnamespace
|
||||
CROSS JOIN LATERAL unnest(con.conkey, con.confkey) WITH ORDINALITY AS k(attnum, fattnum, ord)
|
||||
JOIN pg_catalog.pg_attribute att ON att.attrelid = con.conrelid AND att.attnum = k.attnum
|
||||
JOIN pg_catalog.pg_attribute fatt ON fatt.attrelid = con.confrelid AND fatt.attnum = k.fattnum
|
||||
WHERE con.contype = 'f'
|
||||
AND ns.nspname = $1
|
||||
ORDER BY ns.nspname, cl.relname, con.conname, k.ord
|
||||
`
|
||||
|
||||
rows, err := r.conn.Query(r.ctx, query, schemaName)
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildConnString(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
fields ConnFields
|
||||
mask bool
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "postgres defaults with db",
|
||||
fields: func() ConnFields { f := DefaultConnFields(ConnPostgres); f.Database = "app"; return f }(),
|
||||
want: "postgres://postgres@localhost:5432/app?sslmode=disable",
|
||||
},
|
||||
{
|
||||
name: "postgres password unmasked",
|
||||
fields: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5433", Database: "x", User: "u", Password: "p@ss/w", SSLMode: "require"},
|
||||
want: "postgres://u:p%40ss%2Fw@db:5433/x?sslmode=require",
|
||||
},
|
||||
{
|
||||
name: "postgres password masked",
|
||||
fields: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5432", Database: "x", User: "u", Password: "secret"},
|
||||
mask: true,
|
||||
want: "postgres://u:****@db:5432/x",
|
||||
},
|
||||
{
|
||||
name: "mssql",
|
||||
fields: ConnFields{Kind: ConnMSSQL, Host: "sql", Port: "1433", Database: "shop", User: "sa", Password: "pw", SSLMode: "disable"},
|
||||
want: "sqlserver://sa:pw@sql:1433?database=shop&encrypt=disable",
|
||||
},
|
||||
{
|
||||
name: "sqlite is the plain path",
|
||||
fields: ConnFields{Kind: ConnSQLite, FilePath: "/tmp/a b.db"},
|
||||
want: "/tmp/a b.db",
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := BuildConnString(tt.fields, tt.mask); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaskedBuildHidesPassword(t *testing.T) {
|
||||
f := ConnFields{Kind: ConnMSSQL, Host: "h", User: "u", Password: "hunter2"}
|
||||
if got := BuildConnString(f, true); strings.Contains(got, "hunter2") {
|
||||
t.Errorf("masked string leaks password: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConnString(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
want ConnFields
|
||||
}{
|
||||
{
|
||||
name: "postgres full",
|
||||
in: "postgres://u:p%40ss@db:5433/app?sslmode=require&application_name=x",
|
||||
want: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5433", Database: "app", User: "u", Password: "p@ss", SSLMode: "require"},
|
||||
},
|
||||
{
|
||||
name: "postgresql scheme, default port",
|
||||
in: "postgresql://u@db/app",
|
||||
want: ConnFields{Kind: ConnPostgres, Host: "db", Port: "5432", Database: "app", User: "u"},
|
||||
},
|
||||
{
|
||||
name: "mssql",
|
||||
in: "sqlserver://sa:pw@sql:1444?database=shop&encrypt=true",
|
||||
want: ConnFields{Kind: ConnMSSQL, Host: "sql", Port: "1444", Database: "shop", User: "sa", Password: "pw", SSLMode: "true"},
|
||||
},
|
||||
{
|
||||
name: "sqlite path",
|
||||
in: "/data/app.db",
|
||||
want: ConnFields{Kind: ConnSQLite, FilePath: "/data/app.db"},
|
||||
},
|
||||
{
|
||||
name: "sqlite scheme",
|
||||
in: "sqlite:///data/app.db",
|
||||
want: ConnFields{Kind: ConnSQLite, FilePath: "/data/app.db"},
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := ParseConnString(tt.in, ConnPostgres)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got.Extra = nil
|
||||
if !reflect.DeepEqual(got, tt.want) {
|
||||
t.Errorf("got %+v, want %+v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConnStringEmptyUsesHintDefaults(t *testing.T) {
|
||||
got, err := ParseConnString(" ", ConnMSSQL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.Kind != ConnMSSQL || got.Port != "1433" || got.Host != "localhost" {
|
||||
t.Errorf("unexpected defaults: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseConnStringInvalid(t *testing.T) {
|
||||
if _, err := ParseConnString("postgres://u:p@host:badport/db", ConnPostgres); err == nil {
|
||||
t.Error("expected error for invalid port")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnStringRoundTrip(t *testing.T) {
|
||||
for _, in := range []string{
|
||||
"postgres://u:pw@db:5433/app?application_name=x&sslmode=require",
|
||||
"sqlserver://sa:pw@sql:1433?application+name=x&database=shop&encrypt=false",
|
||||
} {
|
||||
f, err := ParseConnString(in, ConnPostgres)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := BuildConnString(f, false); got != in {
|
||||
t.Errorf("round trip: got %q, want %q", got, in)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTestConnectionSQLite(t *testing.T) {
|
||||
if err := TestConnection(ConnFields{Kind: ConnSQLite}); err == nil {
|
||||
t.Error("expected error for empty path")
|
||||
}
|
||||
if err := TestConnection(ConnFields{Kind: ConnSQLite, FilePath: t.TempDir() + "/missing.db"}); err == nil {
|
||||
t.Error("expected error for missing file")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func TestColumnDataOps(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
|
||||
if se.CreateColumn(5, 0, "x", "int", false, false) != nil || se.CreateColumn(0, 5, "x", "int", false, false) != nil {
|
||||
t.Error("create with bad index must return nil")
|
||||
}
|
||||
col := se.CreateColumn(0, 0, "age", "integer", true, true)
|
||||
if col == nil || col.Type != "integer" || !col.IsPrimaryKey || !col.NotNull {
|
||||
t.Fatalf("create: %+v", col)
|
||||
}
|
||||
if se.GetColumn(0, 0, "age") != col || se.GetColumn(0, 0, "nope") != nil || se.GetColumn(9, 0, "age") != nil {
|
||||
t.Error("get mismatch")
|
||||
}
|
||||
|
||||
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
|
||||
t.Error("create second column")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
si, ti int
|
||||
old, new string
|
||||
want bool
|
||||
}{
|
||||
{"bad table", 0, 9, "age", "age", false},
|
||||
{"missing column", 0, 0, "zzz", "zzz", false},
|
||||
{"in place", 0, 0, "age", "age", true},
|
||||
{"rename", 0, 0, "age", "years", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := se.UpdateColumn(tt.si, tt.ti, tt.old, tt.new, "bigint", false, true, "0", "desc"); got != tt.want {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
got := se.GetColumn(0, 0, "years")
|
||||
if got == nil || got.Name != "years" || got.Type != "bigint" || got.IsPrimaryKey || got.Default != "0" || got.Description != "desc" {
|
||||
t.Errorf("after update: %+v", got)
|
||||
}
|
||||
if se.GetColumn(0, 0, "age") != nil {
|
||||
t.Error("old name must be gone")
|
||||
}
|
||||
|
||||
if len(se.GetAllColumns(0, 0)) != 4 || se.GetAllColumns(0, 9) != nil {
|
||||
t.Error("GetAllColumns")
|
||||
}
|
||||
if se.DeleteColumn(0, 9, "a") || se.DeleteColumn(0, 0, "zzz") {
|
||||
t.Error("delete bad target must fail")
|
||||
}
|
||||
if !se.DeleteColumn(0, 0, "a") || se.DeleteColumn(0, 0, "a") {
|
||||
t.Error("delete should succeed once")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateColumn_NilMap(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
se.db.Schemas[0].Tables[0].Columns = nil
|
||||
if se.CreateColumn(0, 0, "a", "text", false, false) == nil {
|
||||
t.Error("create with nil map")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelationshipDataOps(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
rel := &models.Relationship{Name: "fk_a", FromTable: "users", ToTable: "orders"}
|
||||
|
||||
if se.CreateRelationship(9, 0, rel) != nil || se.CreateRelationship(0, 9, rel) != nil || se.CreateRelationship(0, -1, rel) != nil {
|
||||
t.Error("create bad index")
|
||||
}
|
||||
// Before any relationship exists, update/delete/get/names report nothing.
|
||||
se.db.Schemas[0].Tables[0].Relationships = nil
|
||||
if se.UpdateRelationship(0, 0, "fk_a", rel) || se.DeleteRelationship(0, 0, "fk_a") ||
|
||||
se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationshipNames(0, 0) != nil {
|
||||
t.Error("nil map handling")
|
||||
}
|
||||
|
||||
if se.CreateRelationship(0, 0, rel) != rel {
|
||||
t.Fatal("create")
|
||||
}
|
||||
se.CreateRelationship(0, 0, &models.Relationship{Name: "fk_0"})
|
||||
if got := se.GetRelationshipNames(0, 0); !reflect.DeepEqual(got, []string{"fk_0", "fk_a"}) {
|
||||
t.Errorf("names must be sorted: %v", got)
|
||||
}
|
||||
if se.GetRelationship(0, 0, "fk_a") != rel || se.GetRelationship(0, 0, "none") != nil {
|
||||
t.Error("get")
|
||||
}
|
||||
|
||||
renamed := &models.Relationship{Name: "fk_b"}
|
||||
if !se.UpdateRelationship(0, 0, "fk_a", renamed) {
|
||||
t.Fatal("update")
|
||||
}
|
||||
if se.GetRelationship(0, 0, "fk_a") != nil || se.GetRelationship(0, 0, "fk_b") != renamed {
|
||||
t.Error("rename")
|
||||
}
|
||||
if se.UpdateRelationship(9, 0, "x", renamed) || se.UpdateRelationship(0, 9, "x", renamed) {
|
||||
t.Error("update bad index")
|
||||
}
|
||||
if se.DeleteRelationship(9, 0, "x") || se.DeleteRelationship(0, 9, "x") {
|
||||
t.Error("delete bad index")
|
||||
}
|
||||
if !se.DeleteRelationship(0, 0, "fk_b") || se.GetRelationship(0, 0, "fk_b") != nil {
|
||||
t.Error("delete")
|
||||
}
|
||||
if se.GetRelationship(9, 0, "x") != nil || se.GetRelationship(0, 9, "x") != nil ||
|
||||
se.GetRelationshipNames(9, 0) != nil || se.GetRelationshipNames(0, 9) != nil {
|
||||
t.Error("bad index reads")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaDataOps(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
s := se.CreateSchema("sales", "desc")
|
||||
if s == nil || s.Name != "sales" || s.Description != "desc" || s.Tables == nil || s.Sequences == nil || s.Enums == nil {
|
||||
t.Fatalf("create: %+v", s)
|
||||
}
|
||||
if len(se.GetAllSchemas()) != 2 || se.GetSchema(1) != s || se.GetSchema(2) != nil || se.GetSchema(-1) != nil {
|
||||
t.Error("get")
|
||||
}
|
||||
se.UpdateSchema(1, "billing", "owner", "d2")
|
||||
if s.Name != "billing" || s.Owner != "owner" || s.Description != "d2" {
|
||||
t.Errorf("update: %+v", s)
|
||||
}
|
||||
se.UpdateSchema(9, "x", "x", "x") // no panic
|
||||
if se.DeleteSchema(9) || se.DeleteSchema(-1) {
|
||||
t.Error("delete bad index")
|
||||
}
|
||||
if !se.DeleteSchema(1) || len(se.db.Schemas) != 1 {
|
||||
t.Error("delete")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTableDataOps(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
if se.CreateTable(9, "x", "") != nil {
|
||||
t.Error("create bad schema")
|
||||
}
|
||||
tbl := se.CreateTable(0, "orders", "d")
|
||||
if tbl == nil || tbl.Schema != "public" || tbl.Columns == nil || tbl.Constraints == nil || tbl.Indexes == nil {
|
||||
t.Fatalf("create: %+v", tbl)
|
||||
}
|
||||
if se.GetTable(0, 1) != tbl || se.GetTable(0, 2) != nil || se.GetTable(9, 0) != nil || se.GetTable(0, -1) != nil {
|
||||
t.Error("get")
|
||||
}
|
||||
if len(se.GetAllTables()) != 2 || len(se.GetTablesInSchema(0)) != 2 || se.GetTablesInSchema(9) != nil {
|
||||
t.Error("get all")
|
||||
}
|
||||
se.UpdateTable(0, 1, "orders2", "d2")
|
||||
if tbl.Name != "orders2" || tbl.Description != "d2" {
|
||||
t.Errorf("update: %+v", tbl)
|
||||
}
|
||||
se.UpdateTable(9, 0, "x", "x")
|
||||
se.UpdateTable(0, 9, "x", "x")
|
||||
if se.DeleteTable(9, 0) || se.DeleteTable(0, 9) {
|
||||
t.Error("delete bad index")
|
||||
}
|
||||
if !se.DeleteTable(0, 1) || len(se.db.Schemas[0].Tables) != 1 {
|
||||
t.Error("delete")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateDatabase(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
se.updateDatabase("n", "d", "c", "pgsql", "16")
|
||||
db := se.db
|
||||
if db.Name != "n" || db.Description != "d" || db.Comment != "c" || db.DatabaseType != models.PostgresqlDatabaseType || db.DatabaseVersion != "16" {
|
||||
t.Errorf("%+v", db)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDomainDataOps(t *testing.T) {
|
||||
se := NewSchemaEditor(models.InitDatabase("d"))
|
||||
se.createDomain("a", "da")
|
||||
se.createDomain("b", "db")
|
||||
if len(se.db.Domains) != 2 || se.db.Domains[1].Sequence != 1 {
|
||||
t.Fatalf("create: %+v", se.db.Domains)
|
||||
}
|
||||
se.updateDomain(0, "a2", "da2")
|
||||
se.updateDomain(9, "x", "x")
|
||||
if se.db.Domains[0].Name != "a2" || se.db.Domains[0].Description != "da2" {
|
||||
t.Error("update")
|
||||
}
|
||||
se.deleteDomain(9)
|
||||
se.deleteDomain(-1)
|
||||
se.deleteDomain(0)
|
||||
if len(se.db.Domains) != 1 || se.db.Domains[0].Name != "b" {
|
||||
t.Errorf("delete: %+v", se.db.Domains)
|
||||
}
|
||||
}
|
||||
@@ -207,6 +207,10 @@ func (se *SchemaEditor) showDomainEditor(index int, domain *models.Domain) {
|
||||
se.showDomainList()
|
||||
})
|
||||
|
||||
form.AddButton("Tables", func() {
|
||||
se.showDomainTables(index)
|
||||
})
|
||||
|
||||
form.AddButton("Delete", func() {
|
||||
se.showDeleteDomainConfirm(index)
|
||||
})
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// FileEntry is a single row in the file browser.
|
||||
type FileEntry struct {
|
||||
Name string
|
||||
IsDir bool
|
||||
}
|
||||
|
||||
// formatExtensions maps a UI format name to the file extensions it reads or writes.
|
||||
var formatExtensions = map[string][]string{
|
||||
"dbml": {".dbml"},
|
||||
"dctx": {".dctx"},
|
||||
"drawdb": {".json"},
|
||||
"graphql": {".graphql", ".gql"},
|
||||
"json": {".json"},
|
||||
"yaml": {".yaml", ".yml"},
|
||||
"gorm": {".go"},
|
||||
"bun": {".go"},
|
||||
"drizzle": {".ts"},
|
||||
"prisma": {".prisma"},
|
||||
"typeorm": {".ts"},
|
||||
"pgsql": {".sql"},
|
||||
"sqlite": {".db", ".sqlite", ".sqlite3"},
|
||||
}
|
||||
|
||||
// directoryFormats are formats whose reader/writer accepts a directory.
|
||||
var directoryFormats = map[string]bool{
|
||||
"gorm": true, "bun": true, "drizzle": true, "typeorm": true,
|
||||
}
|
||||
|
||||
// FormatExtensions returns the extensions for a format, or nil (no filter) if unknown.
|
||||
func FormatExtensions(format string) []string {
|
||||
return formatExtensions[format]
|
||||
}
|
||||
|
||||
// IsDirectoryFormat reports whether a format can be loaded from or saved to a directory.
|
||||
func IsDirectoryFormat(format string) bool {
|
||||
return directoryFormats[format]
|
||||
}
|
||||
|
||||
// ExpandHome replaces a leading ~ with the user's home directory.
|
||||
func ExpandHome(p string) string {
|
||||
if strings.HasPrefix(p, "~") {
|
||||
if home, err := os.UserHomeDir(); err == nil {
|
||||
return filepath.Join(home, p[1:])
|
||||
}
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// MatchesExtension reports whether name has one of exts (case-insensitive).
|
||||
// An empty extension list matches everything.
|
||||
func MatchesExtension(name string, exts []string) bool {
|
||||
if len(exts) == 0 {
|
||||
return true
|
||||
}
|
||||
ext := strings.ToLower(filepath.Ext(name))
|
||||
for _, e := range exts {
|
||||
if strings.EqualFold(e, ext) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ListDir returns the entries of dir: directories first, then files that match
|
||||
// exts, each group sorted case-insensitively. Hidden (dot) entries are skipped
|
||||
// unless showHidden is set.
|
||||
func ListDir(dir string, exts []string, showHidden bool) ([]FileEntry, error) {
|
||||
items, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var dirs, files []FileEntry
|
||||
for _, item := range items {
|
||||
name := item.Name()
|
||||
if !showHidden && strings.HasPrefix(name, ".") {
|
||||
continue
|
||||
}
|
||||
isDir := item.IsDir()
|
||||
if !isDir && item.Type()&os.ModeSymlink != 0 {
|
||||
// Follow symlinks so links to directories are navigable.
|
||||
if info, err := os.Stat(filepath.Join(dir, name)); err == nil {
|
||||
isDir = info.IsDir()
|
||||
}
|
||||
}
|
||||
if isDir {
|
||||
dirs = append(dirs, FileEntry{Name: name, IsDir: true})
|
||||
} else if MatchesExtension(name, exts) {
|
||||
files = append(files, FileEntry{Name: name})
|
||||
}
|
||||
}
|
||||
|
||||
byName := func(s []FileEntry) {
|
||||
sort.Slice(s, func(i, j int) bool {
|
||||
return strings.ToLower(s[i].Name) < strings.ToLower(s[j].Name)
|
||||
})
|
||||
}
|
||||
byName(dirs)
|
||||
byName(files)
|
||||
return append(dirs, files...), nil
|
||||
}
|
||||
|
||||
// ResolveStart works out where the browser should open for the current input
|
||||
// value. It returns the directory to show and, if the input named a file, its
|
||||
// base name. Falls back to the working directory.
|
||||
func ResolveStart(input string) (dir, name string) {
|
||||
input = strings.TrimSpace(input)
|
||||
if input != "" {
|
||||
p := ExpandHome(input)
|
||||
if abs, err := filepath.Abs(p); err == nil {
|
||||
p = abs
|
||||
}
|
||||
if info, err := os.Stat(p); err == nil && info.IsDir() {
|
||||
return p, ""
|
||||
}
|
||||
if info, err := os.Stat(filepath.Dir(p)); err == nil && info.IsDir() {
|
||||
return filepath.Dir(p), filepath.Base(p)
|
||||
}
|
||||
}
|
||||
wd, err := os.Getwd()
|
||||
if err != nil {
|
||||
wd = "."
|
||||
}
|
||||
return wd, ""
|
||||
}
|
||||
@@ -0,0 +1,363 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/gdamore/tcell/v2"
|
||||
"github.com/rivo/tview"
|
||||
)
|
||||
|
||||
// FileBrowserMode selects between picking an existing path and choosing a save target.
|
||||
type FileBrowserMode int
|
||||
|
||||
const (
|
||||
FileBrowserLoad FileBrowserMode = iota
|
||||
FileBrowserSave
|
||||
)
|
||||
|
||||
// FileBrowserConfig configures the file browser dialog.
|
||||
type FileBrowserConfig struct {
|
||||
Mode FileBrowserMode
|
||||
StartPath string // current value of the input; may be empty
|
||||
Extensions []string // empty = show all files
|
||||
AllowDir bool // a directory is a valid result (directory-based formats)
|
||||
ReturnPage string // page to switch back to when the dialog closes
|
||||
OnSelect func(path string)
|
||||
}
|
||||
|
||||
// attachFileBrowser makes Enter on the named input open the file browser,
|
||||
// filtered for the currently selected format.
|
||||
func (se *SchemaEditor) attachFileBrowser(form *tview.Form, label, returnPage string, mode FileBrowserMode, format func() string) {
|
||||
item, ok := form.GetFormItemByLabel(label).(*tview.InputField)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() != tcell.KeyEnter {
|
||||
return event
|
||||
}
|
||||
f := format()
|
||||
se.showFileBrowser(FileBrowserConfig{
|
||||
Mode: mode,
|
||||
StartPath: item.GetText(),
|
||||
Extensions: FormatExtensions(f),
|
||||
AllowDir: IsDirectoryFormat(f),
|
||||
ReturnPage: returnPage,
|
||||
OnSelect: func(path string) { item.SetText(path) },
|
||||
})
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// showFileBrowser displays the file browser page. Esc closes it without
|
||||
// calling OnSelect, leaving the originating input unchanged.
|
||||
func (se *SchemaEditor) showFileBrowser(cfg FileBrowserConfig) {
|
||||
const pageName = "file-browser"
|
||||
|
||||
dir, startName := ResolveStart(cfg.StartPath)
|
||||
showHidden := false
|
||||
useFilter := len(cfg.Extensions) > 0
|
||||
var entries []FileEntry // rows shown below the ".." row
|
||||
|
||||
title := tview.NewTextView().
|
||||
SetText("[::b]Select File").
|
||||
SetTextAlign(tview.AlignCenter).
|
||||
SetDynamicColors(true)
|
||||
if cfg.Mode == FileBrowserSave {
|
||||
title.SetText("[::b]Save As")
|
||||
}
|
||||
|
||||
info := tview.NewTextView().SetDynamicColors(true)
|
||||
|
||||
fileTable := tview.NewTable().SetSelectable(true, false).SetFixed(0, 0)
|
||||
fileTable.SetBorder(true)
|
||||
|
||||
nameInput := tview.NewInputField().SetLabel("File name: ").SetFieldWidth(0)
|
||||
nameInput.SetText(startName)
|
||||
|
||||
closeBrowser := func() {
|
||||
se.pages.RemovePage(pageName)
|
||||
se.pages.SwitchToPage(cfg.ReturnPage)
|
||||
}
|
||||
|
||||
finish := func(path string) {
|
||||
closeBrowser()
|
||||
cfg.OnSelect(path)
|
||||
}
|
||||
|
||||
refresh := func() {
|
||||
exts := cfg.Extensions
|
||||
if !useFilter {
|
||||
exts = nil
|
||||
}
|
||||
list, err := ListDir(dir, exts, showHidden)
|
||||
if err != nil {
|
||||
se.showErrorDialog("Error", fmt.Sprintf("Cannot read %s: %v", dir, err))
|
||||
list = nil
|
||||
}
|
||||
entries = list
|
||||
|
||||
fileTable.Clear()
|
||||
fileTable.SetCell(0, 0, tview.NewTableCell("[..]").SetTextColor(tcell.ColorAqua))
|
||||
for i, e := range entries {
|
||||
cell := tview.NewTableCell(e.Name)
|
||||
if e.IsDir {
|
||||
cell.SetText(e.Name + "/").SetTextColor(tcell.ColorAqua)
|
||||
}
|
||||
fileTable.SetCell(i+1, 0, cell)
|
||||
}
|
||||
fileTable.Select(0, 0)
|
||||
if len(entries) > 0 {
|
||||
fileTable.Select(1, 0)
|
||||
}
|
||||
|
||||
filterText := "all files"
|
||||
if useFilter {
|
||||
filterText = fmt.Sprintf("%v", cfg.Extensions)
|
||||
}
|
||||
hiddenText := "hidden: off"
|
||||
if showHidden {
|
||||
hiddenText = "hidden: on"
|
||||
}
|
||||
info.SetText(fmt.Sprintf("%s [yellow](%s, filter: %s)[-]", tview.Escape(dir), hiddenText, tview.Escape(filterText)))
|
||||
fileTable.SetTitle(" Files ")
|
||||
}
|
||||
|
||||
goUp := func() {
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
return
|
||||
}
|
||||
prev := filepath.Base(dir)
|
||||
dir = parent
|
||||
refresh()
|
||||
for i, e := range entries {
|
||||
if e.Name == prev {
|
||||
fileTable.Select(i+1, 0)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
selected := func() (FileEntry, bool) {
|
||||
row, _ := fileTable.GetSelection()
|
||||
if row < 1 || row > len(entries) {
|
||||
return FileEntry{}, false
|
||||
}
|
||||
return entries[row-1], true
|
||||
}
|
||||
|
||||
// confirmOverwrite asks before replacing an existing file (not directories).
|
||||
confirmOverwrite := func(path string) {
|
||||
modal := tview.NewModal().
|
||||
SetText(fmt.Sprintf("File already exists:\n%s\n\nOverwrite it?", path)).
|
||||
AddButtons([]string{"Cancel", "Overwrite"}).
|
||||
SetDoneFunc(func(_ int, label string) {
|
||||
se.pages.RemovePage("overwrite-confirm")
|
||||
se.pages.SwitchToPage(pageName)
|
||||
if label == "Overwrite" {
|
||||
finish(path)
|
||||
}
|
||||
})
|
||||
modal.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() == tcell.KeyEscape {
|
||||
se.pages.RemovePage("overwrite-confirm")
|
||||
se.pages.SwitchToPage(pageName)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
se.pages.AddAndSwitchToPage("overwrite-confirm", modal, true)
|
||||
}
|
||||
|
||||
chooseSave := func() {
|
||||
name := nameInput.GetText()
|
||||
if name == "" {
|
||||
if cfg.AllowDir {
|
||||
finish(dir)
|
||||
return
|
||||
}
|
||||
se.showErrorDialog("Error", "Enter a file name")
|
||||
return
|
||||
}
|
||||
path := filepath.Join(dir, name)
|
||||
if st, err := os.Stat(path); err == nil {
|
||||
if st.IsDir() {
|
||||
se.showErrorDialog("Error", name+" is a directory")
|
||||
return
|
||||
}
|
||||
confirmOverwrite(path)
|
||||
return
|
||||
}
|
||||
finish(path)
|
||||
}
|
||||
|
||||
// chooseHighlighted handles Select: the highlighted entry in load mode, or
|
||||
// the typed name in save mode.
|
||||
chooseHighlighted := func() {
|
||||
if cfg.Mode == FileBrowserSave {
|
||||
chooseSave()
|
||||
return
|
||||
}
|
||||
e, ok := selected()
|
||||
switch {
|
||||
case ok && !e.IsDir:
|
||||
finish(filepath.Join(dir, e.Name))
|
||||
case ok && cfg.AllowDir:
|
||||
finish(filepath.Join(dir, e.Name))
|
||||
case cfg.AllowDir:
|
||||
finish(dir)
|
||||
default:
|
||||
se.showErrorDialog("Error", "Select a file")
|
||||
}
|
||||
}
|
||||
|
||||
activate := func() {
|
||||
row, _ := fileTable.GetSelection()
|
||||
if row == 0 {
|
||||
goUp()
|
||||
return
|
||||
}
|
||||
e, ok := selected()
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if e.IsDir {
|
||||
dir = filepath.Join(dir, e.Name)
|
||||
refresh()
|
||||
return
|
||||
}
|
||||
if cfg.Mode == FileBrowserSave {
|
||||
nameInput.SetText(e.Name)
|
||||
return
|
||||
}
|
||||
finish(filepath.Join(dir, e.Name))
|
||||
}
|
||||
|
||||
toggleHidden := func() { showHidden = !showHidden; refresh() }
|
||||
toggleFilter := func() {
|
||||
if len(cfg.Extensions) > 0 {
|
||||
useFilter = !useFilter
|
||||
refresh()
|
||||
}
|
||||
}
|
||||
|
||||
btnSelect := tview.NewButton("Select [s]").SetSelectedFunc(chooseHighlighted)
|
||||
btnHidden := tview.NewButton("Hidden [h]").SetSelectedFunc(toggleHidden)
|
||||
btnFilter := tview.NewButton("Filter [f]").SetSelectedFunc(toggleFilter)
|
||||
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(closeBrowser)
|
||||
|
||||
btnFlex := tview.NewFlex().
|
||||
AddItem(btnSelect, 0, 1, false).
|
||||
AddItem(btnHidden, 0, 1, false).
|
||||
AddItem(btnFilter, 0, 1, false).
|
||||
AddItem(btnBack, 0, 1, false)
|
||||
|
||||
flex := tview.NewFlex().SetDirection(tview.FlexRow).
|
||||
AddItem(title, 1, 0, false).
|
||||
AddItem(info, 1, 0, false).
|
||||
AddItem(fileTable, 0, 1, true)
|
||||
|
||||
focusOrder := []tview.Primitive{fileTable}
|
||||
if cfg.Mode == FileBrowserSave {
|
||||
flex.AddItem(nameInput, 1, 0, false)
|
||||
focusOrder = append(focusOrder, nameInput)
|
||||
}
|
||||
flex.AddItem(btnFlex, 1, 0, false)
|
||||
focusOrder = append(focusOrder, btnSelect, btnHidden, btnFilter, btnBack)
|
||||
|
||||
// Circular Tab / Shift+Tab across every focusable widget.
|
||||
cycle := func(event *tcell.EventKey) *tcell.EventKey {
|
||||
step := 0
|
||||
switch event.Key() {
|
||||
case tcell.KeyTab:
|
||||
step = 1
|
||||
case tcell.KeyBacktab:
|
||||
step = -1
|
||||
default:
|
||||
return event
|
||||
}
|
||||
for i, p := range focusOrder {
|
||||
if p.HasFocus() {
|
||||
se.app.SetFocus(focusOrder[(i+step+len(focusOrder))%len(focusOrder)])
|
||||
break
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
fileTable.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event = cycle(event); event == nil {
|
||||
return nil
|
||||
}
|
||||
switch event.Key() {
|
||||
case tcell.KeyEscape:
|
||||
closeBrowser()
|
||||
return nil
|
||||
case tcell.KeyEnter:
|
||||
activate()
|
||||
return nil
|
||||
case tcell.KeyBackspace, tcell.KeyBackspace2, tcell.KeyLeft:
|
||||
goUp()
|
||||
return nil
|
||||
}
|
||||
switch event.Rune() {
|
||||
case 's':
|
||||
chooseHighlighted()
|
||||
return nil
|
||||
case 'h':
|
||||
toggleHidden()
|
||||
return nil
|
||||
case 'f':
|
||||
toggleFilter()
|
||||
return nil
|
||||
case 'b':
|
||||
closeBrowser()
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
|
||||
nameInput.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event = cycle(event); event == nil {
|
||||
return nil
|
||||
}
|
||||
switch event.Key() {
|
||||
case tcell.KeyEscape:
|
||||
closeBrowser()
|
||||
return nil
|
||||
case tcell.KeyEnter:
|
||||
chooseSave()
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
|
||||
for _, b := range []*tview.Button{btnSelect, btnHidden, btnFilter, btnBack} {
|
||||
b.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event = cycle(event); event == nil {
|
||||
return nil
|
||||
}
|
||||
if event.Key() == tcell.KeyEscape {
|
||||
closeBrowser()
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
}
|
||||
|
||||
refresh()
|
||||
if startName != "" {
|
||||
for i, e := range entries {
|
||||
if e.Name == startName {
|
||||
fileTable.Select(i+1, 0)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
se.pages.AddAndSwitchToPage(pageName, flex, true)
|
||||
se.app.SetFocus(fileTable)
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func touch(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(path, nil, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func names(entries []FileEntry) []string {
|
||||
var out []string
|
||||
for _, e := range entries {
|
||||
if e.IsDir {
|
||||
out = append(out, e.Name+"/")
|
||||
} else {
|
||||
out = append(out, e.Name)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestMatchesExtension(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
exts []string
|
||||
want bool
|
||||
}{
|
||||
{"a.dbml", []string{".dbml"}, true},
|
||||
{"A.DBML", []string{".dbml"}, true},
|
||||
{"a.json", []string{".dbml"}, false},
|
||||
{"a.yml", []string{".yaml", ".yml"}, true},
|
||||
{"noext", []string{".sql"}, false},
|
||||
{"anything", nil, true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := MatchesExtension(tt.name, tt.exts); got != tt.want {
|
||||
t.Errorf("MatchesExtension(%q, %v) = %v, want %v", tt.name, tt.exts, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestListDirFilterAndHidden(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
touch(t, filepath.Join(dir, "b.dbml"))
|
||||
touch(t, filepath.Join(dir, "A.dbml"))
|
||||
touch(t, filepath.Join(dir, "c.json"))
|
||||
touch(t, filepath.Join(dir, ".hidden.dbml"))
|
||||
touch(t, filepath.Join(dir, "sub", "x.txt"))
|
||||
touch(t, filepath.Join(dir, ".git", "x"))
|
||||
|
||||
got, err := ListDir(dir, FormatExtensions("dbml"), false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if want := []string{"sub/", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||
t.Errorf("filtered: got %v, want %v", names(got), want)
|
||||
}
|
||||
|
||||
got, _ = ListDir(dir, FormatExtensions("dbml"), true)
|
||||
if want := []string{".git/", "sub/", ".hidden.dbml", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||
t.Errorf("hidden: got %v, want %v", names(got), want)
|
||||
}
|
||||
|
||||
got, _ = ListDir(dir, nil, false)
|
||||
if want := []string{"sub/", "A.dbml", "b.dbml", "c.json"}; !reflect.DeepEqual(names(got), want) {
|
||||
t.Errorf("no filter: got %v, want %v", names(got), want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListDirMissing(t *testing.T) {
|
||||
if _, err := ListDir(filepath.Join(t.TempDir(), "nope"), nil, false); err == nil {
|
||||
t.Error("expected error for missing directory")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStart(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
file := filepath.Join(dir, "schema.dbml")
|
||||
touch(t, file)
|
||||
wd, _ := os.Getwd()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
wantDir string
|
||||
wantFileName string
|
||||
}{
|
||||
{"existing file", file, dir, "schema.dbml"},
|
||||
{"directory", dir, dir, ""},
|
||||
{"new file in existing dir", filepath.Join(dir, "new.dbml"), dir, "new.dbml"},
|
||||
{"empty", "", wd, ""},
|
||||
{"nonexistent parent", filepath.Join(dir, "no", "such", "f.dbml"), wd, ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
d, n := ResolveStart(tt.in)
|
||||
if d != tt.wantDir || n != tt.wantFileName {
|
||||
t.Errorf("got (%q, %q), want (%q, %q)", d, n, tt.wantDir, tt.wantFileName)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatExtensions(t *testing.T) {
|
||||
if got := FormatExtensions("yaml"); !reflect.DeepEqual(got, []string{".yaml", ".yml"}) {
|
||||
t.Errorf("yaml: %v", got)
|
||||
}
|
||||
if FormatExtensions("unknown") != nil {
|
||||
t.Error("unknown format should not filter")
|
||||
}
|
||||
if !IsDirectoryFormat("gorm") || IsDirectoryFormat("json") {
|
||||
t.Error("directory format detection wrong")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,308 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/rivo/tview"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
const uiFixtures = "../../tests/assets"
|
||||
|
||||
func newUIEditor() *SchemaEditor {
|
||||
se := NewSchemaEditor(models.InitDatabase("start"))
|
||||
se.db = newTestEditor().db
|
||||
return se
|
||||
}
|
||||
|
||||
func hasPage(se *SchemaEditor, name string) bool {
|
||||
return se.pages.HasPage(name)
|
||||
}
|
||||
|
||||
func TestSortedKeysAndColumnNames(t *testing.T) {
|
||||
if got := sortedKeys(map[string]int{"b": 1, "a": 2, "c": 3}); strings.Join(got, ",") != "a,b,c" {
|
||||
t.Errorf("sortedKeys: %v", got)
|
||||
}
|
||||
if got := sortedKeys[int](nil); len(got) != 0 {
|
||||
t.Errorf("nil map: %v", got)
|
||||
}
|
||||
tbl := models.InitTable("t", "s")
|
||||
tbl.Columns["z"] = models.InitColumn("z", "t", "s")
|
||||
tbl.Columns["a"] = models.InitColumn("a", "t", "s")
|
||||
if got := getColumnNames(tbl); strings.Join(got, ",") != "a,z" {
|
||||
t.Errorf("getColumnNames: %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocations(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
se.db.Schemas = append(se.db.Schemas, models.InitSchema("empty"))
|
||||
sl := se.schemaLocations()
|
||||
if len(sl) != 2 || sl[0].label != "public" || sl[1].schemaIndex != 1 || sl[0].tableIndex != -1 {
|
||||
t.Errorf("schemaLocations: %+v", sl)
|
||||
}
|
||||
tl := se.tableLocations()
|
||||
if len(tl) != 1 || tl[0].label != "public.users" || tl[0].schemaIndex != 0 || tl[0].tableIndex != 0 {
|
||||
t.Errorf("tableLocations: %+v", tl)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseSkipTablesUI(t *testing.T) {
|
||||
if got := parseSkipTablesUI(""); len(got) != 0 {
|
||||
t.Errorf("empty: %v", got)
|
||||
}
|
||||
got := parseSkipTablesUI(" Users , ORDERS ,, ")
|
||||
if len(got) != 2 || !got["users"] || !got["orders"] {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHelpTexts(t *testing.T) {
|
||||
for name, fn := range map[string]func() string{"load": getLoadHelpText, "save": getSaveHelpText, "import": getImportHelpText} {
|
||||
if txt := fn(); !strings.Contains(txt, "dbml") && name != "save" || txt == "" {
|
||||
t.Errorf("%s help text: %q", name, txt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestObjectKinds(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
if err := se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}, Unique: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveView(0, -1, &models.View{Name: "v1", Definition: "select 1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s1", IncrementBy: 1, StartValue: 1}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveScript(0, -1, &models.Script{Name: "sc1", SQL: "select 1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
kinds := map[string]objectKind{
|
||||
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
|
||||
}
|
||||
for page, k := range kinds {
|
||||
t.Run(page, func(t *testing.T) {
|
||||
if k.page != page || k.title == "" || k.singular == "" || len(k.headers) == 0 {
|
||||
t.Fatalf("metadata: %+v", k)
|
||||
}
|
||||
rows := k.rows()
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("rows: %+v", rows)
|
||||
}
|
||||
for _, r := range rows {
|
||||
if len(r.cells) != len(k.headers) {
|
||||
t.Errorf("cells %v do not match headers %v", r.cells, k.headers)
|
||||
}
|
||||
}
|
||||
if len(k.locations()) == 0 {
|
||||
t.Error("no locations")
|
||||
}
|
||||
|
||||
// Editing an existing row without changes keeps it valid.
|
||||
form := tview.NewForm()
|
||||
save := k.buildForm(form, &rows[0])
|
||||
if form.GetFormItemCount() == 0 {
|
||||
t.Error("no form fields")
|
||||
}
|
||||
loc := k.locations()[0]
|
||||
loc.schemaIndex, loc.tableIndex = rows[0].schemaIndex, rows[0].tableIndex
|
||||
if err := save(loc); err != nil {
|
||||
t.Errorf("save unchanged: %v", err)
|
||||
}
|
||||
|
||||
// A blank new form is rejected by validation.
|
||||
blank := tview.NewForm()
|
||||
saveBlank := k.buildForm(blank, nil)
|
||||
if err := saveBlank(k.locations()[0]); err == nil {
|
||||
t.Error("blank form accepted")
|
||||
}
|
||||
|
||||
if !k.remove(rows[0]) || len(k.rows()) != 0 {
|
||||
t.Error("remove failed")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestObjectKind_CreateIndexFromForm(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
k := se.indexKind()
|
||||
form := tview.NewForm()
|
||||
save := k.buildForm(form, nil)
|
||||
form.GetFormItemByLabel("Name").(*tview.InputField).SetText("idx_new")
|
||||
form.GetFormItemByLabel("Columns (comma separated)").(*tview.InputField).SetText("id, email")
|
||||
if err := save(k.locations()[0]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
idx := se.db.Schemas[0].Tables[0].Indexes["idx_new"]
|
||||
if idx == nil || len(idx.Columns) != 2 || idx.Type != "btree" {
|
||||
t.Errorf("index: %+v", idx)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDatabase(t *testing.T) {
|
||||
for _, tt := range []struct{ format, path string }{
|
||||
{"dbml", "dbml/simple.dbml"},
|
||||
{"json", "json/database.json"},
|
||||
{"yaml", "yaml/database.yaml"},
|
||||
{"drawdb", "drawdb/simple.json"},
|
||||
{"dctx", "dctx/p1.dctx"},
|
||||
{"graphql", "graphql/simple.graphql"},
|
||||
{"prisma", "prisma/example.prisma"},
|
||||
{"typeorm", "typeorm/example.ts"},
|
||||
{"drizzle", "drizzle/schema.ts"},
|
||||
{"gorm", "gorm/simple.go"},
|
||||
{"bun", "bun/simple.go"},
|
||||
} {
|
||||
t.Run(tt.format, func(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
se.loadDatabase(tt.format, filepath.Join(uiFixtures, tt.path), "")
|
||||
if hasPage(se, "error-dialog") || !hasPage(se, "success-dialog") {
|
||||
t.Fatalf("expected success dialog (pages: error=%v)", hasPage(se, "error-dialog"))
|
||||
}
|
||||
if se.loadConfig == nil || se.loadConfig.SourceType != tt.format || len(se.db.Schemas) == 0 {
|
||||
t.Errorf("state: %+v db=%+v", se.loadConfig, se.db)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
errCases := []struct {
|
||||
name, format, path, conn string
|
||||
}{
|
||||
{"pgsql no conn", "pgsql", "", ""},
|
||||
{"file required", "json", "", ""},
|
||||
{"unsupported", "nope", "x", ""},
|
||||
{"missing file", "json", filepath.Join(t.TempDir(), "missing.json"), ""},
|
||||
}
|
||||
for _, tt := range errCases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
before := se.db
|
||||
se.loadDatabase(tt.format, tt.path, tt.conn)
|
||||
if !hasPage(se, "error-dialog") {
|
||||
t.Error("expected error dialog")
|
||||
}
|
||||
if se.db != before || se.loadConfig != nil {
|
||||
t.Error("state must be unchanged on error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateNewDatabase(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
se.loadConfig = &LoadConfig{SourceType: "json"}
|
||||
se.createNewDatabase()
|
||||
if se.db.Name != "New Database" || len(se.db.Schemas) != 0 || se.loadConfig != nil || !hasPage(se, "success-dialog") {
|
||||
t.Errorf("state: %+v", se.db)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveDatabase(t *testing.T) {
|
||||
for _, tt := range []struct{ format, file string }{
|
||||
{"json", "o.json"},
|
||||
{"yaml", "o.yaml"},
|
||||
{"dbml", "o.dbml"},
|
||||
{"drawdb", "o.drawdb.json"},
|
||||
{"graphql", "o.graphql"},
|
||||
{"prisma", "o.prisma"},
|
||||
{"typeorm", "o.ts"},
|
||||
{"drizzle", "d.ts"},
|
||||
{"gorm", "g.go"},
|
||||
{"bun", "b.go"},
|
||||
} {
|
||||
t.Run(tt.format, func(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
out := filepath.Join(t.TempDir(), tt.file)
|
||||
se.saveDatabase(tt.format, out)
|
||||
if hasPage(se, "error-dialog") {
|
||||
t.Fatal("unexpected error dialog")
|
||||
}
|
||||
if se.saveConfig == nil || se.saveConfig.FilePath != out || se.saveConfig.TargetType != tt.format {
|
||||
t.Errorf("saveConfig: %+v", se.saveConfig)
|
||||
}
|
||||
if info, err := os.Stat(out); err != nil || info.Size() == 0 {
|
||||
t.Errorf("output: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for name, args := range map[string][2]string{
|
||||
"pgsql unsupported": {"pgsql", "x.sql"},
|
||||
"path required": {"json", ""},
|
||||
"unknown format": {"nope", "x"},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
se.saveDatabase(args[0], args[1])
|
||||
if !hasPage(se, "error-dialog") || se.saveConfig != nil {
|
||||
t.Error("expected error dialog and no saveConfig")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportAndMerge(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
se.importAndMergeDatabase("json", filepath.Join(uiFixtures, "json/database.json"), "", false, false, false, false, false, "")
|
||||
if hasPage(se, "error-dialog") {
|
||||
t.Fatal("unexpected error dialog")
|
||||
}
|
||||
|
||||
for name, args := range map[string][3]string{
|
||||
"pgsql no conn": {"pgsql", "", ""},
|
||||
"file required": {"json", "", ""},
|
||||
"unsupported": {"nope", "x", ""},
|
||||
"missing file": {"json", filepath.Join(t.TempDir(), "missing.json"), ""},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
se.importAndMergeDatabase(args[0], args[1], args[2], false, false, false, false, false, "")
|
||||
if !hasPage(se, "error-dialog") {
|
||||
t.Error("expected error dialog")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerformMerge(t *testing.T) {
|
||||
se := newUIEditor()
|
||||
src := models.InitDatabase("src")
|
||||
s := models.InitSchema("public")
|
||||
tbl := models.InitTable("orders", "public")
|
||||
tbl.Columns["id"] = models.InitColumn("id", "orders", "public")
|
||||
skip := models.InitTable("skipme", "public")
|
||||
s.Tables = append(s.Tables, tbl, skip)
|
||||
src.Schemas = append(src.Schemas, s)
|
||||
|
||||
se.performMerge(src, false, false, false, false, false, "SkipMe")
|
||||
if !hasPage(se, "success-dialog") {
|
||||
t.Error("expected success dialog")
|
||||
}
|
||||
names := map[string]bool{}
|
||||
for _, tb := range se.db.Schemas[0].Tables {
|
||||
names[tb.Name] = true
|
||||
}
|
||||
if !names["users"] || !names["orders"] || names["skipme"] || len(names) != 2 {
|
||||
t.Errorf("tables after merge: %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEditorAccessors(t *testing.T) {
|
||||
db := models.InitDatabase("d")
|
||||
lc, sc := &LoadConfig{SourceType: "json"}, &SaveConfig{TargetType: "yaml"}
|
||||
se := NewSchemaEditorWithConfigs(db, lc, sc)
|
||||
if se.GetDatabase() != db || se.loadConfig != lc || se.saveConfig != sc || se.app == nil || se.pages == nil {
|
||||
t.Errorf("%+v", se)
|
||||
}
|
||||
if se.createMainMenu() == nil {
|
||||
t.Error("main menu")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/rivo/tview"
|
||||
)
|
||||
|
||||
func newDialogTestEditor() *SchemaEditor {
|
||||
se := &SchemaEditor{app: tview.NewApplication(), pages: tview.NewPages()}
|
||||
se.pages.AddPage("origin", tview.NewBox(), true, true)
|
||||
return se
|
||||
}
|
||||
|
||||
func TestFileBrowserOpensOnEachMode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
touch(t, filepath.Join(dir, "a.dbml"))
|
||||
|
||||
for _, mode := range []FileBrowserMode{FileBrowserLoad, FileBrowserSave} {
|
||||
se := newDialogTestEditor()
|
||||
se.showFileBrowser(FileBrowserConfig{
|
||||
Mode: mode,
|
||||
StartPath: filepath.Join(dir, "a.dbml"),
|
||||
Extensions: FormatExtensions("dbml"),
|
||||
ReturnPage: "origin",
|
||||
OnSelect: func(string) { t.Error("OnSelect must not fire without a selection") },
|
||||
})
|
||||
if !se.pages.HasPage("file-browser") {
|
||||
t.Errorf("mode %d: file-browser page missing", mode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnStringBuilderOpensForEachKind(t *testing.T) {
|
||||
for _, in := range []string{
|
||||
"",
|
||||
"postgres://u:pw@db:5432/app?sslmode=disable",
|
||||
"sqlserver://sa:pw@sql:1433?database=shop&encrypt=disable",
|
||||
"/tmp/app.db",
|
||||
"postgres://u:p@host:badport/db", // parse error falls back to defaults
|
||||
} {
|
||||
se := newDialogTestEditor()
|
||||
se.showConnStringBuilder(in, ConnPostgres, "origin", func(string) {
|
||||
t.Error("onDone must not fire without Save")
|
||||
})
|
||||
if !se.pages.HasPage(connBuilderPage) {
|
||||
t.Errorf("%q: builder page missing", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateConnFields(t *testing.T) {
|
||||
if validateConnFields(ConnFields{Kind: ConnSQLite}) == "" {
|
||||
t.Error("sqlite without path should be invalid")
|
||||
}
|
||||
if validateConnFields(ConnFields{Kind: ConnPostgres}) == "" {
|
||||
t.Error("postgres without host should be invalid")
|
||||
}
|
||||
if msg := validateConnFields(DefaultConnFields(ConnMSSQL)); msg != "" {
|
||||
t.Errorf("defaults should be valid, got %q", msg)
|
||||
}
|
||||
}
|
||||
@@ -92,6 +92,9 @@ func (se *SchemaEditor) showLoadScreen() {
|
||||
connString = value
|
||||
})
|
||||
|
||||
se.attachFileBrowser(form, "File Path", "load-database", FileBrowserLoad, func() string { return currentFormat })
|
||||
se.attachConnStringBuilder(form, "Connection String", "load-database", func() string { return currentFormat })
|
||||
|
||||
form.AddTextView("Help", getLoadHelpText(), 0, 5, true, false)
|
||||
|
||||
// Buttons
|
||||
@@ -190,6 +193,8 @@ func (se *SchemaEditor) showSaveScreen() {
|
||||
filePath = value
|
||||
})
|
||||
|
||||
se.attachFileBrowser(form, "File Path", "save-database", FileBrowserSave, func() string { return currentFormat })
|
||||
|
||||
form.AddTextView("Help", getSaveHelpText(), 0, 5, true, false)
|
||||
|
||||
// Buttons
|
||||
@@ -469,6 +474,8 @@ func getLoadHelpText() string {
|
||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm
|
||||
Database formats: pgsql (requires connection string)
|
||||
|
||||
Press Enter in File Path to browse files, or in Connection String to open the builder.
|
||||
|
||||
Examples:
|
||||
- File path: ~/schemas/mydb.dbml or /path/to/schema.json
|
||||
- Connection: postgres://user:pass@localhost/dbname`
|
||||
@@ -520,6 +527,8 @@ func (se *SchemaEditor) showUpdateExistingDatabaseConfirm() {
|
||||
func getSaveHelpText() string {
|
||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql (SQL export)
|
||||
|
||||
Press Enter in File Path to browse for a target.
|
||||
|
||||
Examples:
|
||||
- File: ~/schemas/mydb.dbml
|
||||
- Directory (for code formats): ./models/`
|
||||
@@ -570,6 +579,9 @@ func (se *SchemaEditor) showImportScreen() {
|
||||
connString = value
|
||||
})
|
||||
|
||||
se.attachFileBrowser(form, "File Path", "import-database", FileBrowserLoad, func() string { return currentFormat })
|
||||
se.attachConnStringBuilder(form, "Connection String", "import-database", func() string { return currentFormat })
|
||||
|
||||
form.AddInputField("Skip Tables (comma-separated)", "", 50, nil, func(value string) {
|
||||
skipTables = value
|
||||
})
|
||||
|
||||
@@ -39,6 +39,18 @@ func (se *SchemaEditor) createMainMenu() tview.Primitive {
|
||||
AddItem("Manage Domains", "View, create, edit, and delete domains", 'd', func() {
|
||||
se.showDomainList()
|
||||
}).
|
||||
AddItem("Manage Indexes", "View, create, edit, and delete table indexes", 'x', func() {
|
||||
se.showObjectList(se.indexKind())
|
||||
}).
|
||||
AddItem("Manage Views", "View, create, edit, and delete views", 'v', func() {
|
||||
se.showObjectList(se.viewKind())
|
||||
}).
|
||||
AddItem("Manage Sequences", "View, create, edit, and delete sequences", 'u', func() {
|
||||
se.showObjectList(se.sequenceKind())
|
||||
}).
|
||||
AddItem("Manage Scripts", "View, create, edit, and delete SQL scripts", 'c', func() {
|
||||
se.showObjectList(se.scriptKind())
|
||||
}).
|
||||
AddItem("Import & Merge", "Import and merge schema from another database", 'i', func() {
|
||||
se.showImportScreen()
|
||||
}).
|
||||
|
||||
@@ -0,0 +1,263 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// Data operations for indexes, views, sequences, scripts and domain/table assignment.
|
||||
|
||||
func (se *SchemaEditor) schemaAt(schemaIndex int) (*models.Schema, error) {
|
||||
if schemaIndex < 0 || schemaIndex >= len(se.db.Schemas) {
|
||||
return nil, errors.New("schema not found")
|
||||
}
|
||||
return se.db.Schemas[schemaIndex], nil
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) tableAt(schemaIndex, tableIndex int) (*models.Schema, *models.Table, error) {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if tableIndex < 0 || tableIndex >= len(schema.Tables) {
|
||||
return nil, nil, errors.New("table not found")
|
||||
}
|
||||
return schema, schema.Tables[tableIndex], nil
|
||||
}
|
||||
|
||||
// splitList splits a comma separated list, trimming blanks and dropping empty entries.
|
||||
func splitList(s string) []string {
|
||||
parts := make([]string, 0)
|
||||
for _, p := range strings.Split(s, ",") {
|
||||
if p = strings.TrimSpace(p); p != "" {
|
||||
parts = append(parts, p)
|
||||
}
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
// SaveIndex adds an index to a table. When oldName is non-empty the index of that
|
||||
// name is replaced (and renamed if needed).
|
||||
func (se *SchemaEditor) SaveIndex(schemaIndex, tableIndex int, oldName string, idx *models.Index) error {
|
||||
schema, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
idx.Name = strings.TrimSpace(idx.Name)
|
||||
if idx.Name == "" {
|
||||
return errors.New("index name is required")
|
||||
}
|
||||
if len(idx.Columns) == 0 {
|
||||
return errors.New("index needs at least one column")
|
||||
}
|
||||
for _, c := range idx.Columns {
|
||||
if _, ok := table.Columns[c]; !ok {
|
||||
return fmt.Errorf("column %q not found in table %s", c, table.Name)
|
||||
}
|
||||
}
|
||||
if _, exists := table.Indexes[idx.Name]; exists && idx.Name != oldName {
|
||||
return fmt.Errorf("index %q already exists", idx.Name)
|
||||
}
|
||||
if table.Indexes == nil {
|
||||
table.Indexes = make(map[string]*models.Index)
|
||||
}
|
||||
if oldName != "" {
|
||||
delete(table.Indexes, oldName)
|
||||
}
|
||||
idx.Table = table.Name
|
||||
idx.Schema = schema.Name
|
||||
table.Indexes[idx.Name] = idx
|
||||
table.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteIndex removes an index from a table.
|
||||
func (se *SchemaEditor) DeleteIndex(schemaIndex, tableIndex int, name string) bool {
|
||||
_, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
if _, ok := table.Indexes[name]; !ok {
|
||||
return false
|
||||
}
|
||||
delete(table.Indexes, name)
|
||||
table.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return true
|
||||
}
|
||||
|
||||
// SaveView adds a view to a schema, or replaces the one at position at (use -1 to add).
|
||||
func (se *SchemaEditor) SaveView(schemaIndex, at int, v *models.View) error {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
v.Name = strings.TrimSpace(v.Name)
|
||||
if v.Name == "" {
|
||||
return errors.New("view name is required")
|
||||
}
|
||||
if strings.TrimSpace(v.Definition) == "" {
|
||||
return errors.New("view definition is required")
|
||||
}
|
||||
for i, o := range schema.Views {
|
||||
if i != at && o.Name == v.Name {
|
||||
return fmt.Errorf("view %q already exists", v.Name)
|
||||
}
|
||||
}
|
||||
v.Schema = schema.Name
|
||||
if at >= 0 && at < len(schema.Views) {
|
||||
schema.Views[at] = v
|
||||
} else {
|
||||
schema.Views = append(schema.Views, v)
|
||||
}
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteView removes the view at position at.
|
||||
func (se *SchemaEditor) DeleteView(schemaIndex, at int) bool {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil || at < 0 || at >= len(schema.Views) {
|
||||
return false
|
||||
}
|
||||
schema.Views = append(schema.Views[:at], schema.Views[at+1:]...)
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return true
|
||||
}
|
||||
|
||||
// SaveSequence adds a sequence to a schema, or replaces the one at position at (use -1 to add).
|
||||
func (se *SchemaEditor) SaveSequence(schemaIndex, at int, s *models.Sequence) error {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.Name = strings.TrimSpace(s.Name)
|
||||
if s.Name == "" {
|
||||
return errors.New("sequence name is required")
|
||||
}
|
||||
if s.IncrementBy == 0 {
|
||||
return errors.New("increment must not be zero")
|
||||
}
|
||||
for i, o := range schema.Sequences {
|
||||
if i != at && o.Name == s.Name {
|
||||
return fmt.Errorf("sequence %q already exists", s.Name)
|
||||
}
|
||||
}
|
||||
s.Schema = schema.Name
|
||||
if at >= 0 && at < len(schema.Sequences) {
|
||||
schema.Sequences[at] = s
|
||||
} else {
|
||||
schema.Sequences = append(schema.Sequences, s)
|
||||
}
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteSequence removes the sequence at position at.
|
||||
func (se *SchemaEditor) DeleteSequence(schemaIndex, at int) bool {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil || at < 0 || at >= len(schema.Sequences) {
|
||||
return false
|
||||
}
|
||||
schema.Sequences = append(schema.Sequences[:at], schema.Sequences[at+1:]...)
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return true
|
||||
}
|
||||
|
||||
// SaveScript adds a script to a schema, or replaces the one at position at (use -1 to add).
|
||||
func (se *SchemaEditor) SaveScript(schemaIndex, at int, s *models.Script) error {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.Name = strings.TrimSpace(s.Name)
|
||||
if s.Name == "" {
|
||||
return errors.New("script name is required")
|
||||
}
|
||||
if strings.TrimSpace(s.SQL) == "" {
|
||||
return errors.New("script SQL is required")
|
||||
}
|
||||
for i, o := range schema.Scripts {
|
||||
if i != at && o.Name == s.Name {
|
||||
return fmt.Errorf("script %q already exists", s.Name)
|
||||
}
|
||||
}
|
||||
s.Schema = schema.Name
|
||||
if at >= 0 && at < len(schema.Scripts) {
|
||||
schema.Scripts[at] = s
|
||||
} else {
|
||||
schema.Scripts = append(schema.Scripts, s)
|
||||
}
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteScript removes the script at position at.
|
||||
func (se *SchemaEditor) DeleteScript(schemaIndex, at int) bool {
|
||||
schema, err := se.schemaAt(schemaIndex)
|
||||
if err != nil || at < 0 || at >= len(schema.Scripts) {
|
||||
return false
|
||||
}
|
||||
schema.Scripts = append(schema.Scripts[:at], schema.Scripts[at+1:]...)
|
||||
schema.UpdateDate()
|
||||
se.db.UpdateDate()
|
||||
return true
|
||||
}
|
||||
|
||||
// AssignTableToDomain adds a reference to schemaName.tableName to the domain at domainIndex.
|
||||
func (se *SchemaEditor) AssignTableToDomain(domainIndex int, schemaName, tableName string) error {
|
||||
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||
return errors.New("domain not found")
|
||||
}
|
||||
domain := se.db.Domains[domainIndex]
|
||||
var table *models.Table
|
||||
for _, s := range se.db.Schemas {
|
||||
if s.Name != schemaName {
|
||||
continue
|
||||
}
|
||||
for _, t := range s.Tables {
|
||||
if t.Name == tableName {
|
||||
table = t
|
||||
}
|
||||
}
|
||||
}
|
||||
if table == nil {
|
||||
return fmt.Errorf("table %s.%s not found", schemaName, tableName)
|
||||
}
|
||||
for _, dt := range domain.Tables {
|
||||
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||
return fmt.Errorf("table %s.%s is already in domain %s", schemaName, tableName, domain.Name)
|
||||
}
|
||||
}
|
||||
dt := models.InitDomainTable(tableName, schemaName)
|
||||
dt.RefTable = table
|
||||
dt.Sequence = uint(len(domain.Tables))
|
||||
domain.Tables = append(domain.Tables, dt)
|
||||
se.db.UpdateDate()
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnassignTableFromDomain removes the reference to schemaName.tableName from the domain.
|
||||
func (se *SchemaEditor) UnassignTableFromDomain(domainIndex int, schemaName, tableName string) bool {
|
||||
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||
return false
|
||||
}
|
||||
domain := se.db.Domains[domainIndex]
|
||||
for i, dt := range domain.Tables {
|
||||
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||
domain.Tables = append(domain.Tables[:i], domain.Tables[i+1:]...)
|
||||
se.db.UpdateDate()
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func newTestEditor() *SchemaEditor {
|
||||
db := models.InitDatabase("test")
|
||||
schema := models.InitSchema("public")
|
||||
table := models.InitTable("users", "public")
|
||||
table.Columns["id"] = models.InitColumn("id", "users", "public")
|
||||
table.Columns["email"] = models.InitColumn("email", "users", "public")
|
||||
schema.Tables = append(schema.Tables, table)
|
||||
db.Schemas = append(db.Schemas, schema)
|
||||
return &SchemaEditor{db: db}
|
||||
}
|
||||
|
||||
func TestSaveIndex(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
table := se.db.Schemas[0].Tables[0]
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
old string
|
||||
idx *models.Index
|
||||
wantErr bool
|
||||
}{
|
||||
{"valid", "", &models.Index{Name: "idx_email", Columns: []string{"email"}, Unique: true}, false},
|
||||
{"duplicate", "", &models.Index{Name: "idx_email", Columns: []string{"email"}}, true},
|
||||
{"missing name", "", &models.Index{Columns: []string{"email"}}, true},
|
||||
{"no columns", "", &models.Index{Name: "idx_none"}, true},
|
||||
{"unknown column", "", &models.Index{Name: "idx_bad", Columns: []string{"nope"}}, true},
|
||||
{"rename", "idx_email", &models.Index{Name: "idx_email2", Columns: []string{"email", "id"}}, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if err := se.SaveIndex(0, 0, tt.old, tt.idx); (err != nil) != tt.wantErr {
|
||||
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
if _, ok := table.Indexes["idx_email"]; ok {
|
||||
t.Error("renamed index should be gone under old name")
|
||||
}
|
||||
if idx := table.Indexes["idx_email2"]; idx == nil || idx.Table != "users" || idx.Schema != "public" {
|
||||
t.Errorf("unexpected renamed index: %+v", idx)
|
||||
}
|
||||
if !se.DeleteIndex(0, 0, "idx_email2") || se.DeleteIndex(0, 0, "idx_email2") {
|
||||
t.Error("delete should succeed once")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveViewSequenceScript(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
schema := se.db.Schemas[0]
|
||||
|
||||
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 2"}); err == nil {
|
||||
t.Error("duplicate view accepted")
|
||||
}
|
||||
if err := se.SaveView(0, 0, &models.View{Name: "v", Definition: "select 3"}); err != nil {
|
||||
t.Errorf("editing in place should not conflict: %v", err)
|
||||
}
|
||||
if err := se.SaveView(0, -1, &models.View{Name: "w"}); err == nil {
|
||||
t.Error("view without definition accepted")
|
||||
}
|
||||
if len(schema.Views) != 1 || schema.Views[0].Definition != "select 3" || schema.Views[0].Schema != "public" {
|
||||
t.Errorf("unexpected views: %+v", schema.Views)
|
||||
}
|
||||
if !se.DeleteView(0, 0) || se.DeleteView(0, 0) {
|
||||
t.Error("view delete mismatch")
|
||||
}
|
||||
|
||||
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "z"}); err == nil {
|
||||
t.Error("zero increment accepted")
|
||||
}
|
||||
if !se.DeleteSequence(0, 0) || len(schema.Sequences) != 0 {
|
||||
t.Error("sequence delete failed")
|
||||
}
|
||||
|
||||
if err := se.SaveScript(0, -1, &models.Script{Name: "init", SQL: "select 1"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.SaveScript(0, -1, &models.Script{Name: "empty"}); err == nil {
|
||||
t.Error("script without SQL accepted")
|
||||
}
|
||||
if err := se.SaveScript(5, -1, &models.Script{Name: "x", SQL: "y"}); err == nil {
|
||||
t.Error("bad schema index accepted")
|
||||
}
|
||||
if !se.DeleteScript(0, 0) || len(schema.Scripts) != 0 {
|
||||
t.Error("script delete failed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDomainTableAssignment(t *testing.T) {
|
||||
se := newTestEditor()
|
||||
se.createDomainNoUI("core")
|
||||
|
||||
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := se.AssignTableToDomain(0, "public", "users"); err == nil {
|
||||
t.Error("duplicate assignment accepted")
|
||||
}
|
||||
if err := se.AssignTableToDomain(0, "public", "missing"); err == nil {
|
||||
t.Error("unknown table accepted")
|
||||
}
|
||||
if err := se.AssignTableToDomain(3, "public", "users"); err == nil {
|
||||
t.Error("bad domain index accepted")
|
||||
}
|
||||
dt := se.db.Domains[0].Tables[0]
|
||||
if dt.RefTable != se.db.Schemas[0].Tables[0] {
|
||||
t.Error("RefTable not linked")
|
||||
}
|
||||
if !se.UnassignTableFromDomain(0, "public", "users") || se.UnassignTableFromDomain(0, "public", "users") {
|
||||
t.Error("unassign mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) createDomainNoUI(name string) {
|
||||
se.db.Domains = append(se.db.Domains, models.InitDomain(name))
|
||||
}
|
||||
|
||||
func TestSplitList(t *testing.T) {
|
||||
got := splitList(" a, b,, c ,")
|
||||
if len(got) != 3 || got[0] != "a" || got[2] != "c" {
|
||||
t.Errorf("got %v", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,476 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gdamore/tcell/v2"
|
||||
"github.com/rivo/tview"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// objectLocation identifies where a new object is created: a schema, and for indexes also a table.
|
||||
type objectLocation struct {
|
||||
label string
|
||||
schemaIndex int
|
||||
tableIndex int
|
||||
}
|
||||
|
||||
// objectRow is one existing object shown in an object list.
|
||||
type objectRow struct {
|
||||
cells []string
|
||||
schemaIndex int
|
||||
tableIndex int
|
||||
at int // position within the schema slice (views, sequences, scripts)
|
||||
name string // map key (indexes)
|
||||
}
|
||||
|
||||
// objectKind describes how a kind of schema object is listed and edited.
|
||||
type objectKind struct {
|
||||
page string
|
||||
title string
|
||||
singular string
|
||||
headers []string
|
||||
rows func() []objectRow
|
||||
locations func() []objectLocation
|
||||
// buildForm adds the editable fields to the form for row (nil when creating) and
|
||||
// returns a function that validates and saves the values at the given location.
|
||||
buildForm func(form *tview.Form, row *objectRow) func(loc objectLocation) error
|
||||
remove func(row objectRow) bool
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) schemaLocations() []objectLocation {
|
||||
locs := make([]objectLocation, 0, len(se.db.Schemas))
|
||||
for si, s := range se.db.Schemas {
|
||||
locs = append(locs, objectLocation{label: s.Name, schemaIndex: si, tableIndex: -1})
|
||||
}
|
||||
return locs
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) tableLocations() []objectLocation {
|
||||
locs := make([]objectLocation, 0)
|
||||
for si, s := range se.db.Schemas {
|
||||
for ti, t := range s.Tables {
|
||||
locs = append(locs, objectLocation{label: s.Name + "." + t.Name, schemaIndex: si, tableIndex: ti})
|
||||
}
|
||||
}
|
||||
return locs
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) indexKind() objectKind {
|
||||
return objectKind{
|
||||
page: "indexes",
|
||||
title: "Manage Indexes",
|
||||
singular: "Index",
|
||||
headers: []string{"Name", "Schema", "Table", "Type", "Unique", "Columns"},
|
||||
locations: se.tableLocations,
|
||||
rows: func() []objectRow {
|
||||
var rows []objectRow
|
||||
for si, s := range se.db.Schemas {
|
||||
for ti, t := range s.Tables {
|
||||
for _, name := range sortedKeys(t.Indexes) {
|
||||
idx := t.Indexes[name]
|
||||
rows = append(rows, objectRow{
|
||||
cells: []string{idx.Name, s.Name, t.Name, idx.Type, strconv.FormatBool(idx.Unique), strings.Join(idx.Columns, ",")},
|
||||
schemaIndex: si, tableIndex: ti, name: name,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
return rows
|
||||
},
|
||||
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||
idx := models.InitIndex("", "", "")
|
||||
idx.Type = "btree"
|
||||
if row != nil {
|
||||
idx = se.db.Schemas[row.schemaIndex].Tables[row.tableIndex].Indexes[row.name]
|
||||
}
|
||||
name, columns, typ, where := idx.Name, strings.Join(idx.Columns, ", "), idx.Type, idx.Where
|
||||
unique := idx.Unique
|
||||
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||
form.AddInputField("Columns (comma separated)", columns, 50, nil, func(v string) { columns = v })
|
||||
form.AddInputField("Type", typ, 20, nil, func(v string) { typ = v })
|
||||
form.AddCheckbox("Unique", unique, func(v bool) { unique = v })
|
||||
form.AddInputField("Where", where, 50, nil, func(v string) { where = v })
|
||||
return func(loc objectLocation) error {
|
||||
oldName := ""
|
||||
if row != nil {
|
||||
oldName = row.name
|
||||
}
|
||||
next := *idx
|
||||
next.Name, next.Columns, next.Type, next.Unique, next.Where = name, splitList(columns), typ, unique, where
|
||||
return se.SaveIndex(loc.schemaIndex, loc.tableIndex, oldName, &next)
|
||||
}
|
||||
},
|
||||
remove: func(r objectRow) bool { return se.DeleteIndex(r.schemaIndex, r.tableIndex, r.name) },
|
||||
}
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) viewKind() objectKind {
|
||||
return objectKind{
|
||||
page: "views",
|
||||
title: "Manage Views",
|
||||
singular: "View",
|
||||
headers: []string{"Name", "Schema", "Description"},
|
||||
locations: se.schemaLocations,
|
||||
rows: func() []objectRow {
|
||||
var rows []objectRow
|
||||
for si, s := range se.db.Schemas {
|
||||
for i, v := range s.Views {
|
||||
rows = append(rows, objectRow{cells: []string{v.Name, s.Name, v.Description}, schemaIndex: si, at: i})
|
||||
}
|
||||
}
|
||||
return rows
|
||||
},
|
||||
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||
view := models.InitView("", "")
|
||||
at := -1
|
||||
if row != nil {
|
||||
view, at = se.db.Schemas[row.schemaIndex].Views[row.at], row.at
|
||||
}
|
||||
name, desc, def := view.Name, view.Description, view.Definition
|
||||
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||
form.AddTextArea("Definition (SQL)", def, 60, 8, 0, func(v string) { def = v })
|
||||
return func(loc objectLocation) error {
|
||||
next := *view
|
||||
next.Name, next.Description, next.Definition = name, desc, def
|
||||
return se.SaveView(loc.schemaIndex, at, &next)
|
||||
}
|
||||
},
|
||||
remove: func(r objectRow) bool { return se.DeleteView(r.schemaIndex, r.at) },
|
||||
}
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) sequenceKind() objectKind {
|
||||
return objectKind{
|
||||
page: "sequences",
|
||||
title: "Manage Sequences",
|
||||
singular: "Sequence",
|
||||
headers: []string{"Name", "Schema", "Start", "Increment", "Cycle", "Description"},
|
||||
locations: se.schemaLocations,
|
||||
rows: func() []objectRow {
|
||||
var rows []objectRow
|
||||
for si, s := range se.db.Schemas {
|
||||
for i, q := range s.Sequences {
|
||||
rows = append(rows, objectRow{
|
||||
cells: []string{q.Name, s.Name, strconv.FormatInt(q.StartValue, 10), strconv.FormatInt(q.IncrementBy, 10), strconv.FormatBool(q.Cycle), q.Description}, schemaIndex: si, at: i,
|
||||
})
|
||||
}
|
||||
}
|
||||
return rows
|
||||
},
|
||||
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||
seq := models.InitSequence("", "")
|
||||
at := -1
|
||||
if row != nil {
|
||||
seq, at = se.db.Schemas[row.schemaIndex].Sequences[row.at], row.at
|
||||
}
|
||||
name, desc := seq.Name, seq.Description
|
||||
start, incr := strconv.FormatInt(seq.StartValue, 10), strconv.FormatInt(seq.IncrementBy, 10)
|
||||
minV, maxV := strconv.FormatInt(seq.MinValue, 10), strconv.FormatInt(seq.MaxValue, 10)
|
||||
cycle := seq.Cycle
|
||||
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||
form.AddInputField("Start", start, 20, nil, func(v string) { start = v })
|
||||
form.AddInputField("Increment", incr, 20, nil, func(v string) { incr = v })
|
||||
form.AddInputField("Min (0 = none)", minV, 20, nil, func(v string) { minV = v })
|
||||
form.AddInputField("Max (0 = none)", maxV, 20, nil, func(v string) { maxV = v })
|
||||
form.AddCheckbox("Cycle", cycle, func(v bool) { cycle = v })
|
||||
return func(loc objectLocation) error {
|
||||
next := *seq
|
||||
next.Name, next.Description, next.Cycle = name, desc, cycle
|
||||
for _, f := range []struct {
|
||||
label string
|
||||
text string
|
||||
dst *int64
|
||||
}{{"start", start, &next.StartValue}, {"increment", incr, &next.IncrementBy}, {"min", minV, &next.MinValue}, {"max", maxV, &next.MaxValue}} {
|
||||
n, err := strconv.ParseInt(strings.TrimSpace(f.text), 10, 64)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s must be an integer", f.label)
|
||||
}
|
||||
*f.dst = n
|
||||
}
|
||||
return se.SaveSequence(loc.schemaIndex, at, &next)
|
||||
}
|
||||
},
|
||||
remove: func(r objectRow) bool { return se.DeleteSequence(r.schemaIndex, r.at) },
|
||||
}
|
||||
}
|
||||
|
||||
func (se *SchemaEditor) scriptKind() objectKind {
|
||||
return objectKind{
|
||||
page: "scripts",
|
||||
title: "Manage Scripts",
|
||||
singular: "Script",
|
||||
headers: []string{"Name", "Schema", "Version", "Priority", "Description"},
|
||||
locations: se.schemaLocations,
|
||||
rows: func() []objectRow {
|
||||
var rows []objectRow
|
||||
for si, s := range se.db.Schemas {
|
||||
for i, sc := range s.Scripts {
|
||||
rows = append(rows, objectRow{cells: []string{sc.Name, s.Name, sc.Version, strconv.Itoa(sc.Priority), sc.Description}, schemaIndex: si, at: i})
|
||||
}
|
||||
}
|
||||
return rows
|
||||
},
|
||||
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||
script := models.InitScript("")
|
||||
at := -1
|
||||
if row != nil {
|
||||
script, at = se.db.Schemas[row.schemaIndex].Scripts[row.at], row.at
|
||||
}
|
||||
name, desc, version, sql, rollback := script.Name, script.Description, script.Version, script.SQL, script.Rollback
|
||||
priority, runAfter := strconv.Itoa(script.Priority), strings.Join(script.RunAfter, ", ")
|
||||
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||
form.AddInputField("Version", version, 20, nil, func(v string) { version = v })
|
||||
form.AddInputField("Priority", priority, 10, nil, func(v string) { priority = v })
|
||||
form.AddInputField("Run after (comma separated)", runAfter, 50, nil, func(v string) { runAfter = v })
|
||||
form.AddTextArea("SQL", sql, 60, 8, 0, func(v string) { sql = v })
|
||||
form.AddTextArea("Rollback SQL", rollback, 60, 4, 0, func(v string) { rollback = v })
|
||||
return func(loc objectLocation) error {
|
||||
prio, err := strconv.Atoi(strings.TrimSpace(priority))
|
||||
if err != nil {
|
||||
return fmt.Errorf("priority must be an integer")
|
||||
}
|
||||
next := *script
|
||||
next.Name, next.Description, next.Version, next.Priority = name, desc, version, prio
|
||||
next.RunAfter, next.SQL, next.Rollback = splitList(runAfter), sql, rollback
|
||||
return se.SaveScript(loc.schemaIndex, at, &next)
|
||||
}
|
||||
},
|
||||
remove: func(r objectRow) bool { return se.DeleteScript(r.schemaIndex, r.at) },
|
||||
}
|
||||
}
|
||||
|
||||
func sortedKeys[V any](m map[string]V) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
// showObjectList displays all objects of a kind across schemas.
|
||||
func (se *SchemaEditor) showObjectList(k objectKind) {
|
||||
flex := tview.NewFlex().SetDirection(tview.FlexRow)
|
||||
title := tview.NewTextView().SetText("[::b]" + k.title).SetDynamicColors(true).SetTextAlign(tview.AlignCenter)
|
||||
|
||||
table := tview.NewTable().SetBorders(true).SetSelectable(true, false).SetFixed(1, 0)
|
||||
for i, h := range k.headers {
|
||||
table.SetCell(0, i, tview.NewTableCell(h).SetTextColor(tcell.ColorYellow).SetSelectable(false).SetAlign(tview.AlignLeft))
|
||||
}
|
||||
rows := k.rows()
|
||||
for r, row := range rows {
|
||||
for c, text := range row.cells {
|
||||
table.SetCell(r+1, c, tview.NewTableCell(text).SetSelectable(true))
|
||||
}
|
||||
}
|
||||
table.SetTitle(" " + k.title[len("Manage "):] + " ").SetBorder(true).SetTitleAlign(tview.AlignLeft)
|
||||
|
||||
back := func() {
|
||||
se.pages.SwitchToPage("main")
|
||||
se.pages.RemovePage(k.page)
|
||||
}
|
||||
btnNew := tview.NewButton("New " + k.singular + " [n]").SetSelectedFunc(func() { se.showObjectForm(k, nil) })
|
||||
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(back)
|
||||
btnNew.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
switch event.Key() {
|
||||
case tcell.KeyBacktab:
|
||||
se.app.SetFocus(table)
|
||||
return nil
|
||||
case tcell.KeyTab:
|
||||
se.app.SetFocus(btnBack)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
btnBack.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
switch event.Key() {
|
||||
case tcell.KeyBacktab:
|
||||
se.app.SetFocus(btnNew)
|
||||
return nil
|
||||
case tcell.KeyTab:
|
||||
se.app.SetFocus(table)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
btnFlex := tview.NewFlex().AddItem(btnNew, 0, 1, true).AddItem(btnBack, 0, 1, false)
|
||||
|
||||
table.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
switch {
|
||||
case event.Key() == tcell.KeyEscape, event.Rune() == 'b':
|
||||
back()
|
||||
return nil
|
||||
case event.Key() == tcell.KeyTab:
|
||||
se.app.SetFocus(btnNew)
|
||||
return nil
|
||||
case event.Key() == tcell.KeyEnter:
|
||||
if row, _ := table.GetSelection(); row > 0 && row <= len(rows) {
|
||||
se.showObjectForm(k, &rows[row-1])
|
||||
return nil
|
||||
}
|
||||
case event.Rune() == 'n':
|
||||
se.showObjectForm(k, nil)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
|
||||
flex.AddItem(title, 1, 0, false).AddItem(table, 0, 1, true).AddItem(btnFlex, 1, 0, false)
|
||||
se.pages.AddPage(k.page, flex, true, true)
|
||||
}
|
||||
|
||||
// showObjectForm shows the create (row == nil) or edit form for an object.
|
||||
func (se *SchemaEditor) showObjectForm(k objectKind, row *objectRow) {
|
||||
formPage := k.page + "-form"
|
||||
form := tview.NewForm()
|
||||
errView := tview.NewTextView().SetDynamicColors(true)
|
||||
|
||||
locs := k.locations()
|
||||
loc := objectLocation{schemaIndex: -1, tableIndex: -1}
|
||||
switch {
|
||||
case row != nil:
|
||||
loc = objectLocation{schemaIndex: row.schemaIndex, tableIndex: row.tableIndex}
|
||||
case len(locs) > 0:
|
||||
loc = locs[0]
|
||||
labels := make([]string, len(locs))
|
||||
for i, l := range locs {
|
||||
labels[i] = l.label
|
||||
}
|
||||
form.AddDropDown("Location", labels, 0, func(_ string, i int) { loc = locs[i] })
|
||||
}
|
||||
|
||||
save := k.buildForm(form, row)
|
||||
|
||||
closeForm := func() {
|
||||
se.pages.RemovePage(formPage)
|
||||
se.pages.RemovePage(k.page)
|
||||
se.showObjectList(k)
|
||||
}
|
||||
form.AddButton("Save", func() {
|
||||
if err := save(loc); err != nil {
|
||||
errView.SetText("[red]" + tview.Escape(err.Error()))
|
||||
return
|
||||
}
|
||||
closeForm()
|
||||
})
|
||||
if row != nil {
|
||||
form.AddButton("Delete", func() {
|
||||
modal := tview.NewModal().
|
||||
SetText(fmt.Sprintf("Delete %s '%s'? This action cannot be undone.", strings.ToLower(k.singular), row.cells[0])).
|
||||
AddButtons([]string{"Cancel", "Delete"}).
|
||||
SetDoneFunc(func(_ int, label string) {
|
||||
se.pages.RemovePage(formPage + "-delete")
|
||||
if label == "Delete" {
|
||||
k.remove(*row)
|
||||
closeForm()
|
||||
}
|
||||
})
|
||||
se.pages.AddAndSwitchToPage(formPage+"-delete", modal, true)
|
||||
})
|
||||
}
|
||||
form.AddButton("Back", closeForm)
|
||||
|
||||
verb := "New"
|
||||
if row != nil {
|
||||
verb = "Edit"
|
||||
}
|
||||
form.SetBorder(true).SetTitle(" " + verb + " " + k.singular + " ").SetTitleAlign(tview.AlignLeft)
|
||||
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() == tcell.KeyEscape {
|
||||
se.showExitConfirmation(formPage, k.page)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
|
||||
if len(locs) == 0 && row == nil {
|
||||
errView.SetText("[red]No schema/table available. Create one first.")
|
||||
}
|
||||
flex := tview.NewFlex().SetDirection(tview.FlexRow).AddItem(form, 0, 1, true).AddItem(errView, 1, 0, false)
|
||||
se.pages.AddPage(formPage, flex, true, true)
|
||||
}
|
||||
|
||||
// showDomainTables lists the tables assigned to a domain and allows assigning/unassigning.
|
||||
func (se *SchemaEditor) showDomainTables(domainIndex int) {
|
||||
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||
return
|
||||
}
|
||||
domain := se.db.Domains[domainIndex]
|
||||
page := "domain-tables"
|
||||
list := tview.NewList().ShowSecondaryText(true)
|
||||
refresh := func() {
|
||||
se.pages.RemovePage(page)
|
||||
se.showDomainTables(domainIndex)
|
||||
}
|
||||
|
||||
for _, dt := range domain.Tables {
|
||||
dt := dt
|
||||
list.AddItem(dt.SchemaName+"."+dt.TableName, "Enter to remove from domain", 0, func() {
|
||||
se.UnassignTableFromDomain(domainIndex, dt.SchemaName, dt.TableName)
|
||||
refresh()
|
||||
})
|
||||
}
|
||||
list.AddItem("[Assign Table]", "Add a table to this domain", 'a', func() {
|
||||
se.showAssignDomainTable(domainIndex, refresh)
|
||||
})
|
||||
list.AddItem("[Back]", "Return to domain", 'b', func() {
|
||||
se.pages.RemovePage(page)
|
||||
})
|
||||
list.SetBorder(true).SetTitle(" Domain " + domain.Name + " - Tables ").SetTitleAlign(tview.AlignLeft)
|
||||
list.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() == tcell.KeyEscape {
|
||||
se.pages.RemovePage(page)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
se.pages.AddPage(page, list, true, true)
|
||||
}
|
||||
|
||||
// showAssignDomainTable shows a form to pick a table not yet in the domain.
|
||||
func (se *SchemaEditor) showAssignDomainTable(domainIndex int, done func()) {
|
||||
page := "assign-domain-table"
|
||||
domain := se.db.Domains[domainIndex]
|
||||
var options []string
|
||||
var refs []models.DomainTable
|
||||
for _, s := range se.db.Schemas {
|
||||
for _, t := range s.Tables {
|
||||
taken := false
|
||||
for _, dt := range domain.Tables {
|
||||
taken = taken || (dt.SchemaName == s.Name && dt.TableName == t.Name)
|
||||
}
|
||||
if !taken {
|
||||
options = append(options, s.Name+"."+t.Name)
|
||||
refs = append(refs, models.DomainTable{SchemaName: s.Name, TableName: t.Name})
|
||||
}
|
||||
}
|
||||
}
|
||||
form := tview.NewForm()
|
||||
selected := 0
|
||||
form.AddDropDown("Table", options, 0, func(_ string, i int) { selected = i })
|
||||
form.AddButton("Assign", func() {
|
||||
if len(refs) > 0 {
|
||||
_ = se.AssignTableToDomain(domainIndex, refs[selected].SchemaName, refs[selected].TableName)
|
||||
}
|
||||
se.pages.RemovePage(page)
|
||||
done()
|
||||
})
|
||||
form.AddButton("Back", func() { se.pages.RemovePage(page) })
|
||||
form.SetBorder(true).SetTitle(" Assign Table ").SetTitleAlign(tview.AlignLeft)
|
||||
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||
if event.Key() == tcell.KeyEscape {
|
||||
se.pages.RemovePage(page)
|
||||
return nil
|
||||
}
|
||||
return event
|
||||
})
|
||||
se.pages.AddPage(page, form, true, true)
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
// richEditor returns an editor whose database has one of every object the screens render.
|
||||
func richEditor(t *testing.T) *SchemaEditor {
|
||||
t.Helper()
|
||||
se := NewSchemaEditor(newTestEditor().db)
|
||||
db := se.db
|
||||
tbl := db.Schemas[0].Tables[0]
|
||||
col := tbl.Columns["id"]
|
||||
col.Type, col.IsPrimaryKey, col.NotNull = "integer", true, true
|
||||
tbl.Relationships["fk_self"] = &models.Relationship{Name: "fk_self", FromTable: "users", ToTable: "users", FromColumns: []string{"id"}, ToColumns: []string{"id"}}
|
||||
se.createDomainNoUI("core")
|
||||
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = se.SaveIndex(0, 0, "", &models.Index{Name: "idx_e", Columns: []string{"email"}})
|
||||
_ = se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"})
|
||||
_ = se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1})
|
||||
_ = se.SaveScript(0, -1, &models.Script{Name: "sc", SQL: "select 1"})
|
||||
return se
|
||||
}
|
||||
|
||||
// TestScreensRender builds every screen and dialog against a populated database
|
||||
// and checks that none panics and that each registers a page.
|
||||
func TestScreensRender(t *testing.T) {
|
||||
col := func(se *SchemaEditor) *models.Column { return se.db.Schemas[0].Tables[0].Columns["id"] }
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
page string
|
||||
run func(se *SchemaEditor)
|
||||
}{
|
||||
{"schema list", "schemas", func(se *SchemaEditor) { se.showSchemaList() }},
|
||||
{"schema editor", "schema-editor", func(se *SchemaEditor) { se.showSchemaEditor(0, se.db.Schemas[0]) }},
|
||||
{"new schema", "new-schema", func(se *SchemaEditor) { se.showNewSchemaDialog() }},
|
||||
{"edit schema", "edit-schema", func(se *SchemaEditor) { se.showEditSchemaDialog(0) }},
|
||||
{"table list", "tables", func(se *SchemaEditor) { se.showTableList() }},
|
||||
{"table editor", "table-editor", func(se *SchemaEditor) { se.showTableEditor(0, 0, se.db.Schemas[0].Tables[0]) }},
|
||||
{"new table", "new-table", func(se *SchemaEditor) { se.showNewTableDialog(0) }},
|
||||
{"new table from list", "new-table-from-list", func(se *SchemaEditor) { se.showNewTableDialogFromList() }},
|
||||
{"edit table", "edit-table", func(se *SchemaEditor) { se.showEditTableDialog(0, 0) }},
|
||||
{"column editor", "column-editor", func(se *SchemaEditor) { se.showColumnEditor(0, 0, 0, col(se)) }},
|
||||
{"new column", "new-column", func(se *SchemaEditor) { se.showNewColumnDialog(0, 0) }},
|
||||
{"relationship list", "relationships", func(se *SchemaEditor) { se.showRelationshipList(0, 0) }},
|
||||
{"new relationship", "new-relationship", func(se *SchemaEditor) { se.showNewRelationshipDialog(0, 0) }},
|
||||
{"edit relationship", "edit-relationship", func(se *SchemaEditor) { se.showEditRelationshipDialog(0, 0, "fk_self") }},
|
||||
{"delete relationship", "delete-relationship-confirm", func(se *SchemaEditor) { se.showDeleteRelationshipConfirm(0, 0, "fk_self") }},
|
||||
{"domain list", "domains", func(se *SchemaEditor) { se.showDomainList() }},
|
||||
{"new domain", "new-domain", func(se *SchemaEditor) { se.showNewDomainDialog() }},
|
||||
{"domain editor", "edit-domain", func(se *SchemaEditor) { se.showDomainEditor(0, se.db.Domains[0]) }},
|
||||
{"delete domain", "delete-domain-confirm", func(se *SchemaEditor) { se.showDeleteDomainConfirm(0) }},
|
||||
{"domain tables", "domain-tables", func(se *SchemaEditor) { se.showDomainTables(0) }},
|
||||
{"assign domain table", "assign-domain-table", func(se *SchemaEditor) { se.showAssignDomainTable(0, func() {}) }},
|
||||
{"edit database", "edit-database", func(se *SchemaEditor) { se.showEditDatabaseForm() }},
|
||||
{"exit confirm", "exit-confirm", func(se *SchemaEditor) { se.showExitConfirmation("a", "main") }},
|
||||
{"exit editor confirm", "exit-editor-confirm", func(se *SchemaEditor) { se.showExitEditorConfirm() }},
|
||||
{"delete schema confirm", "confirm-delete-schema", func(se *SchemaEditor) { se.showDeleteSchemaConfirm(0) }},
|
||||
{"delete table confirm", "confirm-delete-table", func(se *SchemaEditor) { se.showDeleteTableConfirm(0, 0) }},
|
||||
{"delete column confirm", "confirm-delete-column", func(se *SchemaEditor) { se.showDeleteColumnConfirm(0, 0, "id") }},
|
||||
{"load screen", "load-database", func(se *SchemaEditor) { se.showLoadScreen() }},
|
||||
{"save screen", "save-database", func(se *SchemaEditor) { se.showSaveScreen() }},
|
||||
{"import screen", "import-database", func(se *SchemaEditor) { se.showImportScreen() }},
|
||||
{"update existing confirm", "update-confirm", func(se *SchemaEditor) {
|
||||
se.loadConfig = &LoadConfig{SourceType: "json", FilePath: "x.json"}
|
||||
se.showUpdateExistingDatabaseConfirm()
|
||||
}},
|
||||
{"import confirm", "import-confirm", func(se *SchemaEditor) {
|
||||
se.showImportConfirmation(models.InitDatabase("src"), false, false, false, false, false, "")
|
||||
}},
|
||||
{"conn builder", "", func(se *SchemaEditor) { se.showConnStringBuilder("", "", "main", func(string) {}) }},
|
||||
}
|
||||
for _, tt := range cases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
se := richEditor(t)
|
||||
before := len(se.pages.GetPageNames(false))
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
t.Fatalf("panic: %v", r)
|
||||
}
|
||||
}()
|
||||
tt.run(se)
|
||||
if tt.page != "" && !se.pages.HasPage(tt.page) {
|
||||
t.Errorf("page %q not registered; pages: %v", tt.page, se.pages.GetPageNames(false))
|
||||
}
|
||||
if tt.page == "" && len(se.pages.GetPageNames(false)) <= before {
|
||||
t.Error("no page added")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestObjectScreensRender(t *testing.T) {
|
||||
se := richEditor(t)
|
||||
for name, k := range map[string]objectKind{
|
||||
"indexes": se.indexKind(), "views": se.viewKind(), "sequences": se.sequenceKind(), "scripts": se.scriptKind(),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
se.showObjectList(k)
|
||||
if !se.pages.HasPage(k.page) {
|
||||
t.Errorf("list page %q missing; pages: %v", k.page, se.pages.GetPageNames(false))
|
||||
}
|
||||
rows := k.rows()
|
||||
se.showObjectForm(k, nil)
|
||||
se.showObjectForm(k, &rows[0])
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -205,12 +205,24 @@ Organize UI code into these files:
|
||||
- **column_screens.go** - Column editor, new column dialog
|
||||
- **domain_screens.go** - Domain list, domain editor, new/edit domain dialogs
|
||||
- **dialogs.go** - Confirmation dialogs (exit, delete)
|
||||
- **filebrowser_screens.go** - File browser dialog (`file-browser`), opened with Enter on File Path inputs
|
||||
- **connstring_screens.go** - Connection string builder dialog (`conn-builder`), opened with Enter on Connection String inputs
|
||||
|
||||
### Data Operations Files (Business Logic)
|
||||
|
||||
- **schema_dataops.go** - Schema CRUD operations (Create, Read, Update, Delete)
|
||||
- **table_dataops.go** - Table CRUD operations
|
||||
- **column_dataops.go** - Column CRUD operations
|
||||
- **filebrowser.go** - Directory listing, extension filtering and start-path resolution (no tview)
|
||||
- **connstring.go**, **connstring_check.go** - Connection string build/parse/mask and connection test (no tview)
|
||||
|
||||
### Input Dialogs
|
||||
|
||||
- **File browser** - Enter on a File Path input opens it. Up/Down move, Enter opens a directory or selects a file,
|
||||
Backspace/`[..]` goes to the parent, `h` toggles hidden files, `f` toggles the format extension filter,
|
||||
`s` selects, Esc/`b` cancels (input unchanged). Save mode adds a file name field and asks before overwriting.
|
||||
- **Connection string builder** - Enter on a Connection String input opens it. Fields per type (PostgreSQL, MSSQL,
|
||||
SQLite), masked password and preview, F2 save, F3 test connection, Esc cancels (input unchanged).
|
||||
|
||||
## Code Separation Rules
|
||||
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
// Package updatecheck looks up the latest RelSpec release on the project's
|
||||
// Gitea instance and compares it against the running version.
|
||||
package updatecheck
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// DefaultAPIURL is the Gitea endpoint that returns the latest release.
|
||||
const DefaultAPIURL = "https://git.warky.dev/api/v1/repos/wdevs/relspecgo/releases/latest"
|
||||
|
||||
// Asset is a downloadable file attached to a release.
|
||||
type Asset struct {
|
||||
Name string `json:"name"`
|
||||
URL string `json:"browser_download_url"`
|
||||
}
|
||||
|
||||
// Release is the subset of the Gitea release payload that is needed.
|
||||
type Release struct {
|
||||
Tag string `json:"tag_name"`
|
||||
URL string `json:"html_url"`
|
||||
Assets []Asset `json:"assets"`
|
||||
}
|
||||
|
||||
// Latest fetches the latest release from apiURL. A nil client uses a client
|
||||
// with a 10 second timeout.
|
||||
func Latest(ctx context.Context, client *http.Client, apiURL string) (*Release, error) {
|
||||
if client == nil {
|
||||
client = &http.Client{Timeout: 10 * time.Second}
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("checking for updates: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("checking for updates: unexpected status %s", resp.Status)
|
||||
}
|
||||
|
||||
var rel Release
|
||||
if err := json.NewDecoder(resp.Body).Decode(&rel); err != nil {
|
||||
return nil, fmt.Errorf("decoding release: %w", err)
|
||||
}
|
||||
if rel.Tag == "" {
|
||||
return nil, fmt.Errorf("release has no tag")
|
||||
}
|
||||
return &rel, nil
|
||||
}
|
||||
|
||||
// FindAsset returns the first asset with the given name.
|
||||
func (r *Release) FindAsset(name string) (Asset, bool) {
|
||||
for _, a := range r.Assets {
|
||||
if a.Name == name {
|
||||
return a, true
|
||||
}
|
||||
}
|
||||
return Asset{}, false
|
||||
}
|
||||
|
||||
// IsNewer reports whether latest is a higher version than current. Versions
|
||||
// that are not dotted numbers (for example "dev" or a commit hash) are never
|
||||
// considered outdated, so development builds are not nagged.
|
||||
func IsNewer(current, latest string) bool {
|
||||
cur, ok := parseVersion(current)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
lat, ok := parseVersion(latest)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
for i := range cur {
|
||||
if lat[i] != cur[i] {
|
||||
return lat[i] > cur[i]
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// parseVersion parses "v1.2.3" style versions into three numeric parts.
|
||||
// Missing parts are zero; any pre-release or build suffix is ignored.
|
||||
func parseVersion(v string) ([3]int, bool) {
|
||||
var out [3]int
|
||||
v = strings.TrimPrefix(strings.TrimSpace(v), "v")
|
||||
if i := strings.IndexAny(v, "-+ "); i >= 0 {
|
||||
v = v[:i]
|
||||
}
|
||||
if v == "" {
|
||||
return out, false
|
||||
}
|
||||
parts := strings.Split(v, ".")
|
||||
if len(parts) > 3 {
|
||||
return out, false
|
||||
}
|
||||
for i, p := range parts {
|
||||
n, err := strconv.Atoi(p)
|
||||
if err != nil || n < 0 {
|
||||
return out, false
|
||||
}
|
||||
out[i] = n
|
||||
}
|
||||
return out, true
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package updatecheck
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsNewer(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
current, latest string
|
||||
want bool
|
||||
}{
|
||||
{"patch bump", "v1.0.85", "v1.0.86", true},
|
||||
{"minor beats patch", "v1.0.99", "v1.1.0", true},
|
||||
{"major bump", "v1.9.9", "v2.0.0", true},
|
||||
{"numeric not lexical", "v1.0.9", "v1.0.10", true},
|
||||
{"same", "v1.0.85", "v1.0.85", false},
|
||||
{"older latest", "v1.0.86", "v1.0.85", false},
|
||||
{"no v prefix", "1.0.1", "v1.0.2", true},
|
||||
{"short version", "v1.0", "v1.0.1", true},
|
||||
{"prerelease suffix ignored", "v1.0.1-rc1", "v1.0.2", true},
|
||||
{"dev build", "dev", "v9.9.9", false},
|
||||
{"commit hash", "abc1234", "v9.9.9", false},
|
||||
{"empty current", "", "v1.0.0", false},
|
||||
{"bad latest", "v1.0.0", "latest", false},
|
||||
{"too many parts", "v1.0.0.0", "v1.0.1", false},
|
||||
{"negative", "v1.-1.0", "v1.0.0", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := IsNewer(tt.current, tt.latest); got != tt.want {
|
||||
t.Errorf("IsNewer(%q,%q) = %v, want %v", tt.current, tt.latest, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLatest(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
status int
|
||||
body string
|
||||
wantTag string
|
||||
wantErr bool
|
||||
}{
|
||||
{"ok", 200, `{"tag_name":"v1.2.3","html_url":"https://x/r","assets":[{"name":"a.exe","browser_download_url":"https://x/a.exe"}]}`, "v1.2.3", false},
|
||||
{"not found", 404, `{}`, "", true},
|
||||
{"bad json", 200, `{`, "", true},
|
||||
{"missing tag", 200, `{"html_url":"u"}`, "", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(tt.status)
|
||||
_, _ = w.Write([]byte(tt.body))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
rel, err := Latest(context.Background(), nil, srv.URL)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if err == nil {
|
||||
if rel.Tag != tt.wantTag {
|
||||
t.Errorf("tag = %q", rel.Tag)
|
||||
}
|
||||
if a, ok := rel.FindAsset("a.exe"); !ok || a.URL != "https://x/a.exe" {
|
||||
t.Errorf("asset = %+v %v", a, ok)
|
||||
}
|
||||
if _, ok := rel.FindAsset("missing"); ok {
|
||||
t.Error("unexpected asset")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLatestUnreachable(t *testing.T) {
|
||||
srv := httptest.NewServer(http.NotFoundHandler())
|
||||
url := srv.URL
|
||||
srv.Close()
|
||||
if _, err := Latest(context.Background(), nil, url); err == nil {
|
||||
t.Error("expected error for unreachable server")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLatestBadURL(t *testing.T) {
|
||||
if _, err := Latest(context.Background(), nil, "://bad"); err == nil {
|
||||
t.Error("expected error for bad URL")
|
||||
}
|
||||
}
|
||||
@@ -129,6 +129,9 @@ type User struct {
|
||||
- `default` - Default value
|
||||
- `rel` - Relationship definition
|
||||
- `type` - Explicit SQL type
|
||||
- `scanonly` - Excluded from INSERT/UPDATE: `GENERATED ... STORED` columns, and `GENERATED ALWAYS AS IDENTITY` columns that are not the primary key
|
||||
- `generated` - Marks a `GENERATED ... STORED` column
|
||||
- `identity` - Marks a `GENERATED ALWAYS AS IDENTITY` column (a primary key gets only this marker, since Bun drops `scanonly` fields from the PK list)
|
||||
|
||||
## Type Mapping
|
||||
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package bun
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func TestBuildBunTag_GeneratedColumn(t *testing.T) {
|
||||
mapper := NewTypeMapper("", "")
|
||||
|
||||
generated := models.InitColumn("full_name", "users", "public")
|
||||
generated.Type = "text"
|
||||
generated.Generated = true
|
||||
generated.GenerationExpression = "first || ' ' || last"
|
||||
|
||||
plain := models.InitColumn("first", "users", "public")
|
||||
plain.Type = "text"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
column *models.Column
|
||||
want bool
|
||||
}{
|
||||
{"generated column is scanonly and marked", generated, true},
|
||||
{"ordinary column is untouched", plain, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tag := mapper.BuildBunTag(tt.column, nil)
|
||||
parts := strings.Split(tag, ",")
|
||||
has := func(s string) bool {
|
||||
for _, p := range parts {
|
||||
if p == s {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
if has("scanonly") != tt.want || has("generated") != tt.want {
|
||||
t.Errorf("tag %q: scanonly/generated present = %v/%v, want %v", tag, has("scanonly"), has("generated"), tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package bun
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
)
|
||||
|
||||
func TestBuildBunTag_IdentityColumn(t *testing.T) {
|
||||
mapper := NewTypeMapper("", "")
|
||||
|
||||
build := func(pk, identity bool, generation string) *models.Column {
|
||||
col := models.InitColumn("id", "users", "public")
|
||||
col.Type = "bigint"
|
||||
col.IsPrimaryKey = pk
|
||||
col.Identity = identity
|
||||
col.IdentityGeneration = generation
|
||||
return col
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
column *models.Column
|
||||
wantFlag bool
|
||||
wantMarker bool
|
||||
}{
|
||||
{"always identity non-pk is write-blocked", build(false, true, "ALWAYS"), true, true},
|
||||
{"always identity pk is only marked", build(true, true, "ALWAYS"), false, true},
|
||||
{"by default identity is writable", build(false, true, "BY DEFAULT"), false, false},
|
||||
{"non-identity is untouched", build(false, false, ""), false, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tag := mapper.BuildBunTag(tt.column, nil)
|
||||
has := func(s string) bool {
|
||||
for _, p := range strings.Split(tag, ",") {
|
||||
if p == s {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
if has("scanonly") != tt.wantFlag || has("identity") != tt.wantMarker {
|
||||
t.Errorf("tag %q: scanonly/identity = %v/%v, want %v/%v", tag, has("scanonly"), has("identity"), tt.wantFlag, tt.wantMarker)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user