Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
08e1417393 | ||
|
|
70282fff73 | ||
|
|
43265dac0f | ||
|
|
66b90ca54b | ||
|
|
47108809aa | ||
|
|
720476fd6e | ||
|
|
572d03fe42 | ||
|
|
f1b9079b2d | ||
|
|
bb671c3680 | ||
|
|
bc8284db25 | ||
|
|
29e747393d | ||
|
|
df980a3434 | ||
|
|
778379538b | ||
|
|
7a9219b6e3 | ||
|
|
fd9c37cd25 | ||
|
|
0235a28add | ||
|
|
d961536186 | ||
|
|
d36806047b | ||
|
|
948419ffd3 | ||
|
|
b38f53c603 | ||
|
|
ccba53c494 | ||
|
|
53327b9a5a | ||
|
|
734b14d48d | ||
|
|
938f0ed51f | ||
|
|
6e2e7eb19e |
@@ -1,4 +1,4 @@
|
|||||||
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version godoc vet fmt fmt-check staticcheck govulncheck check
|
.PHONY: all build test test-unit test-integration lint coverage clean install help docker-up docker-down docker-test docker-test-integration start stop release release-version rerelease godoc vet fmt fmt-check staticcheck govulncheck check
|
||||||
|
|
||||||
# Binary name
|
# Binary name
|
||||||
BINARY_NAME=relspec
|
BINARY_NAME=relspec
|
||||||
@@ -260,5 +260,13 @@ release-version: lint fmt-check ## Run lint and format check, then auto-incremen
|
|||||||
git push origin HEAD "$$NEXT"; \
|
git push origin HEAD "$$NEXT"; \
|
||||||
echo "Pushed $$NEXT — release workflow triggered"
|
echo "Pushed $$NEXT — release workflow triggered"
|
||||||
|
|
||||||
|
rerelease: lint fmt-check ## Move the latest tag to HEAD and force push it
|
||||||
|
@TAG=$$(git describe --tags --abbrev=0 2>/dev/null); \
|
||||||
|
if [ -z "$$TAG" ]; then echo "No existing tags found"; exit 1; fi; \
|
||||||
|
echo "Moving $$TAG to $$(git rev-parse --short HEAD)"; \
|
||||||
|
git tag -f -a "$$TAG" -m "Release $$TAG" HEAD; \
|
||||||
|
git push --force origin "$$TAG"; \
|
||||||
|
echo "Pushed $$TAG — release workflow triggered"
|
||||||
|
|
||||||
help: ## Display this help screen
|
help: ## Display this help screen
|
||||||
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
|
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-20s\033[0m %s\n", $$1, $$2}'
|
||||||
|
|||||||
@@ -23,6 +23,8 @@ go install -v git.warky.dev/wdevs/relspecgo/cmd/relspec@latest
|
|||||||
| **Readers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqldir` `sqlite` `typeorm` `yaml` |
|
| **Readers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqldir` `sqlite` `typeorm` `yaml` |
|
||||||
| **Writers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqlexec` `sqlite` `template` `typeorm` `yaml` |
|
| **Writers** | `bun` `dbml` `dctx` `drawdb` `drizzle` `gorm` `graphql` `json` `mssql` `pgsql` `prisma` `sqlexec` `sqlite` `template` `typeorm` `yaml` |
|
||||||
|
|
||||||
|
See [docs/FORMAT_EXAMPLES.md](docs/FORMAT_EXAMPLES.md) for usage examples covering every format.
|
||||||
|
|
||||||
## Commands
|
## Commands
|
||||||
|
|
||||||
### `convert` — Schema conversion
|
### `convert` — Schema conversion
|
||||||
@@ -40,6 +42,26 @@ relspec convert --from pgsql --from-conn "postgres://..." --to sqlite --to-path
|
|||||||
|
|
||||||
# Multiple input files merged
|
# Multiple input files merged
|
||||||
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
||||||
|
|
||||||
|
# Watch mode: regenerate whenever the source file(s) change (Ctrl-C to stop)
|
||||||
|
relspec convert --from dbml --from-path schema.dbml --to gorm --to-path models/ --package models --watch
|
||||||
|
```
|
||||||
|
|
||||||
|
`--watch` works with `--from-path` and `--from-list` (not live database
|
||||||
|
connections or `--dry-run`). Source files are polled every `--watch-interval`
|
||||||
|
(default 500ms), a directory source is watched recursively, and the output path
|
||||||
|
is ignored so generating into the source tree does not loop. Conversion errors
|
||||||
|
are printed and watching continues.
|
||||||
|
|
||||||
|
### `batch` — Convert many inputs in one run
|
||||||
|
|
||||||
|
Converts each input independently (one output per input, unlike `--from-list`
|
||||||
|
which merges). `--input` takes paths or globs; outputs go to `--to-dir`.
|
||||||
|
Use `--keep-going` to continue past failures (exit is still non-zero) and
|
||||||
|
`--dry-run` to validate without writing. For named workflows see `relspec job run`.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec batch --from dbml --input "schemas/*.dbml" --to json --to-dir out/
|
||||||
```
|
```
|
||||||
|
|
||||||
PostgreSQL connections opened by relspec set `application_name` by default to
|
PostgreSQL connections opened by relspec set `application_name` by default to
|
||||||
@@ -216,6 +238,23 @@ see [`bun`'s `--array-nullable`](./pkg/writers/bun/README.md#nullablearrays)
|
|||||||
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
flag for nullable-array handling. The `SqlXxxArray` wrapper types remain
|
||||||
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
available in `pkg/sqltypes` and are still used by the `gorm` writer.
|
||||||
|
|
||||||
|
#### Custom type mapping
|
||||||
|
|
||||||
|
Override the built-in SQL → Go mapping of the `bun` and `gorm` writers with the
|
||||||
|
repeatable `--type-map sqltype=gotype` flag:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec convert --from pgsql --from-conn "$DSN" --to gorm --to-path models.go \
|
||||||
|
--type-map uuid=string --type-map jsonb=json.RawMessage
|
||||||
|
```
|
||||||
|
|
||||||
|
SQL type names are matched case-insensitively on the base type (modifiers such
|
||||||
|
as `(10,2)` are ignored; aliases like `int4` resolve to `integer`). NOT NULL
|
||||||
|
columns use the Go type verbatim, nullable columns get a `*` prefix (unless the
|
||||||
|
type is already a pointer, slice, map or `any`), and arrays become `[]gotype`.
|
||||||
|
Unmapped types keep their defaults. The flag does not add imports: use types
|
||||||
|
that need none, or add the import afterwards (e.g. with `goimports`).
|
||||||
|
|
||||||
## Contributing
|
## Contributing
|
||||||
|
|
||||||
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
1. Register or sign in with GitHub at [git.warky.dev](https://git.warky.dev)
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
batchSourceType string
|
||||||
|
batchInputs []string
|
||||||
|
batchTargetType string
|
||||||
|
batchTargetDir string
|
||||||
|
batchPackageName string
|
||||||
|
batchSchemaFilter string
|
||||||
|
batchFlattenSchema bool
|
||||||
|
batchNullableTypes string
|
||||||
|
batchNullableArrays string
|
||||||
|
batchContinueOnError bool
|
||||||
|
batchKeepGoing bool
|
||||||
|
batchDryRun bool
|
||||||
|
)
|
||||||
|
|
||||||
|
var batchCmd = &cobra.Command{
|
||||||
|
Use: "batch",
|
||||||
|
Short: "Convert many input files to a target format in one run",
|
||||||
|
Long: `Convert each input file independently to the target format.
|
||||||
|
|
||||||
|
Unlike 'convert --from-list', which merges all inputs into one output, batch
|
||||||
|
mode writes one output per input into --to-dir. The output is named after the
|
||||||
|
input file (without its extension). Directory-style targets (gorm, bun,
|
||||||
|
drizzle) get a sub-directory per input.
|
||||||
|
|
||||||
|
Inputs are given with --input, which accepts file paths and glob patterns and
|
||||||
|
may be repeated or comma-separated. Inputs are processed in sorted order and
|
||||||
|
duplicates are removed. The command exits non-zero if any input fails.
|
||||||
|
|
||||||
|
For named, multi-step workflows use 'relspec job run' instead.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
# Convert every DBML file in a directory to JSON
|
||||||
|
relspec batch --from dbml --input "schemas/*.dbml" --to json --to-dir out/
|
||||||
|
|
||||||
|
# Convert specific files to GORM models, one package directory per input
|
||||||
|
relspec batch --from json --input a.json,b.json \
|
||||||
|
--to gorm --to-dir models/ --package models
|
||||||
|
|
||||||
|
# Validate everything first, writing nothing
|
||||||
|
relspec batch --from yaml --input "specs/*.yaml" --to pgsql --to-dir sql/ --dry-run
|
||||||
|
|
||||||
|
# Report all failures instead of stopping at the first
|
||||||
|
relspec batch --from json --input "*.json" --to yaml --to-dir out/ --keep-going`,
|
||||||
|
RunE: runBatch,
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
batchCmd.Flags().StringVar(&batchSourceType, "from", "", "Source format for every input (dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, sqlite)")
|
||||||
|
batchCmd.Flags().StringSliceVar(&batchInputs, "input", nil, "Input file path or glob pattern (repeatable, comma-separated)")
|
||||||
|
batchCmd.Flags().StringVar(&batchTargetType, "to", "", "Target format")
|
||||||
|
batchCmd.Flags().StringVar(&batchTargetDir, "to-dir", "", "Output directory; one output per input is written here")
|
||||||
|
batchCmd.Flags().StringVar(&batchPackageName, "package", "", "Package name (for code generation formats like gorm/bun)")
|
||||||
|
batchCmd.Flags().StringVar(&batchSchemaFilter, "schema", "", "Filter to a specific schema by name")
|
||||||
|
batchCmd.Flags().BoolVar(&batchFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table")
|
||||||
|
batchCmd.Flags().StringVar(&batchNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm)")
|
||||||
|
batchCmd.Flags().StringVar(&batchNullableArrays, "array-nullable", "", "Nullable array representation for the Bun writer")
|
||||||
|
batchCmd.Flags().BoolVar(&batchContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL (pgsql output only)")
|
||||||
|
batchCmd.Flags().BoolVar(&batchKeepGoing, "keep-going", false, "Process remaining inputs after a failure; still exits non-zero")
|
||||||
|
batchCmd.Flags().BoolVar(&batchDryRun, "dry-run", false, "Read and validate every input and print the plan without writing any output")
|
||||||
|
|
||||||
|
for _, f := range []string{"from", "input", "to", "to-dir"} {
|
||||||
|
if err := batchCmd.MarkFlagRequired(f); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "Error marking %s flag as required: %v\n", f, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// batchDirTargets are writers that emit a directory of files rather than one file.
|
||||||
|
var batchDirTargets = map[string]bool{"gorm": true, "bun": true, "drizzle": true}
|
||||||
|
|
||||||
|
// batchExtensions maps single-file target formats to their output extension.
|
||||||
|
var batchExtensions = map[string]string{
|
||||||
|
"dbml": ".dbml", "dctx": ".dctx", "drawdb": ".ddb", "json": ".json",
|
||||||
|
"yaml": ".yaml", "yml": ".yaml", "pgsql": ".sql", "postgres": ".sql",
|
||||||
|
"postgresql": ".sql", "sql": ".sql", "mssql": ".sql", "sqlserver": ".sql",
|
||||||
|
"mssql2016": ".sql", "mssql2017": ".sql", "mssql2019": ".sql", "mssql2022": ".sql",
|
||||||
|
"sqlite": ".sql", "sqlite3": ".sql", "prisma": ".prisma", "typeorm": ".ts",
|
||||||
|
"graphql": ".graphql", "gql": ".graphql",
|
||||||
|
}
|
||||||
|
|
||||||
|
// expandBatchInputs resolves paths and glob patterns into a sorted,
|
||||||
|
// de-duplicated file list. A pattern that matches nothing is an error.
|
||||||
|
func expandBatchInputs(patterns []string) ([]string, error) {
|
||||||
|
seen := map[string]bool{}
|
||||||
|
var files []string
|
||||||
|
for _, p := range patterns {
|
||||||
|
p = strings.TrimSpace(p)
|
||||||
|
if p == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
matches, err := filepath.Glob(p)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid pattern %q: %w", p, err)
|
||||||
|
}
|
||||||
|
if len(matches) == 0 {
|
||||||
|
return nil, fmt.Errorf("no files match %q", p)
|
||||||
|
}
|
||||||
|
for _, m := range matches {
|
||||||
|
if !seen[m] {
|
||||||
|
seen[m] = true
|
||||||
|
files = append(files, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(files) == 0 {
|
||||||
|
return nil, fmt.Errorf("no input files given")
|
||||||
|
}
|
||||||
|
sort.Strings(files)
|
||||||
|
return files, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// batchOutputPaths returns the output path for each input. It errors when two
|
||||||
|
// inputs would collide on the same output name.
|
||||||
|
func batchOutputPaths(files []string, targetType, dir string) ([]string, error) {
|
||||||
|
key := strings.ToLower(targetType)
|
||||||
|
ext := ""
|
||||||
|
if !batchDirTargets[key] {
|
||||||
|
var ok bool
|
||||||
|
if ext, ok = batchExtensions[key]; !ok {
|
||||||
|
return nil, fmt.Errorf("unsupported target format: %s", targetType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
outs := make([]string, len(files))
|
||||||
|
owner := map[string]string{}
|
||||||
|
for i, f := range files {
|
||||||
|
stem := strings.TrimSuffix(filepath.Base(f), filepath.Ext(f))
|
||||||
|
out := filepath.Join(dir, stem+ext)
|
||||||
|
if prev, dup := owner[out]; dup {
|
||||||
|
return nil, fmt.Errorf("inputs %s and %s would both write %s", prev, f, out)
|
||||||
|
}
|
||||||
|
owner[out] = f
|
||||||
|
outs[i] = out
|
||||||
|
}
|
||||||
|
return outs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func runBatch(cmd *cobra.Command, args []string) error {
|
||||||
|
files, err := expandBatchInputs(batchInputs)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
outs, err := batchOutputPaths(files, batchTargetType, batchTargetDir)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== RelSpec Batch Converter ===\n")
|
||||||
|
fmt.Fprintf(os.Stderr, "Started at: %s\n", getCurrentTimestamp())
|
||||||
|
fmt.Fprintf(os.Stderr, "Inputs: %d file(s), %s -> %s\n\n", len(files), batchSourceType, batchTargetType)
|
||||||
|
|
||||||
|
out := outWriter(cmd)
|
||||||
|
if batchDryRun {
|
||||||
|
fmt.Fprintf(out, "RelSpec batch plan (dry run - nothing written):\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
var failed []string
|
||||||
|
for i, f := range files {
|
||||||
|
fmt.Fprintf(os.Stderr, "[%d/%d] %s -> %s\n", i+1, len(files), f, outs[i])
|
||||||
|
if err := processBatchItem(cmd, f, outs[i]); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, " ✗ %v\n", err)
|
||||||
|
failed = append(failed, fmt.Sprintf("%s: %v", f, err))
|
||||||
|
if !batchKeepGoing {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
fmt.Fprintf(os.Stderr, " ✓ done\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Batch Complete: %d ok, %d failed ===\n", len(files)-len(failed), len(failed))
|
||||||
|
if len(failed) > 0 {
|
||||||
|
return fmt.Errorf("batch finished with %d failure(s):\n %s", len(failed), strings.Join(failed, "\n "))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func processBatchItem(cmd *cobra.Command, in, outPath string) error {
|
||||||
|
db, err := readDatabaseForConvert(batchSourceType, in, "")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to read source: %w", err)
|
||||||
|
}
|
||||||
|
finalizeCommentedRefs(db, stderrWarn)
|
||||||
|
|
||||||
|
if batchDryRun {
|
||||||
|
if err := validateWriteTarget(db, batchTargetType, batchPackageName, batchSchemaFilter, ""); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
w := outWriter(cmd)
|
||||||
|
fmt.Fprintf(w, " %s -> %s (database '%s')\n", in, outPath, db.Name)
|
||||||
|
printDryRunPlan(w, db)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.MkdirAll(batchTargetDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("failed to create output directory: %w", err)
|
||||||
|
}
|
||||||
|
if err := writeDatabase(db, batchTargetType, outPath, batchPackageName, batchSchemaFilter, batchFlattenSchema, batchNullableTypes, batchNullableArrays, batchContinueOnError, ""); err != nil {
|
||||||
|
return fmt.Errorf("failed to write target: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func saveBatchState(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
a, b, c, d, e, f, g := batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun
|
||||||
|
t.Cleanup(func() {
|
||||||
|
batchSourceType, batchInputs, batchTargetType, batchTargetDir, batchPackageName, batchKeepGoing, batchDryRun = a, b, c, d, e, f, g
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_ConvertsEachInput(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "b.json"), []string{"posts"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchKeepGoing, batchDryRun = "", false, false
|
||||||
|
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||||
|
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
if err := runBatch(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("batch: %v", err)
|
||||||
|
}
|
||||||
|
for _, name := range []string{"a.yaml", "b.yaml"} {
|
||||||
|
if _, err := os.Stat(filepath.Join(outDir, name)); err != nil {
|
||||||
|
t.Errorf("expected %s: %v", name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_DryRunWritesNothing(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchKeepGoing, batchDryRun = "", false, true
|
||||||
|
batchInputs = []string{filepath.Join(dir, "a.json")}
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runBatch(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(outDir); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output directory")
|
||||||
|
}
|
||||||
|
if !strings.Contains(buf.String(), "users") || !strings.Contains(buf.String(), "a.yaml") {
|
||||||
|
t.Errorf("plan incomplete:\n%s", buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunBatch_FailureHandling(t *testing.T) {
|
||||||
|
saveBatchState(t)
|
||||||
|
dir := t.TempDir()
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "a.json"), []string{"users"})
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "b.json"), []byte("{not json"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
writeTestJSON(t, filepath.Join(dir, "c.json"), []string{"posts"})
|
||||||
|
outDir := filepath.Join(dir, "out")
|
||||||
|
|
||||||
|
batchSourceType, batchTargetType, batchTargetDir = "json", "yaml", outDir
|
||||||
|
batchPackageName, batchDryRun = "", false
|
||||||
|
batchInputs = []string{filepath.Join(dir, "*.json")}
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
|
||||||
|
// Default: stop at first failure.
|
||||||
|
batchKeepGoing = false
|
||||||
|
err := runBatch(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "b.json") {
|
||||||
|
t.Fatalf("expected failure naming b.json, got %v", err)
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); !os.IsNotExist(statErr) {
|
||||||
|
t.Error("c.json should not be processed without --keep-going")
|
||||||
|
}
|
||||||
|
|
||||||
|
// --keep-going: remaining inputs are processed, exit still fails.
|
||||||
|
batchKeepGoing = true
|
||||||
|
if err := runBatch(cmd, nil); err == nil {
|
||||||
|
t.Fatal("expected non-zero result with --keep-going")
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(filepath.Join(outDir, "c.yaml")); statErr != nil {
|
||||||
|
t.Errorf("c.yaml should be written with --keep-going: %v", statErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExpandBatchInputs(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
for _, n := range []string{"b.json", "a.json"} {
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, n), []byte("{}"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
got, err := expandBatchInputs([]string{filepath.Join(dir, "*.json"), filepath.Join(dir, "a.json")})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(got) != 2 || filepath.Base(got[0]) != "a.json" || filepath.Base(got[1]) != "b.json" {
|
||||||
|
t.Errorf("want sorted deduped [a b], got %v", got)
|
||||||
|
}
|
||||||
|
if _, err := expandBatchInputs([]string{filepath.Join(dir, "*.nope")}); err == nil {
|
||||||
|
t.Error("unmatched pattern should error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchOutputPaths(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
files []string
|
||||||
|
target string
|
||||||
|
want []string
|
||||||
|
wantErr string
|
||||||
|
}{
|
||||||
|
{"file target", []string{"x/a.dbml"}, "json", []string{"out/a.json"}, ""},
|
||||||
|
{"dir target", []string{"x/a.json"}, "gorm", []string{"out/a"}, ""},
|
||||||
|
{"collision", []string{"x/a.json", "y/a.json"}, "yaml", nil, "both write"},
|
||||||
|
{"unsupported", []string{"a.json"}, "nope", nil, "unsupported target"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := batchOutputPaths(tt.files, tt.target, "out")
|
||||||
|
if tt.wantErr != "" {
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||||
|
t.Fatalf("want error %q, got %v", tt.wantErr, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil || len(got) != len(tt.want) || got[0] != filepath.FromSlash(tt.want[0]) {
|
||||||
|
t.Fatalf("got %v, %v; want %v", got, err, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+140
-15
@@ -3,6 +3,7 @@ package main
|
|||||||
import (
|
import (
|
||||||
stdjson "encoding/json"
|
stdjson "encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -21,6 +22,7 @@ import (
|
|||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/graphql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/graphql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/json"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/mssql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/mssql"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/mysql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/pgsql"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/prisma"
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
"git.warky.dev/wdevs/relspecgo/pkg/readers/sqlite"
|
||||||
@@ -36,6 +38,7 @@ import (
|
|||||||
wgraphql "git.warky.dev/wdevs/relspecgo/pkg/writers/graphql"
|
wgraphql "git.warky.dev/wdevs/relspecgo/pkg/writers/graphql"
|
||||||
wjson "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
|
wjson "git.warky.dev/wdevs/relspecgo/pkg/writers/json"
|
||||||
wmssql "git.warky.dev/wdevs/relspecgo/pkg/writers/mssql"
|
wmssql "git.warky.dev/wdevs/relspecgo/pkg/writers/mssql"
|
||||||
|
wmysql "git.warky.dev/wdevs/relspecgo/pkg/writers/mysql"
|
||||||
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
|
wpgsql "git.warky.dev/wdevs/relspecgo/pkg/writers/pgsql"
|
||||||
wprisma "git.warky.dev/wdevs/relspecgo/pkg/writers/prisma"
|
wprisma "git.warky.dev/wdevs/relspecgo/pkg/writers/prisma"
|
||||||
wsqlite "git.warky.dev/wdevs/relspecgo/pkg/writers/sqlite"
|
wsqlite "git.warky.dev/wdevs/relspecgo/pkg/writers/sqlite"
|
||||||
@@ -57,6 +60,9 @@ var (
|
|||||||
convertNullableArrays string
|
convertNullableArrays string
|
||||||
convertContinueOnError bool
|
convertContinueOnError bool
|
||||||
convertExtraFields string
|
convertExtraFields string
|
||||||
|
convertDryRun bool
|
||||||
|
convertWatch bool
|
||||||
|
convertWatchInterval time.Duration
|
||||||
)
|
)
|
||||||
|
|
||||||
var convertCmd = &cobra.Command{
|
var convertCmd = &cobra.Command{
|
||||||
@@ -165,7 +171,11 @@ Examples:
|
|||||||
|
|
||||||
# Convert SQLite to PostgreSQL SQL
|
# Convert SQLite to PostgreSQL SQL
|
||||||
relspec convert --from sqlite --from-path database.db \
|
relspec convert --from sqlite --from-path database.db \
|
||||||
--to pgsql --to-path schema.sql`,
|
--to pgsql --to-path schema.sql
|
||||||
|
|
||||||
|
# Regenerate GORM models every time the DBML file changes
|
||||||
|
relspec convert --from dbml --from-path schema.dbml \
|
||||||
|
--to gorm --to-path models/ --package models --watch`,
|
||||||
RunE: runConvert,
|
RunE: runConvert,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -185,6 +195,11 @@ func init() {
|
|||||||
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
convertCmd.Flags().BoolVar(&convertContinueOnError, "continue-on-error", false, "Prepend \\set ON_ERROR_STOP off to generated SQL so psql continues past errors (pgsql output only)")
|
||||||
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
convertCmd.Flags().StringVar(&convertExtraFields, "extra-fields", "", "Path to JSON file containing extra Bun model fields to inject (bun output only); fields support target_table, name, type, bun_tag, json_tag, comment")
|
||||||
|
|
||||||
|
convertCmd.Flags().BoolVar(&convertDryRun, "dry-run", false, "Read and validate the input and print the plan without writing any output")
|
||||||
|
|
||||||
|
convertCmd.Flags().BoolVar(&convertWatch, "watch", false, "Watch the source files (--from-path or --from-list) and regenerate the output whenever they change")
|
||||||
|
convertCmd.Flags().DurationVar(&convertWatchInterval, "watch-interval", 500*time.Millisecond, "Polling interval used by --watch")
|
||||||
|
|
||||||
err := convertCmd.MarkFlagRequired("from")
|
err := convertCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
||||||
@@ -200,6 +215,13 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func runConvert(cmd *cobra.Command, args []string) error {
|
func runConvert(cmd *cobra.Command, args []string) error {
|
||||||
|
if convertWatch {
|
||||||
|
return runConvertWatch(cmd.Context(), os.Stderr, func() error { return runConvertOnce(cmd) })
|
||||||
|
}
|
||||||
|
return runConvertOnce(cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runConvertOnce(cmd *cobra.Command) error {
|
||||||
fmt.Fprintf(os.Stderr, "\n=== RelSpec Schema Converter ===\n")
|
fmt.Fprintf(os.Stderr, "\n=== RelSpec Schema Converter ===\n")
|
||||||
fmt.Fprintf(os.Stderr, "Started at: %s\n\n", getCurrentTimestamp())
|
fmt.Fprintf(os.Stderr, "Started at: %s\n\n", getCurrentTimestamp())
|
||||||
|
|
||||||
@@ -240,6 +262,22 @@ func runConvert(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
fmt.Fprintf(os.Stderr, " Found: %d table(s)\n\n", totalTables)
|
fmt.Fprintf(os.Stderr, " Found: %d table(s)\n\n", totalTables)
|
||||||
|
|
||||||
|
if convertDryRun {
|
||||||
|
if err := validateWriteTarget(db, convertTargetType, convertPackageName, convertSchemaFilter, convertExtraFields); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec convert plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Input: %s database '%s'\n", convertSourceType, db.Name)
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", convertTargetType, convertTargetPath)
|
||||||
|
if convertSchemaFilter != "" {
|
||||||
|
fmt.Fprintf(out, " Schema filter: %s\n", convertSchemaFilter)
|
||||||
|
}
|
||||||
|
printDryRunPlan(out, db)
|
||||||
|
fmt.Fprintf(os.Stderr, "=== Dry run complete: no output written ===\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Write to target format
|
// Write to target format
|
||||||
fmt.Fprintf(os.Stderr, "[2/2] Writing to target format...\n")
|
fmt.Fprintf(os.Stderr, "[2/2] Writing to target format...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", convertTargetType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", convertTargetType)
|
||||||
@@ -380,6 +418,12 @@ func readDatabaseForConvert(dbType, filePath, connString string) (*models.Databa
|
|||||||
}
|
}
|
||||||
reader = mssql.NewReader(newReaderOptions("", connString))
|
reader = mssql.NewReader(newReaderOptions("", connString))
|
||||||
|
|
||||||
|
case "mysql", "mariadb":
|
||||||
|
if connString == "" {
|
||||||
|
return nil, fmt.Errorf("connection string is required for MySQL format")
|
||||||
|
}
|
||||||
|
reader = mysql.NewReader(newReaderOptions("", connString))
|
||||||
|
|
||||||
case "sqlite", "sqlite3":
|
case "sqlite", "sqlite3":
|
||||||
// SQLite can use either file path or connection string
|
// SQLite can use either file path or connection string
|
||||||
dbPath := filePath
|
dbPath := filePath
|
||||||
@@ -408,23 +452,12 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
|
|
||||||
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
writerOpts := newWriterOptions(outputPath, packageName, flattenSchema, nullableTypes, nullableArrays, continueOnError)
|
||||||
if extraFields != "" {
|
if extraFields != "" {
|
||||||
if !strings.EqualFold(dbType, "bun") {
|
extraFieldsJSON, err := loadExtraFields(dbType, extraFields)
|
||||||
return fmt.Errorf("--extra-fields is only supported for Bun output")
|
|
||||||
}
|
|
||||||
extraFieldsJSON, err := os.ReadFile(extraFields)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to read --extra-fields file %q: %w", extraFields, err)
|
return err
|
||||||
}
|
|
||||||
|
|
||||||
var parsed []wbun.ExtraFieldConfig
|
|
||||||
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
|
||||||
return fmt.Errorf("invalid --extra-fields JSON in %q: %w", extraFields, err)
|
|
||||||
}
|
|
||||||
if len(parsed) == 0 {
|
|
||||||
return fmt.Errorf("--extra-fields must contain at least one field")
|
|
||||||
}
|
}
|
||||||
writerOpts.Metadata = map[string]interface{}{
|
writerOpts.Metadata = map[string]interface{}{
|
||||||
"extra_fields": string(extraFieldsJSON),
|
"extra_fields": extraFieldsJSON,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -465,6 +498,9 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
case "mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022":
|
case "mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022":
|
||||||
writer = wmssql.NewWriter(writerOpts)
|
writer = wmssql.NewWriter(writerOpts)
|
||||||
|
|
||||||
|
case "mysql", "mariadb":
|
||||||
|
writer = wmysql.NewWriter(writerOpts)
|
||||||
|
|
||||||
case "sqlite", "sqlite3":
|
case "sqlite", "sqlite3":
|
||||||
writer = wsqlite.NewWriter(writerOpts)
|
writer = wsqlite.NewWriter(writerOpts)
|
||||||
|
|
||||||
@@ -529,6 +565,95 @@ func writeDatabase(db *models.Database, dbType, outputPath, packageName, schemaF
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// loadExtraFields reads and validates the --extra-fields JSON file, returning
|
||||||
|
// its raw content.
|
||||||
|
func loadExtraFields(dbType, path string) (string, error) {
|
||||||
|
if !strings.EqualFold(dbType, "bun") {
|
||||||
|
return "", fmt.Errorf("--extra-fields is only supported for Bun output")
|
||||||
|
}
|
||||||
|
extraFieldsJSON, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to read --extra-fields file %q: %w", path, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var parsed []wbun.ExtraFieldConfig
|
||||||
|
if err := stdjson.Unmarshal(extraFieldsJSON, &parsed); err != nil {
|
||||||
|
return "", fmt.Errorf("invalid --extra-fields JSON in %q: %w", path, err)
|
||||||
|
}
|
||||||
|
if len(parsed) == 0 {
|
||||||
|
return "", fmt.Errorf("--extra-fields must contain at least one field")
|
||||||
|
}
|
||||||
|
return string(extraFieldsJSON), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateWriteTarget performs the checks writeDatabase makes before writing,
|
||||||
|
// without constructing a writer or touching the output path. Used by --dry-run.
|
||||||
|
func validateWriteTarget(db *models.Database, dbType, packageName, schemaFilter, extraFields string) error {
|
||||||
|
if extraFields != "" {
|
||||||
|
if _, err := loadExtraFields(dbType, extraFields); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch strings.ToLower(dbType) {
|
||||||
|
case "dbml", "dctx", "drawdb", "json", "yaml", "yml", "drizzle",
|
||||||
|
"pgsql", "postgres", "postgresql", "sql",
|
||||||
|
"mssql", "sqlserver", "mssql2016", "mssql2017", "mssql2019", "mssql2022",
|
||||||
|
"sqlite", "sqlite3", "prisma", "typeorm", "graphql", "gql":
|
||||||
|
case "gorm", "bun":
|
||||||
|
if packageName == "" {
|
||||||
|
return fmt.Errorf("package name is required for %s format (use --package flag)", dbType)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported target format: %s", dbType)
|
||||||
|
}
|
||||||
|
|
||||||
|
if schemaFilter != "" {
|
||||||
|
for _, schema := range db.Schemas {
|
||||||
|
if schema.Name == schemaFilter {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Errorf("schema '%s' not found in database. Available schemas: %v",
|
||||||
|
schemaFilter, getSchemaNames(db))
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.EqualFold(dbType, "dctx") {
|
||||||
|
if len(db.Schemas) == 0 {
|
||||||
|
return fmt.Errorf("no schemas found in database")
|
||||||
|
}
|
||||||
|
if len(db.Schemas) > 1 {
|
||||||
|
return fmt.Errorf("multiple schemas found, please specify which schema to export using --schema flag. Available schemas: %v",
|
||||||
|
getSchemaNames(db))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// outWriter returns the command's stdout, falling back to os.Stdout when the
|
||||||
|
// command is nil (tests call the run functions directly).
|
||||||
|
func outWriter(cmd *cobra.Command) io.Writer {
|
||||||
|
if cmd == nil {
|
||||||
|
return os.Stdout
|
||||||
|
}
|
||||||
|
return cmd.OutOrStdout()
|
||||||
|
}
|
||||||
|
|
||||||
|
// printDryRunPlan prints the schemas and tables that would be written.
|
||||||
|
func printDryRunPlan(out io.Writer, db *models.Database) {
|
||||||
|
for _, schema := range db.Schemas {
|
||||||
|
names := make([]string, 0, len(schema.Tables))
|
||||||
|
for _, t := range schema.Tables {
|
||||||
|
names = append(names, t.Name)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(out, " schema %q: %d table(s)", schema.Name, len(schema.Tables))
|
||||||
|
if len(names) > 0 {
|
||||||
|
fmt.Fprintf(out, " [%s]", strings.Join(names, ", "))
|
||||||
|
}
|
||||||
|
fmt.Fprintln(out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// getSchemaNames returns a slice of schema names from a database
|
// getSchemaNames returns a slice of schema names from a database
|
||||||
func getSchemaNames(db *models.Database) []string {
|
func getSchemaNames(db *models.Database) []string {
|
||||||
names := make([]string, len(db.Schemas))
|
names := make([]string, len(db.Schemas))
|
||||||
|
|||||||
@@ -0,0 +1,158 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newDryRunCmd() (*cobra.Command, *bytes.Buffer) {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
cmd := &cobra.Command{}
|
||||||
|
cmd.SetOut(&buf)
|
||||||
|
return cmd, &buf
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvert_DryRunWritesNothing(t *testing.T) {
|
||||||
|
defer func(a, b, c, d string, e bool) {
|
||||||
|
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||||
|
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "out.json")
|
||||||
|
writeTestJSON(t, in, []string{"users", "posts"})
|
||||||
|
|
||||||
|
convertSourceType, convertSourcePath = "json", in
|
||||||
|
convertTargetType, convertTargetPath = "json", out
|
||||||
|
convertDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runConvert(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||||
|
if !strings.Contains(buf.String(), want) {
|
||||||
|
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Normal behavior is unchanged.
|
||||||
|
convertDryRun = false
|
||||||
|
if err := runConvert(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("real run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); err != nil {
|
||||||
|
t.Fatalf("real run should write output: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvert_DryRunValidatesTarget(t *testing.T) {
|
||||||
|
defer func(a, b, c, d string, e bool) {
|
||||||
|
convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun = a, b, c, d, e
|
||||||
|
}(convertSourceType, convertSourcePath, convertTargetType, convertTargetPath, convertDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "models")
|
||||||
|
writeTestJSON(t, in, []string{"users"})
|
||||||
|
|
||||||
|
convertSourceType, convertSourcePath = "json", in
|
||||||
|
convertDryRun = true
|
||||||
|
|
||||||
|
// gorm without --package must fail validation, as a real run would.
|
||||||
|
convertTargetType, convertTargetPath = "gorm", out
|
||||||
|
cmd, _ := newDryRunCmd()
|
||||||
|
err := runConvert(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "package name is required") {
|
||||||
|
t.Fatalf("expected package validation error, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
convertTargetType = "nope"
|
||||||
|
err = runConvert(cmd, nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "unsupported target format") {
|
||||||
|
t.Fatalf("expected unsupported format error, got %v", err)
|
||||||
|
}
|
||||||
|
if _, statErr := os.Stat(out); !os.IsNotExist(statErr) {
|
||||||
|
t.Fatal("dry run must not create the output path")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunSplit_DryRunWritesNothing(t *testing.T) {
|
||||||
|
defer func(a, b, c, d, e string, f bool) {
|
||||||
|
splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun = a, b, c, d, e, f
|
||||||
|
}(splitSourceType, splitSourcePath, splitTargetType, splitTargetPath, splitTables, splitDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
in := filepath.Join(dir, "in.json")
|
||||||
|
out := filepath.Join(dir, "subset.json")
|
||||||
|
writeTestJSON(t, in, []string{"users", "posts", "comments"})
|
||||||
|
|
||||||
|
splitSourceType, splitSourcePath = "json", in
|
||||||
|
splitTargetType, splitTargetPath = "json", out
|
||||||
|
splitTables = "users,posts"
|
||||||
|
splitDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runSplit(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
got := buf.String()
|
||||||
|
if !strings.Contains(got, "2 table(s)") || strings.Contains(got, "comments") {
|
||||||
|
t.Errorf("plan should show only the 2 selected tables:\n%s", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A selection that matches nothing fails validation in dry-run too.
|
||||||
|
splitTables = "does_not_exist"
|
||||||
|
if err := runSplit(cmd, nil); err == nil {
|
||||||
|
t.Fatal("expected error for empty selection")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunMerge_DryRunWritesNothing(t *testing.T) {
|
||||||
|
saved := saveMergeState()
|
||||||
|
defer restoreMergeState(saved)
|
||||||
|
defer func(v bool) { mergeDryRun = v }(mergeDryRun)
|
||||||
|
|
||||||
|
dir := t.TempDir()
|
||||||
|
target := filepath.Join(dir, "target.json")
|
||||||
|
source := filepath.Join(dir, "source.json")
|
||||||
|
out := filepath.Join(dir, "merged.json")
|
||||||
|
writeTestJSON(t, target, []string{"users"})
|
||||||
|
writeTestJSON(t, source, []string{"posts"})
|
||||||
|
|
||||||
|
mergeTargetType, mergeTargetPath, mergeTargetConn = "json", target, ""
|
||||||
|
mergeSourceType, mergeSourcePath, mergeSourceConn = "json", source, ""
|
||||||
|
mergeFromList = nil
|
||||||
|
mergeOutputType, mergeOutputPath, mergeOutputConn = "json", out, ""
|
||||||
|
mergeSkipTables, mergeReportPath = "", ""
|
||||||
|
mergeDryRun = true
|
||||||
|
|
||||||
|
cmd, buf := newDryRunCmd()
|
||||||
|
if err := runMerge(cmd, nil); err != nil {
|
||||||
|
t.Fatalf("dry run: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := os.Stat(out); !os.IsNotExist(err) {
|
||||||
|
t.Fatal("dry run must not create the output file")
|
||||||
|
}
|
||||||
|
for _, want := range []string{"dry run", "users", "posts", out} {
|
||||||
|
if !strings.Contains(buf.String(), want) {
|
||||||
|
t.Errorf("plan missing %q:\n%s", want, buf.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
mergeOutputType = "nope"
|
||||||
|
if err := runMerge(cmd, nil); err == nil || !strings.Contains(err.Error(), "unsupported format") {
|
||||||
|
t.Fatalf("expected unsupported output format error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -61,6 +61,7 @@ var (
|
|||||||
mergeReportPath string // Path to write merge report
|
mergeReportPath string // Path to write merge report
|
||||||
mergeFullDDL bool // Execute full DDL instead of diffing the live pgsql database
|
mergeFullDDL bool // Execute full DDL instead of diffing the live pgsql database
|
||||||
mergeFlattenSchema bool
|
mergeFlattenSchema bool
|
||||||
|
mergeDryRun bool
|
||||||
)
|
)
|
||||||
|
|
||||||
var mergeCmd = &cobra.Command{
|
var mergeCmd = &cobra.Command{
|
||||||
@@ -130,6 +131,7 @@ func init() {
|
|||||||
mergeCmd.Flags().BoolVar(&mergeVerbose, "verbose", false, "Show verbose output")
|
mergeCmd.Flags().BoolVar(&mergeVerbose, "verbose", false, "Show verbose output")
|
||||||
mergeCmd.Flags().BoolVar(&mergeFullDDL, "full-ddl", false, "pgsql database output: execute the full DDL instead of only the differences from the live database")
|
mergeCmd.Flags().BoolVar(&mergeFullDDL, "full-ddl", false, "pgsql database output: execute the full DDL instead of only the differences from the live database")
|
||||||
mergeCmd.Flags().StringVar(&mergeReportPath, "merge-report", "", "Path to write merge report (JSON format)")
|
mergeCmd.Flags().StringVar(&mergeReportPath, "merge-report", "", "Path to write merge report (JSON format)")
|
||||||
|
mergeCmd.Flags().BoolVar(&mergeDryRun, "dry-run", false, "Read and merge in memory, then print the merge plan without writing any output")
|
||||||
mergeCmd.Flags().BoolVar(&mergeFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
mergeCmd.Flags().BoolVar(&mergeFlattenSchema, "flatten-schema", false, "Flatten schema.table names to schema_table (useful for databases like SQLite that do not support schemas)")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -264,6 +266,29 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
|||||||
merge.GetColumnTypeConflictSummary(result, 10))
|
merge.GetColumnTypeConflictSummary(result, 10))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if mergeDryRun {
|
||||||
|
if !isMergeOutputFormat(mergeOutputType) {
|
||||||
|
return fmt.Errorf("dry run validation failed: Output: unsupported format '%s'", mergeOutputType)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec merge plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Target: %s database '%s'\n", mergeTargetType, targetDB.Name)
|
||||||
|
fmt.Fprintf(out, " Source: %s database '%s'\n", mergeSourceType, sourceDB.Name)
|
||||||
|
switch {
|
||||||
|
case mergeOutputPath != "":
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", mergeOutputType, mergeOutputPath)
|
||||||
|
case mergeOutputConn != "":
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", mergeOutputType, maskPassword(mergeOutputConn))
|
||||||
|
default:
|
||||||
|
fmt.Fprintf(out, " Output: %s\n", mergeOutputType)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(out, " Result:\n")
|
||||||
|
printDryRunPlan(out, targetDB)
|
||||||
|
fmt.Fprintf(out, "\n%s\n", merge.GetMergeSummary(result))
|
||||||
|
fmt.Fprintf(os.Stderr, "\n=== Dry run complete: no output written ===\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Step 4: Write output
|
// Step 4: Write output
|
||||||
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
|
fmt.Fprintf(os.Stderr, "\n[4/4] Writing output...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", mergeOutputType)
|
||||||
@@ -285,6 +310,16 @@ func runMerge(cmd *cobra.Command, args []string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isMergeOutputFormat reports whether writeDatabaseForMerge supports dbType.
|
||||||
|
func isMergeOutputFormat(dbType string) bool {
|
||||||
|
switch strings.ToLower(dbType) {
|
||||||
|
case "dbml", "dctx", "drawdb", "graphql", "json", "yaml", "gorm", "bun",
|
||||||
|
"drizzle", "prisma", "typeorm", "sqlite", "sqlite3", "pgsql":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
func readDatabaseForMerge(dbType, filePath, connString, label string) (*models.Database, error) {
|
func readDatabaseForMerge(dbType, filePath, connString, label string) (*models.Database, error) {
|
||||||
var reader readers.Reader
|
var reader readers.Reader
|
||||||
|
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ func newWriterOptions(outputPath, packageName string, flattenSchema bool, nullab
|
|||||||
FlattenSchema: flattenSchema,
|
FlattenSchema: flattenSchema,
|
||||||
NullableTypes: nullableTypes,
|
NullableTypes: nullableTypes,
|
||||||
NullableArrays: nullableArrays,
|
NullableArrays: nullableArrays,
|
||||||
|
TypeMappings: typeMappings,
|
||||||
Prisma7: prisma7,
|
Prisma7: prisma7,
|
||||||
ContinueOnError: continueOnError,
|
ContinueOnError: continueOnError,
|
||||||
StrictDirectives: strictDirectives,
|
StrictDirectives: strictDirectives,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
"git.warky.dev/wdevs/relspecgo/pkg/buildinfo"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
)
|
)
|
||||||
|
|
||||||
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
// version/buildDate mirror pkg/buildinfo so existing call sites keep working.
|
||||||
@@ -17,6 +18,8 @@ var (
|
|||||||
noVersion bool
|
noVersion bool
|
||||||
silent bool
|
silent bool
|
||||||
strictDirectives bool
|
strictDirectives bool
|
||||||
|
typeMapFlags []string
|
||||||
|
typeMappings map[string]string
|
||||||
)
|
)
|
||||||
|
|
||||||
var rootCmd = &cobra.Command{
|
var rootCmd = &cobra.Command{
|
||||||
@@ -28,10 +31,16 @@ bidirectional conversion between various database schema formats.
|
|||||||
It reads database schemas from multiple sources (live databases, DBML,
|
It reads database schemas from multiple sources (live databases, DBML,
|
||||||
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
DCTX, DrawDB, etc.) and writes them to various formats (GORM, Bun,
|
||||||
JSON, YAML, SQL, etc.).`,
|
JSON, YAML, SQL, etc.).`,
|
||||||
|
PersistentPreRunE: func(cmd *cobra.Command, args []string) error {
|
||||||
|
var err error
|
||||||
|
typeMappings, err = writers.ParseTypeMappings(typeMapFlags)
|
||||||
|
return err
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
rootCmd.AddCommand(convertCmd)
|
rootCmd.AddCommand(convertCmd)
|
||||||
|
rootCmd.AddCommand(batchCmd)
|
||||||
rootCmd.AddCommand(diffCmd)
|
rootCmd.AddCommand(diffCmd)
|
||||||
rootCmd.AddCommand(inspectCmd)
|
rootCmd.AddCommand(inspectCmd)
|
||||||
rootCmd.AddCommand(scriptsCmd)
|
rootCmd.AddCommand(scriptsCmd)
|
||||||
@@ -44,6 +53,7 @@ func init() {
|
|||||||
rootCmd.AddCommand(versionCmd)
|
rootCmd.AddCommand(versionCmd)
|
||||||
rootCmd.AddCommand(reportCmd)
|
rootCmd.AddCommand(reportCmd)
|
||||||
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
rootCmd.PersistentFlags().BoolVar(&prisma7, "prisma7", false, "Use Prisma 7 generator conventions when reading/writing Prisma schemas")
|
||||||
|
rootCmd.PersistentFlags().StringArrayVar(&typeMapFlags, "type-map", nil, "Override a SQL-to-Go type mapping for bun/gorm output as sqltype=gotype (repeatable), e.g. --type-map uuid=uuid.UUID --type-map numeric=decimal.Decimal")
|
||||||
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
rootCmd.PersistentFlags().BoolVar(&noVersion, "no-version", false, "Suppress the RelSpec version header")
|
||||||
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
rootCmd.PersistentFlags().BoolVar(&silent, "silent", false, "Suppress progress and status messages (errors are still shown)")
|
||||||
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
rootCmd.PersistentFlags().BoolVar(&strictDirectives, "strict-directives", false, "Fail on unknown or untranslatable DBML dialect directives (@postgres:, @sqlite:, …)")
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ var (
|
|||||||
splitExcludeTables string
|
splitExcludeTables string
|
||||||
splitNullableTypes string
|
splitNullableTypes string
|
||||||
splitNullableArrays string
|
splitNullableArrays string
|
||||||
|
splitDryRun bool
|
||||||
)
|
)
|
||||||
|
|
||||||
var splitCmd = &cobra.Command{
|
var splitCmd = &cobra.Command{
|
||||||
@@ -115,6 +116,8 @@ func init() {
|
|||||||
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
splitCmd.Flags().StringVar(&splitNullableTypes, "types", "", "Nullable type package for code-gen writers (bun/gorm): 'baselib' (default, Go pointer types), 'stdlib' (database/sql), or 'sqltypes'")
|
||||||
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
splitCmd.Flags().StringVar(&splitNullableArrays, "array-nullable", "", "Nullable PostgreSQL array representation for the Bun writer in stdlib/baselib --types mode: 'slice' (default, plain slice) or 'pointer_slice' (*[]T, distinguishes NULL from '{}')")
|
||||||
|
|
||||||
|
splitCmd.Flags().BoolVar(&splitDryRun, "dry-run", false, "Read, filter and validate the selection and print the plan without writing any output")
|
||||||
|
|
||||||
err := splitCmd.MarkFlagRequired("from")
|
err := splitCmd.MarkFlagRequired("from")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
fmt.Fprintf(os.Stderr, "Error marking from flag as required: %v\n", err)
|
||||||
@@ -174,6 +177,26 @@ func runSplit(cmd *cobra.Command, args []string) error {
|
|||||||
}
|
}
|
||||||
fmt.Fprintf(os.Stderr, " ✓ Filtered to: %d schema(s), %d table(s)\n\n", len(filteredDB.Schemas), filteredTables)
|
fmt.Fprintf(os.Stderr, " ✓ Filtered to: %d schema(s), %d table(s)\n\n", len(filteredDB.Schemas), filteredTables)
|
||||||
|
|
||||||
|
if splitDryRun {
|
||||||
|
if err := validateWriteTarget(filteredDB, splitTargetType, splitPackageName, "", ""); err != nil {
|
||||||
|
return fmt.Errorf("dry run validation failed: %w", err)
|
||||||
|
}
|
||||||
|
out := outWriter(cmd)
|
||||||
|
fmt.Fprintf(out, "RelSpec split plan (dry run - nothing written):\n")
|
||||||
|
fmt.Fprintf(out, " Input: %s database '%s'\n", splitSourceType, db.Name)
|
||||||
|
fmt.Fprintf(out, " Output: %s -> %s\n", splitTargetType, splitTargetPath)
|
||||||
|
fmt.Fprintf(out, " Selection: %s\n", splitSelection{
|
||||||
|
Schemas: parseCommaSeparated(splitSchemas),
|
||||||
|
Tables: parseCommaSeparated(splitTables),
|
||||||
|
ExcludeSchemas: parseCommaSeparated(splitExcludeSchema),
|
||||||
|
ExcludeTables: parseCommaSeparated(splitExcludeTables),
|
||||||
|
DatabaseName: splitDatabaseName,
|
||||||
|
}.summary())
|
||||||
|
printDryRunPlan(out, filteredDB)
|
||||||
|
fmt.Fprintf(os.Stderr, "=== Dry run complete: no output written ===\n\n")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Write to target format
|
// Write to target format
|
||||||
fmt.Fprintf(os.Stderr, "[3/3] Writing to target format...\n")
|
fmt.Fprintf(os.Stderr, "[3/3] Writing to target format...\n")
|
||||||
fmt.Fprintf(os.Stderr, " Format: %s\n", splitTargetType)
|
fmt.Fprintf(os.Stderr, " Format: %s\n", splitTargetType)
|
||||||
|
|||||||
@@ -0,0 +1,152 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"io/fs"
|
||||||
|
"os"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// watchSnapshot maps a file path to its modification time and size.
|
||||||
|
type watchSnapshot map[string]string
|
||||||
|
|
||||||
|
// takeWatchSnapshot records the state of every file under the given paths.
|
||||||
|
// Directories are walked recursively. Anything at or below the excluded path
|
||||||
|
// (typically the output path) is skipped so regenerating output does not
|
||||||
|
// retrigger the watcher. Missing paths are simply absent from the snapshot, so
|
||||||
|
// creating them later counts as a change.
|
||||||
|
func takeWatchSnapshot(paths []string, exclude string) watchSnapshot {
|
||||||
|
snap := watchSnapshot{}
|
||||||
|
exclude = absPathOrSelf(exclude)
|
||||||
|
for _, root := range paths {
|
||||||
|
_ = filepath.WalkDir(root, func(p string, d fs.DirEntry, err error) error {
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if exclude != "" && isWithin(absPathOrSelf(p), exclude) {
|
||||||
|
if d.IsDir() {
|
||||||
|
return filepath.SkipDir
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if d.IsDir() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, err := d.Info()
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
snap[p] = fmt.Sprintf("%d-%d", info.ModTime().UnixNano(), info.Size())
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return snap
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s watchSnapshot) equal(o watchSnapshot) bool {
|
||||||
|
if len(s) != len(o) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for k, v := range s {
|
||||||
|
if o[k] != v {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func absPathOrSelf(p string) string {
|
||||||
|
if p == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if abs, err := filepath.Abs(p); err == nil {
|
||||||
|
return abs
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// isWithin reports whether path equals dir or is located below it.
|
||||||
|
func isWithin(path, dir string) bool {
|
||||||
|
if path == dir {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return strings.HasPrefix(path, dir+string(filepath.Separator))
|
||||||
|
}
|
||||||
|
|
||||||
|
// watchLoop runs fn once immediately and again whenever the watched paths
|
||||||
|
// change, until ctx is cancelled. Changes are debounced: fn runs only after
|
||||||
|
// the snapshot has stayed unchanged for one poll interval. Errors from fn are
|
||||||
|
// reported to w and do not stop the loop.
|
||||||
|
func watchLoop(ctx context.Context, w io.Writer, paths []string, exclude string, interval time.Duration, fn func() error) {
|
||||||
|
run := func() {
|
||||||
|
if err := fn(); err != nil {
|
||||||
|
fmt.Fprintf(w, "Error: %v\n", err)
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, "Watching for changes (Ctrl-C to stop)...\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
last := takeWatchSnapshot(paths, exclude)
|
||||||
|
run()
|
||||||
|
|
||||||
|
ticker := time.NewTicker(interval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
|
cur := takeWatchSnapshot(paths, exclude)
|
||||||
|
if cur.equal(last) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Debounce: wait until writes settle.
|
||||||
|
for settled := false; !settled; {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-time.After(interval):
|
||||||
|
}
|
||||||
|
next := takeWatchSnapshot(paths, exclude)
|
||||||
|
settled = next.equal(cur)
|
||||||
|
cur = next
|
||||||
|
}
|
||||||
|
fmt.Fprintf(w, "\nChange detected, regenerating...\n")
|
||||||
|
last = cur
|
||||||
|
run()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// runConvertWatch runs the conversion once and then again whenever the source
|
||||||
|
// files change, until interrupted.
|
||||||
|
func runConvertWatch(parent context.Context, w io.Writer, run func() error) error {
|
||||||
|
var paths []string
|
||||||
|
switch {
|
||||||
|
case len(convertFromList) > 0:
|
||||||
|
paths = convertFromList
|
||||||
|
case convertSourcePath != "":
|
||||||
|
paths = []string{convertSourcePath}
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("--watch requires --from-path or --from-list (live database connections cannot be watched)")
|
||||||
|
}
|
||||||
|
if convertDryRun {
|
||||||
|
return fmt.Errorf("--watch cannot be combined with --dry-run")
|
||||||
|
}
|
||||||
|
if convertWatchInterval <= 0 {
|
||||||
|
return fmt.Errorf("--watch-interval must be positive")
|
||||||
|
}
|
||||||
|
if parent == nil {
|
||||||
|
parent = context.Background()
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, stop := signal.NotifyContext(parent, os.Interrupt, syscall.SIGTERM)
|
||||||
|
defer stop()
|
||||||
|
watchLoop(ctx, w, paths, convertTargetPath, convertWatchInterval, run)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,92 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWatchSnapshotExcludesOutput(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
out := filepath.Join(dir, "out")
|
||||||
|
if err := os.MkdirAll(out, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
src := filepath.Join(dir, "schema.dbml")
|
||||||
|
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
before := takeWatchSnapshot([]string{dir}, out)
|
||||||
|
if err := os.WriteFile(filepath.Join(out, "gen.go"), []byte("x"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if after := takeWatchSnapshot([]string{dir}, out); !before.equal(after) {
|
||||||
|
t.Errorf("writing into the excluded output path changed the snapshot")
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(src, []byte("changed"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if after := takeWatchSnapshot([]string{dir}, out); before.equal(after) {
|
||||||
|
t.Errorf("modifying a source file did not change the snapshot")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchLoopRerunsOnChange(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
src := filepath.Join(dir, "schema.dbml")
|
||||||
|
if err := os.WriteFile(src, []byte("a"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var runs atomic.Int32
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(done)
|
||||||
|
watchLoop(ctx, io.Discard, []string{src}, "", 10*time.Millisecond, func() error {
|
||||||
|
runs.Add(1)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
|
||||||
|
waitFor(t, func() bool { return runs.Load() == 1 })
|
||||||
|
if err := os.WriteFile(src, []byte("changed content"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
waitFor(t, func() bool { return runs.Load() == 2 })
|
||||||
|
cancel()
|
||||||
|
<-done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunConvertWatchValidation(t *testing.T) {
|
||||||
|
oldPath, oldList, oldDry, oldInt := convertSourcePath, convertFromList, convertDryRun, convertWatchInterval
|
||||||
|
defer func() {
|
||||||
|
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = oldPath, oldList, oldDry, oldInt
|
||||||
|
}()
|
||||||
|
convertSourcePath, convertFromList, convertDryRun, convertWatchInterval = "", nil, false, time.Second
|
||||||
|
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||||
|
t.Error("expected error without --from-path/--from-list")
|
||||||
|
}
|
||||||
|
convertSourcePath, convertDryRun = "x.dbml", true
|
||||||
|
if err := runConvertWatch(context.Background(), &bytes.Buffer{}, nil); err == nil {
|
||||||
|
t.Error("expected error combining --watch with --dry-run")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitFor(t *testing.T, cond func() bool) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(5 * time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if cond() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatal("condition not met in time")
|
||||||
|
}
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
# Format Usage Examples
|
||||||
|
|
||||||
|
Examples for `relspec convert` covering the file-based reader and writer
|
||||||
|
formats. The "Writers" and "Readers" sections below were run against
|
||||||
|
`examples/test_schema.dbml`. The cross-format and live-database examples were
|
||||||
|
not run; they follow the flags shown in `relspec convert --help` and require
|
||||||
|
matching input files or reachable databases.
|
||||||
|
|
||||||
|
Any reader can be combined with any writer: pick `--from`/`--from-path` for the
|
||||||
|
source and `--to`/`--to-path` for the target. Add `--silent` to suppress progress
|
||||||
|
output.
|
||||||
|
|
||||||
|
## Writers: DBML to every format
|
||||||
|
|
||||||
|
```bash
|
||||||
|
S="--from dbml --from-path examples/test_schema.dbml"
|
||||||
|
|
||||||
|
relspec convert $S --to json --to-path schema.json
|
||||||
|
relspec convert $S --to yaml --to-path schema.yaml
|
||||||
|
relspec convert $S --to dctx --to-path schema.dctx
|
||||||
|
relspec convert $S --to drawdb --to-path schema.drawdb.json
|
||||||
|
relspec convert $S --to graphql --to-path schema.graphql
|
||||||
|
relspec convert $S --to prisma --to-path schema.prisma
|
||||||
|
relspec convert $S --to pgsql --to-path schema.pg.sql
|
||||||
|
relspec convert $S --to mssql --to-path schema.mssql.sql
|
||||||
|
relspec convert $S --to sqlite --to-path schema.sqlite.sql
|
||||||
|
relspec convert $S --to drizzle --to-path schema.ts
|
||||||
|
relspec convert $S --to typeorm --to-path entities.ts
|
||||||
|
relspec convert $S --to gorm --to-path models.go --package models
|
||||||
|
relspec convert $S --to bun --to-path models.go --package models
|
||||||
|
```
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
|
||||||
|
- Code-generation writers (`gorm`, `bun`) take `--package`. They also accept
|
||||||
|
`--types baselib|stdlib|sqltypes` to choose the nullable type package.
|
||||||
|
- When `--to-path` is a directory it must already exist.
|
||||||
|
- `sqlite` output automatically flattens `schema.table` names. Use
|
||||||
|
`--flatten-schema` for other formats if the target has no schema support.
|
||||||
|
- `dctx` supports a single schema only; use `--schema <name>` to select one.
|
||||||
|
|
||||||
|
## Readers: file-based formats into DBML (or JSON where noted)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
relspec convert --from json --from-path schema.json --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from yaml --from-path schema.yaml --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from dctx --from-path schema.dctx --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from drawdb --from-path schema.drawdb.json --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from graphql --from-path schema.graphql --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from prisma --from-path schema.prisma --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from drizzle --from-path schema.ts --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from typeorm --from-path entities.ts --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from bun --from-path models.go --to dbml --to-path out.dbml
|
||||||
|
relspec convert --from gorm --from-path models.go --to json --to-path out.json
|
||||||
|
```
|
||||||
|
|
||||||
|
Code-first readers (`gorm`, `bun`, `drizzle`, `typeorm`) accept a single file or a
|
||||||
|
directory of model files.
|
||||||
|
|
||||||
|
> Known issue: reading GORM models and writing DBML currently panics in the DBML
|
||||||
|
> writer (`pkg/writers/dbml/writer.go`, `constraintToDBML`). Use another target
|
||||||
|
> such as JSON until this is fixed.
|
||||||
|
|
||||||
|
## Cross-format combinations
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# ORM models to SQL DDL
|
||||||
|
relspec convert --from gorm --from-path models.go --to pgsql --to-path schema.sql
|
||||||
|
|
||||||
|
# Prisma to Drizzle
|
||||||
|
relspec convert --from prisma --from-path schema.prisma --to drizzle --to-path schema.ts
|
||||||
|
|
||||||
|
# DrawDB diagram to GraphQL
|
||||||
|
relspec convert --from drawdb --from-path diagram.json --to graphql --to-path schema.graphql
|
||||||
|
|
||||||
|
# Merge several files while converting
|
||||||
|
relspec convert --from json --from-list "a.json,b.json" --to yaml --to-path merged.yaml
|
||||||
|
```
|
||||||
|
|
||||||
|
## Live databases
|
||||||
|
|
||||||
|
These need a reachable database:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# PostgreSQL
|
||||||
|
relspec convert --from pgsql --from-conn "postgres://user:pass@localhost:5432/mydb" \
|
||||||
|
--to dbml --to-path schema.dbml
|
||||||
|
|
||||||
|
# SQL Server
|
||||||
|
relspec convert --from mssql --from-conn "<mssql connection string>" \
|
||||||
|
--to json --to-path schema.json
|
||||||
|
|
||||||
|
# SQLite database file (--from-conn takes the file path)
|
||||||
|
relspec convert --from sqlite --from-conn ./app.db --to dbml --to-path schema.dbml
|
||||||
|
```
|
||||||
|
|
||||||
|
## Formats outside `convert`
|
||||||
|
|
||||||
|
- `sqldir` (SQL script directory reader) and `sqlexec` (SQL execution writer) are
|
||||||
|
used by `relspec scripts` and `relspec job`, and `sqldir` by `relspec diff`.
|
||||||
|
See [SCRIPTS_COMMAND.md](SCRIPTS_COMMAND.md) and [JOB_FILES.md](JOB_FILES.md).
|
||||||
|
- The `template` writer is exposed through `relspec templ`. See
|
||||||
|
[TEMPLATE_MODE.md](TEMPLATE_MODE.md).
|
||||||
@@ -0,0 +1,212 @@
|
|||||||
|
# TUI mouse support plan
|
||||||
|
|
||||||
|
Issue: #46
|
||||||
|
Status: design only; this document does not implement mouse input.
|
||||||
|
|
||||||
|
## 1. Current implementation and scope
|
||||||
|
|
||||||
|
The editor is created in `cmd/relspec/edit.go` by
|
||||||
|
`ui.NewSchemaEditorWithConfigs(...).Run()`. `pkg/ui/editor.go` owns the
|
||||||
|
`tview.Application`, `tview.Pages`, and application lifecycle. The current
|
||||||
|
code never calls `Application.EnableMouse`, so tcell mouse reporting is off.
|
||||||
|
The module uses tview v0.42.0 and tcell/v2 v2.13.9.
|
||||||
|
|
||||||
|
The first implementation should add `--no-mouse` to the `edit` Cobra command
|
||||||
|
only. The flag is a local boolean, defaulting to false, and should be passed
|
||||||
|
explicitly into the editor (prefer an options/config field rather than a
|
||||||
|
package-global or environment variable). It must not affect convert, inspect,
|
||||||
|
merge, or other commands. There is no environment-variable or persistent
|
||||||
|
configuration setting in this issue: a command-line opt-out is predictable,
|
||||||
|
visible in `edit --help`, and avoids adding configuration precedence rules.
|
||||||
|
|
||||||
|
At startup, the editor should call `app.EnableMouse(!noMouse)` before
|
||||||
|
`Run()`. `--no-mouse` must mean that the application does not enable terminal
|
||||||
|
mouse reporting and that no custom mouse handlers are relied upon. Keyboard
|
||||||
|
behavior must remain identical in both modes.
|
||||||
|
|
||||||
|
Likely implementation files are `cmd/relspec/edit.go`,
|
||||||
|
`pkg/ui/editor.go`, focused TUI mouse helpers/tests under `pkg/ui`, and a
|
||||||
|
short user-facing note in the command help or TUI documentation. Do not
|
||||||
|
refactor unrelated screens or data operations.
|
||||||
|
|
||||||
|
## 2. Widget and screen coverage
|
||||||
|
|
||||||
|
The application composes `Pages`, `Flex`, `TextView`, `List`, `Table`, `Form`,
|
||||||
|
`Button`, `InputField`, `DropDown`, `TextArea`, `CheckBox`, and `Modal`.
|
||||||
|
Vendored tview confirms mouse handlers exist for all of those relevant
|
||||||
|
primitives, including focus on left-down, list/table selection, button clicks,
|
||||||
|
form child dispatch, dropdown opening/drag selection, text-area cursor and
|
||||||
|
scrolling, and modal button dispatch. `Pages`, `Flex`, and `Form` forward events
|
||||||
|
to their children.
|
||||||
|
|
||||||
|
The coverage plan is:
|
||||||
|
|
||||||
|
* Main menu (`pkg/ui/main_menu.go`): left click focuses/selects a list entry;
|
||||||
|
second activation opens it; buttons and exit confirmation remain reachable.
|
||||||
|
* Schema, table, domain, object, relation, and database screens: click a row
|
||||||
|
to select it; double-click the row to perform the same action as the
|
||||||
|
keyboard Enter/selected callback where opening is meaningful; scroll lists
|
||||||
|
and tables; click each action button.
|
||||||
|
* Tables (`schema_screens.go`, `table_screens.go`, and object/relation tables):
|
||||||
|
tview's table handler provides selection and scrolling, but it does not
|
||||||
|
provide application-specific double-click activation. Add a small reusable
|
||||||
|
wrapper/helper for the table instances that need it. It must preserve the
|
||||||
|
existing selected row/column behavior and invoke the same callback as Enter,
|
||||||
|
not duplicate mutation logic.
|
||||||
|
* Forms (`load_save_screens.go` and the form-building screen files): click an
|
||||||
|
input to focus it, click buttons to activate them, click a dropdown to open
|
||||||
|
it and choose an option, scroll multiline help/text areas, and retain all
|
||||||
|
existing keyboard Tab/Shift-Tab, shortcut, Enter, and Escape behavior.
|
||||||
|
* Dialogs (`pkg/ui/dialogs.go` plus confirmation/error/success modals): modal
|
||||||
|
buttons are clickable and the modal keeps focus above the underlying page.
|
||||||
|
Clicking outside a modal must not activate the hidden page or dismiss a
|
||||||
|
destructive confirmation. Escape and the existing button-key behavior stay
|
||||||
|
authoritative.
|
||||||
|
* The planned file browser and connection-string builder from issue #44 must
|
||||||
|
use the same contracts: clickable entries/buttons and scrolling, with
|
||||||
|
keyboard navigation and explicit cancel/accept paths. #46 should not
|
||||||
|
implement #44's widgets; it should define the integration point and test
|
||||||
|
them when #44 lands.
|
||||||
|
|
||||||
|
Do not promise drag semantics for every widget. Drag is appropriate for text
|
||||||
|
selection/cursor movement and dropdown selection where tview already supports
|
||||||
|
it. For ordinary list/table navigation, a click selects and the wheel scrolls;
|
||||||
|
row dragging should not mutate data.
|
||||||
|
|
||||||
|
## 3. Exact mouse action contract
|
||||||
|
|
||||||
|
| Widget/type | Left down/click | Double click | Wheel/drag | Keyboard fallback |
|
||||||
|
| --- | --- | --- | --- | --- |
|
||||||
|
| Main/list menu | focus and select row | invoke row selected callback | scroll list | arrows, Enter, shortcuts |
|
||||||
|
| Data table | focus and select cell/row | invoke the screen's existing open/edit action for the selected row | vertical/horizontal scroll as supported by tview | arrows, PageUp/PageDown, Enter, existing shortcuts |
|
||||||
|
| Button | focus | same as one activation, never duplicate the callback | none | Tab/Shift-Tab, Enter/Space and existing shortcut |
|
||||||
|
| Input field | focus; place cursor if supported | no destructive action | text-area behavior if provided by tview | typing, arrows, Home/End, Tab, Escape |
|
||||||
|
| Text area/help | focus and position cursor | select word only where tview supports it; no application action | scroll; drag text selection if supported | arrows, PageUp/PageDown, standard editing keys |
|
||||||
|
| Dropdown | focus/open and choose the hit option | same as click; no duplicate selection | drag through options only while open | arrows, Enter, Escape, Tab |
|
||||||
|
| Checkbox | toggle on click | no second toggle | none | Space and existing form navigation |
|
||||||
|
| Modal | focus/click visible button | same button action once | no underlying-page scrolling | Tab/Shift-Tab, Enter, Escape, existing button keys |
|
||||||
|
| Blank/border/title area | focus containing primitive where useful | none | no mutation | current screen shortcuts |
|
||||||
|
|
||||||
|
Right and middle clicks should have no application action in the first
|
||||||
|
release. Wheel events should be consumed only by the scrollable primitive
|
||||||
|
under the pointer. Double-click timing/translation should come from tview/
|
||||||
|
tcell; custom code must not fire the action once for both the click and the
|
||||||
|
double-click. Any custom table wrapper needs a small state machine or tview
|
||||||
|
mouse action handling that is tested for this property.
|
||||||
|
|
||||||
|
## 4. tview gaps and implementation boundaries
|
||||||
|
|
||||||
|
Enabling mouse support is not sufficient for the desired behavior. tview's
|
||||||
|
built-in Table handler selects cells and scrolls but has no repository-specific
|
||||||
|
row-open callback on double click. Existing screen code also wires keyboard
|
||||||
|
input captures directly on individual widgets, so mouse actions must call the
|
||||||
|
same screen callbacks rather than route through synthetic key events.
|
||||||
|
|
||||||
|
Use tview's `MouseHandler`/`WrapMouseHandler` contracts and `setFocus` rather
|
||||||
|
than reading terminal coordinates in each screen. A reusable table adapter
|
||||||
|
may embed `*tview.Table`, delegate ordinary actions to the original table
|
||||||
|
handler, and add the screen's double-click callback. Keep the adapter in
|
||||||
|
`pkg/ui` and use it only where a row-opening action exists. Do not modify the
|
||||||
|
vendored tview copy.
|
||||||
|
|
||||||
|
The `Pages`/`Modal` dispatch order must be verified: a visible modal consumes
|
||||||
|
its click before the page below it. Page transitions should happen only in the
|
||||||
|
existing callbacks, so a stale hidden page cannot receive a click.
|
||||||
|
|
||||||
|
## 5. Keyboard, terminal, and copy/paste behavior
|
||||||
|
|
||||||
|
Mouse is an enhancement, never a requirement. Every acceptance path must be
|
||||||
|
reachable with the existing keyboard controls, including load/save, navigation,
|
||||||
|
editing, confirmations, cancel, and exit. `--no-mouse` is the regression mode
|
||||||
|
for proving this contract.
|
||||||
|
|
||||||
|
Mouse reporting is terminal capability dependent. On local terminals it is
|
||||||
|
negotiated by tcell; tmux and SSH can suppress, translate, or fail to pass
|
||||||
|
mouse reporting depending on their configuration. The application must still
|
||||||
|
start and remain keyboard usable if mouse reporting is unavailable or broken.
|
||||||
|
Documentation should state that terminal/tmux configuration may be required,
|
||||||
|
and that SSH behavior depends on the remote terminal path. Windows Terminal and
|
||||||
|
other Windows console hosts should be treated as supported only insofar as the
|
||||||
|
selected tcell backend reports mouse events; the CLI must not assume POSIX
|
||||||
|
escape sequences or add platform-specific code in this issue.
|
||||||
|
|
||||||
|
Enabling mouse capture normally prevents terminal-native selection/copy from
|
||||||
|
seeing ordinary button-drag events. Document the standard workaround: hold the
|
||||||
|
terminal's bypass modifier (commonly Shift, terminal-dependent) for selection,
|
||||||
|
or use `--no-mouse` when native copy/paste is the priority. Do not implement a
|
||||||
|
second clipboard protocol. Input-field/text-area copy/paste must continue to
|
||||||
|
use tview/tcell paste handling and keyboard shortcuts; verify that enabling
|
||||||
|
mouse does not intercept paste events.
|
||||||
|
|
||||||
|
## 6. Test strategy using tcell simulation
|
||||||
|
|
||||||
|
Add focused tests rather than attempting a full interactive end-to-end test.
|
||||||
|
Use `tcell.NewSimulationScreen("")`, `screen.Init()`, construct the editor or
|
||||||
|
an isolated primitive tree, and inject events with the actual API:
|
||||||
|
`SimulationScreen.InjectMouse(x, y, buttons, mod)` and `InjectKey(...)`.
|
||||||
|
Coordinates must be derived from the primitive's drawn rectangle or fixed by a
|
||||||
|
small deterministic test layout; do not use arbitrary coordinates without
|
||||||
|
checking the rendered screen.
|
||||||
|
|
||||||
|
Minimum cases:
|
||||||
|
|
||||||
|
1. Default editor configuration enables mouse; the explicit disabled option
|
||||||
|
leaves it disabled. If the Application API is not observable directly,
|
||||||
|
test through the simulation screen's event path plus a constructor-level
|
||||||
|
option assertion.
|
||||||
|
2. A list click changes focus/selection, and double-click invokes the existing
|
||||||
|
selected action exactly once.
|
||||||
|
3. A table click selects the expected row/cell, wheel events change the visible
|
||||||
|
offset, and double-click invokes the row action exactly once.
|
||||||
|
4. Form button, input field, checkbox, and dropdown clicks match their
|
||||||
|
keyboard callbacks.
|
||||||
|
5. A modal button click acts on the modal and cannot activate the underlying
|
||||||
|
page; Escape still cancels.
|
||||||
|
6. `--no-mouse` leaves keyboard selection/activation unchanged and mouse
|
||||||
|
injection has no application effect.
|
||||||
|
7. Existing dialogs and screen transitions do not leave a stale mouse capture
|
||||||
|
after a page is removed.
|
||||||
|
|
||||||
|
Prefer callback counters and selected-index assertions over screen-text-only
|
||||||
|
assertions. Run the relevant `pkg/ui` tests with `go test -race ./pkg/ui` and
|
||||||
|
run the full repository test suite if time/resources permit.
|
||||||
|
|
||||||
|
## 7. Rollout and acceptance criteria
|
||||||
|
|
||||||
|
Implementation is ready for review when:
|
||||||
|
|
||||||
|
* `relspec edit --help` documents `--no-mouse` and mouse is enabled by default.
|
||||||
|
* Only the edit TUI is affected; non-TUI commands have no changed behavior.
|
||||||
|
* Main screens, tables, lists, forms, dropdowns, buttons, text areas, and
|
||||||
|
visible dialogs support the action contract above.
|
||||||
|
* Keyboard-only operation is complete and verified with `--no-mouse`.
|
||||||
|
* Modal clicks cannot fall through to an underlying page.
|
||||||
|
* Table double-click behavior is explicit, tested, and does not duplicate
|
||||||
|
activation.
|
||||||
|
* tcell simulation tests cover default-on, opt-out, selection, scrolling,
|
||||||
|
activation, dialog focus, and keyboard fallback.
|
||||||
|
* `go test -race ./pkg/ui`, appropriate command tests, `go test ./...`,
|
||||||
|
formatting, and `git diff --check` pass (or any limitation is recorded).
|
||||||
|
* User-facing docs explain tmux/SSH/Windows variability and the terminal
|
||||||
|
modifier workaround for native copy/paste.
|
||||||
|
|
||||||
|
Roll out in two implementation slices if needed: first application option,
|
||||||
|
standard tview handlers, tests, and documentation; second only the reusable
|
||||||
|
table double-click adapter and screen wiring. Do not block the first slice on
|
||||||
|
issue #44, but do not claim #44's future widgets are covered until they use the
|
||||||
|
same contract.
|
||||||
|
|
||||||
|
## 8. Open decisions and dependencies
|
||||||
|
|
||||||
|
* Confirm whether the project wants a public editor options type or a small
|
||||||
|
`SetMouseEnabled`/constructor parameter; avoid a global flag.
|
||||||
|
* Confirm the preferred terminal copy modifier in project documentation, since
|
||||||
|
tmux, SSH clients, and Windows Terminal differ.
|
||||||
|
* Decide whether horizontal wheel events should be supported where tview/table
|
||||||
|
exposes them; vertical scrolling is mandatory, horizontal is optional.
|
||||||
|
* Decide whether double-click opens every data table or only tables with an
|
||||||
|
unambiguous row action. The plan recommends the latter.
|
||||||
|
* Confirm #44's file-browser and connection-builder primitive choices before
|
||||||
|
wiring their mouse tests.
|
||||||
|
* Confirm CI has a stable non-terminal environment for simulation-screen tests;
|
||||||
|
no real terminal, tmux session, database, or network should be required.
|
||||||
@@ -4,6 +4,7 @@ go 1.25.13
|
|||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/gdamore/tcell/v2 v2.13.9
|
github.com/gdamore/tcell/v2 v2.13.9
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/jackc/pgx/v5 v5.9.2
|
github.com/jackc/pgx/v5 v5.9.2
|
||||||
github.com/microsoft/go-mssqldb v1.10.0
|
github.com/microsoft/go-mssqldb v1.10.0
|
||||||
@@ -18,6 +19,7 @@ require (
|
|||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
filippo.io/edwards25519 v1.1.0 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/gdamore/encoding v1.0.1 // indirect
|
github.com/gdamore/encoding v1.0.1 // indirect
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
|
||||||
|
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1 h1:jHb/wfvRikGdxMXYV3QG/SzUOPYN9KEUUuC0Yd0/vC0=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
github.com/Azure/azure-sdk-for-go/sdk/azcore v1.21.1/go.mod h1:pzBXCYn05zvYIrwLgtK8Ap8QcjRg+0i76tMQdWN6wOk=
|
||||||
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
github.com/Azure/azure-sdk-for-go/sdk/azidentity v1.13.1 h1:Hk5QBxZQC1jb2Fwj6mpzme37xbCDdNTxU7O9eb5+LB4=
|
||||||
@@ -21,6 +23,8 @@ github.com/gdamore/encoding v1.0.1 h1:YzKZckdBL6jVt2Gc+5p82qhrGiqMdG/eNs6Wy0u3Uh
|
|||||||
github.com/gdamore/encoding v1.0.1/go.mod h1:0Z0cMFinngz9kS1QfMjCP8TY7em3bZYeeklsSDPivEo=
|
github.com/gdamore/encoding v1.0.1/go.mod h1:0Z0cMFinngz9kS1QfMjCP8TY7em3bZYeeklsSDPivEo=
|
||||||
github.com/gdamore/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
|
github.com/gdamore/tcell/v2 v2.13.9 h1:uI5l3DYPcFvHINKlGft+en23evOKL+dwtD21QR8ejVA=
|
||||||
github.com/gdamore/tcell/v2 v2.13.9/go.mod h1:+Wfe208WDdB7INEtCsNrAN6O2m+wsTPk1RAovjaILlo=
|
github.com/gdamore/tcell/v2 v2.13.9/go.mod h1:+Wfe208WDdB7INEtCsNrAN6O2m+wsTPk1RAovjaILlo=
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo=
|
||||||
|
github.com/go-sql-driver/mysql v1.9.3/go.mod h1:qn46aNg1333BRMNU69Lq93t8du/dwxI64Gl8i5p1WMU=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
|
||||||
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
|
||||||
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
github.com/golang-sql/civil v0.0.0-20220223132316-b832511892a9 h1:au07oEsX2xN0ktxqI+Sida1w446QrXBRJ0nee3SNZlA=
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ const (
|
|||||||
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
|
PostgresqlDatabaseType DatabaseType = "pgsql" // PostgreSQL database
|
||||||
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
|
MSSQLDatabaseType DatabaseType = "mssql" // Microsoft SQL Server database
|
||||||
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
|
SqlLiteDatabaseType DatabaseType = "sqlite" // SQLite database
|
||||||
|
MySQLDatabaseType DatabaseType = "mysql" // MySQL/MariaDB database
|
||||||
)
|
)
|
||||||
|
|
||||||
// Database represents the complete database schema
|
// Database represents the complete database schema
|
||||||
|
|||||||
@@ -0,0 +1,254 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
_ "github.com/go-sql-driver/mysql"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Reader struct {
|
||||||
|
options *readers.ReaderOptions
|
||||||
|
db *sql.DB
|
||||||
|
ctx context.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewReader(options *readers.ReaderOptions) *Reader {
|
||||||
|
return &Reader{options: options, ctx: context.Background()}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadDatabase() (*models.Database, error) {
|
||||||
|
if r.options == nil || r.options.ConnectionString == "" {
|
||||||
|
return nil, fmt.Errorf("connection string is required")
|
||||||
|
}
|
||||||
|
if err := r.connect(); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to connect: %w", err)
|
||||||
|
}
|
||||||
|
defer r.close()
|
||||||
|
var name, version string
|
||||||
|
if err := r.db.QueryRowContext(r.ctx, "SELECT DATABASE()").Scan(&name); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to get database name: %w", err)
|
||||||
|
}
|
||||||
|
_ = r.db.QueryRowContext(r.ctx, "SELECT VERSION()").Scan(&version)
|
||||||
|
db := models.InitDatabase(name)
|
||||||
|
db.DatabaseType, db.SourceFormat, db.DatabaseVersion = models.MySQLDatabaseType, "mysql", version
|
||||||
|
schemas, err := r.querySchemas(name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to query schemas: %w", err)
|
||||||
|
}
|
||||||
|
for _, schema := range schemas {
|
||||||
|
tables, err := r.queryTables(schema.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
schema.Tables = tables
|
||||||
|
for _, table := range tables {
|
||||||
|
table.Columns, err = r.queryColumns(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.Constraints, err = r.queryConstraints(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.Indexes, err = r.queryIndexes(schema.Name, table.Name)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
table.RefSchema = schema
|
||||||
|
for _, c := range table.Constraints {
|
||||||
|
if c.Type == models.ForeignKeyConstraint {
|
||||||
|
r.deriveRelationship(table, c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
schema.RefDatabase = db
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
}
|
||||||
|
return db, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadSchema() (*models.Schema, error) {
|
||||||
|
db, err := r.ReadDatabase()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(db.Schemas) == 0 {
|
||||||
|
return nil, fmt.Errorf("no schemas found in database")
|
||||||
|
}
|
||||||
|
return db.Schemas[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) ReadTable() (*models.Table, error) {
|
||||||
|
s, err := r.ReadSchema()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(s.Tables) == 0 {
|
||||||
|
return nil, fmt.Errorf("no tables found in schema")
|
||||||
|
}
|
||||||
|
return s.Tables[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) connect() error {
|
||||||
|
db, err := sql.Open("mysql", r.options.ConnectionString)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err = db.PingContext(r.ctx); err != nil {
|
||||||
|
db.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
r.db = db
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) close() {
|
||||||
|
if r.db != nil {
|
||||||
|
_ = r.db.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
func (r *Reader) mapDataType(t string) string { return mariadb.ConvertMariaDBToCanonical(t) }
|
||||||
|
|
||||||
|
func (r *Reader) querySchemas(current string) ([]*models.Schema, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, "SELECT SCHEMA_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME = ?", current)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var out []*models.Schema
|
||||||
|
for rows.Next() {
|
||||||
|
var n string
|
||||||
|
if err := rows.Scan(&n); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out = append(out, models.InitSchema(n))
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryTables(schema string) ([]*models.Table, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, "SELECT TABLE_NAME FROM information_schema.TABLES WHERE TABLE_SCHEMA = ? AND TABLE_TYPE = 'BASE TABLE' ORDER BY TABLE_NAME", schema)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var out []*models.Table
|
||||||
|
for rows.Next() {
|
||||||
|
var n string
|
||||||
|
if err := rows.Scan(&n); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out = append(out, models.InitTable(n, schema))
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryColumns(schema, table string) (map[string]*models.Column, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT COLUMN_NAME, COLUMN_TYPE, IS_NULLABLE, COLUMN_DEFAULT, ORDINAL_POSITION, EXTRA, COLUMN_COMMENT FROM information_schema.COLUMNS WHERE TABLE_SCHEMA = ? AND TABLE_NAME = ? ORDER BY ORDINAL_POSITION`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Column{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, typ, nullable, extra, comment string
|
||||||
|
var def sql.NullString
|
||||||
|
var pos int
|
||||||
|
if err := rows.Scan(&name, &typ, &nullable, &def, &pos, &extra, &comment); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c := models.InitColumn(name, table, schema)
|
||||||
|
c.Type = r.mapDataType(typ)
|
||||||
|
c.NotNull = strings.EqualFold(nullable, "NO")
|
||||||
|
c.Sequence = uint(pos)
|
||||||
|
c.Comment = comment
|
||||||
|
if def.Valid {
|
||||||
|
c.Default = def.String
|
||||||
|
}
|
||||||
|
c.AutoIncrement = strings.Contains(strings.ToLower(extra), "auto_increment")
|
||||||
|
out[name] = c
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryConstraints(schema, table string) (map[string]*models.Constraint, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT CONSTRAINT_NAME, CONSTRAINT_TYPE, COLUMN_NAME, REFERENCED_TABLE_SCHEMA, REFERENCED_TABLE_NAME, REFERENCED_COLUMN_NAME, ORDINAL_POSITION FROM information_schema.KEY_COLUMN_USAGE k JOIN information_schema.TABLE_CONSTRAINTS t USING (CONSTRAINT_SCHEMA, TABLE_NAME, CONSTRAINT_NAME) WHERE k.TABLE_SCHEMA=? AND k.TABLE_NAME=? ORDER BY CONSTRAINT_NAME, ORDINAL_POSITION`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Constraint{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, typ, col, rs, rt, rc string
|
||||||
|
var pos int
|
||||||
|
if err := rows.Scan(&name, &typ, &col, &rs, &rt, &rc, &pos); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c := out[name]
|
||||||
|
if c == nil {
|
||||||
|
ct := models.UniqueConstraint
|
||||||
|
if typ == "PRIMARY KEY" {
|
||||||
|
ct = models.PrimaryKeyConstraint
|
||||||
|
}
|
||||||
|
if typ == "FOREIGN KEY" {
|
||||||
|
ct = models.ForeignKeyConstraint
|
||||||
|
}
|
||||||
|
c = models.InitConstraint(name, ct)
|
||||||
|
c.Schema = schema
|
||||||
|
c.Table = table
|
||||||
|
c.ReferencedSchema = rs
|
||||||
|
c.ReferencedTable = rt
|
||||||
|
out[name] = c
|
||||||
|
}
|
||||||
|
c.Columns = append(c.Columns, col)
|
||||||
|
if rc != "" {
|
||||||
|
c.ReferencedColumns = append(c.ReferencedColumns, rc)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) queryIndexes(schema, table string) (map[string]*models.Index, error) {
|
||||||
|
rows, err := r.db.QueryContext(r.ctx, `SELECT INDEX_NAME, NON_UNIQUE, COLUMN_NAME, SEQ_IN_INDEX, INDEX_TYPE FROM information_schema.STATISTICS WHERE TABLE_SCHEMA=? AND TABLE_NAME=? AND INDEX_NAME <> 'PRIMARY' ORDER BY INDEX_NAME, SEQ_IN_INDEX`, schema, table)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := map[string]*models.Index{}
|
||||||
|
for rows.Next() {
|
||||||
|
var name, col, typ string
|
||||||
|
var non, seq int
|
||||||
|
if err := rows.Scan(&name, &non, &col, &seq, &typ); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
i := out[name]
|
||||||
|
if i == nil {
|
||||||
|
i = models.InitIndex(name, table, schema)
|
||||||
|
i.Unique = non == 0
|
||||||
|
i.Type = strings.ToLower(typ)
|
||||||
|
out[name] = i
|
||||||
|
}
|
||||||
|
i.Columns = append(i.Columns, col)
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Reader) deriveRelationship(t *models.Table, c *models.Constraint) {
|
||||||
|
n := fmt.Sprintf("%s_to_%s", t.Name, c.ReferencedTable)
|
||||||
|
rel := models.InitRelationship(n, models.OneToMany)
|
||||||
|
rel.FromTable = t.Name
|
||||||
|
rel.FromSchema = t.Schema
|
||||||
|
rel.FromColumns = append([]string(nil), c.Columns...)
|
||||||
|
rel.ToTable = c.ReferencedTable
|
||||||
|
rel.ToSchema = c.ReferencedSchema
|
||||||
|
rel.ToColumns = append([]string(nil), c.ReferencedColumns...)
|
||||||
|
rel.ForeignKey = c.Name
|
||||||
|
t.Relationships[n] = rel
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestReaderMapDataType(t *testing.T) {
|
||||||
|
r := NewReader(&readers.ReaderOptions{})
|
||||||
|
for _, tc := range []struct{ input, want string }{{"varchar(64)", "string"}, {"bigint unsigned", "int64"}, {"datetime", "timestamp"}, {"json", "json"}} {
|
||||||
|
if got := r.mapDataType(tc.input); got != tc.want {
|
||||||
|
t.Errorf("mapDataType(%q) = %q, want %q", tc.input, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReaderRequiresConnectionString(t *testing.T) {
|
||||||
|
if _, err := NewReader(&readers.ReaderOptions{}).ReadDatabase(); err == nil {
|
||||||
|
t.Fatal("expected missing connection string error")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,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,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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -207,6 +207,10 @@ func (se *SchemaEditor) showDomainEditor(index int, domain *models.Domain) {
|
|||||||
se.showDomainList()
|
se.showDomainList()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
form.AddButton("Tables", func() {
|
||||||
|
se.showDomainTables(index)
|
||||||
|
})
|
||||||
|
|
||||||
form.AddButton("Delete", func() {
|
form.AddButton("Delete", func() {
|
||||||
se.showDeleteDomainConfirm(index)
|
se.showDeleteDomainConfirm(index)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -0,0 +1,134 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileEntry is a single row in the file browser.
|
||||||
|
type FileEntry struct {
|
||||||
|
Name string
|
||||||
|
IsDir bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// formatExtensions maps a UI format name to the file extensions it reads or writes.
|
||||||
|
var formatExtensions = map[string][]string{
|
||||||
|
"dbml": {".dbml"},
|
||||||
|
"dctx": {".dctx"},
|
||||||
|
"drawdb": {".json"},
|
||||||
|
"graphql": {".graphql", ".gql"},
|
||||||
|
"json": {".json"},
|
||||||
|
"yaml": {".yaml", ".yml"},
|
||||||
|
"gorm": {".go"},
|
||||||
|
"bun": {".go"},
|
||||||
|
"drizzle": {".ts"},
|
||||||
|
"prisma": {".prisma"},
|
||||||
|
"typeorm": {".ts"},
|
||||||
|
"pgsql": {".sql"},
|
||||||
|
"sqlite": {".db", ".sqlite", ".sqlite3"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// directoryFormats are formats whose reader/writer accepts a directory.
|
||||||
|
var directoryFormats = map[string]bool{
|
||||||
|
"gorm": true, "bun": true, "drizzle": true, "typeorm": true,
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatExtensions returns the extensions for a format, or nil (no filter) if unknown.
|
||||||
|
func FormatExtensions(format string) []string {
|
||||||
|
return formatExtensions[format]
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsDirectoryFormat reports whether a format can be loaded from or saved to a directory.
|
||||||
|
func IsDirectoryFormat(format string) bool {
|
||||||
|
return directoryFormats[format]
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExpandHome replaces a leading ~ with the user's home directory.
|
||||||
|
func ExpandHome(p string) string {
|
||||||
|
if strings.HasPrefix(p, "~") {
|
||||||
|
if home, err := os.UserHomeDir(); err == nil {
|
||||||
|
return filepath.Join(home, p[1:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchesExtension reports whether name has one of exts (case-insensitive).
|
||||||
|
// An empty extension list matches everything.
|
||||||
|
func MatchesExtension(name string, exts []string) bool {
|
||||||
|
if len(exts) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
ext := strings.ToLower(filepath.Ext(name))
|
||||||
|
for _, e := range exts {
|
||||||
|
if strings.EqualFold(e, ext) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListDir returns the entries of dir: directories first, then files that match
|
||||||
|
// exts, each group sorted case-insensitively. Hidden (dot) entries are skipped
|
||||||
|
// unless showHidden is set.
|
||||||
|
func ListDir(dir string, exts []string, showHidden bool) ([]FileEntry, error) {
|
||||||
|
items, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var dirs, files []FileEntry
|
||||||
|
for _, item := range items {
|
||||||
|
name := item.Name()
|
||||||
|
if !showHidden && strings.HasPrefix(name, ".") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
isDir := item.IsDir()
|
||||||
|
if !isDir && item.Type()&os.ModeSymlink != 0 {
|
||||||
|
// Follow symlinks so links to directories are navigable.
|
||||||
|
if info, err := os.Stat(filepath.Join(dir, name)); err == nil {
|
||||||
|
isDir = info.IsDir()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if isDir {
|
||||||
|
dirs = append(dirs, FileEntry{Name: name, IsDir: true})
|
||||||
|
} else if MatchesExtension(name, exts) {
|
||||||
|
files = append(files, FileEntry{Name: name})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
byName := func(s []FileEntry) {
|
||||||
|
sort.Slice(s, func(i, j int) bool {
|
||||||
|
return strings.ToLower(s[i].Name) < strings.ToLower(s[j].Name)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
byName(dirs)
|
||||||
|
byName(files)
|
||||||
|
return append(dirs, files...), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResolveStart works out where the browser should open for the current input
|
||||||
|
// value. It returns the directory to show and, if the input named a file, its
|
||||||
|
// base name. Falls back to the working directory.
|
||||||
|
func ResolveStart(input string) (dir, name string) {
|
||||||
|
input = strings.TrimSpace(input)
|
||||||
|
if input != "" {
|
||||||
|
p := ExpandHome(input)
|
||||||
|
if abs, err := filepath.Abs(p); err == nil {
|
||||||
|
p = abs
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(p); err == nil && info.IsDir() {
|
||||||
|
return p, ""
|
||||||
|
}
|
||||||
|
if info, err := os.Stat(filepath.Dir(p)); err == nil && info.IsDir() {
|
||||||
|
return filepath.Dir(p), filepath.Base(p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wd, err := os.Getwd()
|
||||||
|
if err != nil {
|
||||||
|
wd = "."
|
||||||
|
}
|
||||||
|
return wd, ""
|
||||||
|
}
|
||||||
@@ -0,0 +1,363 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"github.com/gdamore/tcell/v2"
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileBrowserMode selects between picking an existing path and choosing a save target.
|
||||||
|
type FileBrowserMode int
|
||||||
|
|
||||||
|
const (
|
||||||
|
FileBrowserLoad FileBrowserMode = iota
|
||||||
|
FileBrowserSave
|
||||||
|
)
|
||||||
|
|
||||||
|
// FileBrowserConfig configures the file browser dialog.
|
||||||
|
type FileBrowserConfig struct {
|
||||||
|
Mode FileBrowserMode
|
||||||
|
StartPath string // current value of the input; may be empty
|
||||||
|
Extensions []string // empty = show all files
|
||||||
|
AllowDir bool // a directory is a valid result (directory-based formats)
|
||||||
|
ReturnPage string // page to switch back to when the dialog closes
|
||||||
|
OnSelect func(path string)
|
||||||
|
}
|
||||||
|
|
||||||
|
// attachFileBrowser makes Enter on the named input open the file browser,
|
||||||
|
// filtered for the currently selected format.
|
||||||
|
func (se *SchemaEditor) attachFileBrowser(form *tview.Form, label, returnPage string, mode FileBrowserMode, format func() string) {
|
||||||
|
item, ok := form.GetFormItemByLabel(label).(*tview.InputField)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
item.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() != tcell.KeyEnter {
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
f := format()
|
||||||
|
se.showFileBrowser(FileBrowserConfig{
|
||||||
|
Mode: mode,
|
||||||
|
StartPath: item.GetText(),
|
||||||
|
Extensions: FormatExtensions(f),
|
||||||
|
AllowDir: IsDirectoryFormat(f),
|
||||||
|
ReturnPage: returnPage,
|
||||||
|
OnSelect: func(path string) { item.SetText(path) },
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// showFileBrowser displays the file browser page. Esc closes it without
|
||||||
|
// calling OnSelect, leaving the originating input unchanged.
|
||||||
|
func (se *SchemaEditor) showFileBrowser(cfg FileBrowserConfig) {
|
||||||
|
const pageName = "file-browser"
|
||||||
|
|
||||||
|
dir, startName := ResolveStart(cfg.StartPath)
|
||||||
|
showHidden := false
|
||||||
|
useFilter := len(cfg.Extensions) > 0
|
||||||
|
var entries []FileEntry // rows shown below the ".." row
|
||||||
|
|
||||||
|
title := tview.NewTextView().
|
||||||
|
SetText("[::b]Select File").
|
||||||
|
SetTextAlign(tview.AlignCenter).
|
||||||
|
SetDynamicColors(true)
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
title.SetText("[::b]Save As")
|
||||||
|
}
|
||||||
|
|
||||||
|
info := tview.NewTextView().SetDynamicColors(true)
|
||||||
|
|
||||||
|
fileTable := tview.NewTable().SetSelectable(true, false).SetFixed(0, 0)
|
||||||
|
fileTable.SetBorder(true)
|
||||||
|
|
||||||
|
nameInput := tview.NewInputField().SetLabel("File name: ").SetFieldWidth(0)
|
||||||
|
nameInput.SetText(startName)
|
||||||
|
|
||||||
|
closeBrowser := func() {
|
||||||
|
se.pages.RemovePage(pageName)
|
||||||
|
se.pages.SwitchToPage(cfg.ReturnPage)
|
||||||
|
}
|
||||||
|
|
||||||
|
finish := func(path string) {
|
||||||
|
closeBrowser()
|
||||||
|
cfg.OnSelect(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
refresh := func() {
|
||||||
|
exts := cfg.Extensions
|
||||||
|
if !useFilter {
|
||||||
|
exts = nil
|
||||||
|
}
|
||||||
|
list, err := ListDir(dir, exts, showHidden)
|
||||||
|
if err != nil {
|
||||||
|
se.showErrorDialog("Error", fmt.Sprintf("Cannot read %s: %v", dir, err))
|
||||||
|
list = nil
|
||||||
|
}
|
||||||
|
entries = list
|
||||||
|
|
||||||
|
fileTable.Clear()
|
||||||
|
fileTable.SetCell(0, 0, tview.NewTableCell("[..]").SetTextColor(tcell.ColorAqua))
|
||||||
|
for i, e := range entries {
|
||||||
|
cell := tview.NewTableCell(e.Name)
|
||||||
|
if e.IsDir {
|
||||||
|
cell.SetText(e.Name + "/").SetTextColor(tcell.ColorAqua)
|
||||||
|
}
|
||||||
|
fileTable.SetCell(i+1, 0, cell)
|
||||||
|
}
|
||||||
|
fileTable.Select(0, 0)
|
||||||
|
if len(entries) > 0 {
|
||||||
|
fileTable.Select(1, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
filterText := "all files"
|
||||||
|
if useFilter {
|
||||||
|
filterText = fmt.Sprintf("%v", cfg.Extensions)
|
||||||
|
}
|
||||||
|
hiddenText := "hidden: off"
|
||||||
|
if showHidden {
|
||||||
|
hiddenText = "hidden: on"
|
||||||
|
}
|
||||||
|
info.SetText(fmt.Sprintf("%s [yellow](%s, filter: %s)[-]", tview.Escape(dir), hiddenText, tview.Escape(filterText)))
|
||||||
|
fileTable.SetTitle(" Files ")
|
||||||
|
}
|
||||||
|
|
||||||
|
goUp := func() {
|
||||||
|
parent := filepath.Dir(dir)
|
||||||
|
if parent == dir {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
prev := filepath.Base(dir)
|
||||||
|
dir = parent
|
||||||
|
refresh()
|
||||||
|
for i, e := range entries {
|
||||||
|
if e.Name == prev {
|
||||||
|
fileTable.Select(i+1, 0)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
selected := func() (FileEntry, bool) {
|
||||||
|
row, _ := fileTable.GetSelection()
|
||||||
|
if row < 1 || row > len(entries) {
|
||||||
|
return FileEntry{}, false
|
||||||
|
}
|
||||||
|
return entries[row-1], true
|
||||||
|
}
|
||||||
|
|
||||||
|
// confirmOverwrite asks before replacing an existing file (not directories).
|
||||||
|
confirmOverwrite := func(path string) {
|
||||||
|
modal := tview.NewModal().
|
||||||
|
SetText(fmt.Sprintf("File already exists:\n%s\n\nOverwrite it?", path)).
|
||||||
|
AddButtons([]string{"Cancel", "Overwrite"}).
|
||||||
|
SetDoneFunc(func(_ int, label string) {
|
||||||
|
se.pages.RemovePage("overwrite-confirm")
|
||||||
|
se.pages.SwitchToPage(pageName)
|
||||||
|
if label == "Overwrite" {
|
||||||
|
finish(path)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
modal.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage("overwrite-confirm")
|
||||||
|
se.pages.SwitchToPage(pageName)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddAndSwitchToPage("overwrite-confirm", modal, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
chooseSave := func() {
|
||||||
|
name := nameInput.GetText()
|
||||||
|
if name == "" {
|
||||||
|
if cfg.AllowDir {
|
||||||
|
finish(dir)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
se.showErrorDialog("Error", "Enter a file name")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
path := filepath.Join(dir, name)
|
||||||
|
if st, err := os.Stat(path); err == nil {
|
||||||
|
if st.IsDir() {
|
||||||
|
se.showErrorDialog("Error", name+" is a directory")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
confirmOverwrite(path)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
// chooseHighlighted handles Select: the highlighted entry in load mode, or
|
||||||
|
// the typed name in save mode.
|
||||||
|
chooseHighlighted := func() {
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
chooseSave()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e, ok := selected()
|
||||||
|
switch {
|
||||||
|
case ok && !e.IsDir:
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
case ok && cfg.AllowDir:
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
case cfg.AllowDir:
|
||||||
|
finish(dir)
|
||||||
|
default:
|
||||||
|
se.showErrorDialog("Error", "Select a file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
activate := func() {
|
||||||
|
row, _ := fileTable.GetSelection()
|
||||||
|
if row == 0 {
|
||||||
|
goUp()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e, ok := selected()
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if e.IsDir {
|
||||||
|
dir = filepath.Join(dir, e.Name)
|
||||||
|
refresh()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
nameInput.SetText(e.Name)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(filepath.Join(dir, e.Name))
|
||||||
|
}
|
||||||
|
|
||||||
|
toggleHidden := func() { showHidden = !showHidden; refresh() }
|
||||||
|
toggleFilter := func() {
|
||||||
|
if len(cfg.Extensions) > 0 {
|
||||||
|
useFilter = !useFilter
|
||||||
|
refresh()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
btnSelect := tview.NewButton("Select [s]").SetSelectedFunc(chooseHighlighted)
|
||||||
|
btnHidden := tview.NewButton("Hidden [h]").SetSelectedFunc(toggleHidden)
|
||||||
|
btnFilter := tview.NewButton("Filter [f]").SetSelectedFunc(toggleFilter)
|
||||||
|
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(closeBrowser)
|
||||||
|
|
||||||
|
btnFlex := tview.NewFlex().
|
||||||
|
AddItem(btnSelect, 0, 1, false).
|
||||||
|
AddItem(btnHidden, 0, 1, false).
|
||||||
|
AddItem(btnFilter, 0, 1, false).
|
||||||
|
AddItem(btnBack, 0, 1, false)
|
||||||
|
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow).
|
||||||
|
AddItem(title, 1, 0, false).
|
||||||
|
AddItem(info, 1, 0, false).
|
||||||
|
AddItem(fileTable, 0, 1, true)
|
||||||
|
|
||||||
|
focusOrder := []tview.Primitive{fileTable}
|
||||||
|
if cfg.Mode == FileBrowserSave {
|
||||||
|
flex.AddItem(nameInput, 1, 0, false)
|
||||||
|
focusOrder = append(focusOrder, nameInput)
|
||||||
|
}
|
||||||
|
flex.AddItem(btnFlex, 1, 0, false)
|
||||||
|
focusOrder = append(focusOrder, btnSelect, btnHidden, btnFilter, btnBack)
|
||||||
|
|
||||||
|
// Circular Tab / Shift+Tab across every focusable widget.
|
||||||
|
cycle := func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
step := 0
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyTab:
|
||||||
|
step = 1
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
step = -1
|
||||||
|
default:
|
||||||
|
return event
|
||||||
|
}
|
||||||
|
for i, p := range focusOrder {
|
||||||
|
if p.HasFocus() {
|
||||||
|
se.app.SetFocus(focusOrder[(i+step+len(focusOrder))%len(focusOrder)])
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
fileTable.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyEscape:
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyEnter:
|
||||||
|
activate()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyBackspace, tcell.KeyBackspace2, tcell.KeyLeft:
|
||||||
|
goUp()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Rune() {
|
||||||
|
case 's':
|
||||||
|
chooseHighlighted()
|
||||||
|
return nil
|
||||||
|
case 'h':
|
||||||
|
toggleHidden()
|
||||||
|
return nil
|
||||||
|
case 'f':
|
||||||
|
toggleFilter()
|
||||||
|
return nil
|
||||||
|
case 'b':
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
nameInput.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyEscape:
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
case tcell.KeyEnter:
|
||||||
|
chooseSave()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, b := range []*tview.Button{btnSelect, btnHidden, btnFilter, btnBack} {
|
||||||
|
b.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event = cycle(event); event == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
closeBrowser()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
refresh()
|
||||||
|
if startName != "" {
|
||||||
|
for i, e := range entries {
|
||||||
|
if e.Name == startName {
|
||||||
|
fileTable.Select(i+1, 0)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
se.pages.AddAndSwitchToPage(pageName, flex, true)
|
||||||
|
se.app.SetFocus(fileTable)
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func touch(t *testing.T, path string) {
|
||||||
|
t.Helper()
|
||||||
|
if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(path, nil, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func names(entries []FileEntry) []string {
|
||||||
|
var out []string
|
||||||
|
for _, e := range entries {
|
||||||
|
if e.IsDir {
|
||||||
|
out = append(out, e.Name+"/")
|
||||||
|
} else {
|
||||||
|
out = append(out, e.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchesExtension(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
exts []string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"a.dbml", []string{".dbml"}, true},
|
||||||
|
{"A.DBML", []string{".dbml"}, true},
|
||||||
|
{"a.json", []string{".dbml"}, false},
|
||||||
|
{"a.yml", []string{".yaml", ".yml"}, true},
|
||||||
|
{"noext", []string{".sql"}, false},
|
||||||
|
{"anything", nil, true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := MatchesExtension(tt.name, tt.exts); got != tt.want {
|
||||||
|
t.Errorf("MatchesExtension(%q, %v) = %v, want %v", tt.name, tt.exts, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListDirFilterAndHidden(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
touch(t, filepath.Join(dir, "b.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "A.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "c.json"))
|
||||||
|
touch(t, filepath.Join(dir, ".hidden.dbml"))
|
||||||
|
touch(t, filepath.Join(dir, "sub", "x.txt"))
|
||||||
|
touch(t, filepath.Join(dir, ".git", "x"))
|
||||||
|
|
||||||
|
got, err := ListDir(dir, FormatExtensions("dbml"), false)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if want := []string{"sub/", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("filtered: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, _ = ListDir(dir, FormatExtensions("dbml"), true)
|
||||||
|
if want := []string{".git/", "sub/", ".hidden.dbml", "A.dbml", "b.dbml"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("hidden: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, _ = ListDir(dir, nil, false)
|
||||||
|
if want := []string{"sub/", "A.dbml", "b.dbml", "c.json"}; !reflect.DeepEqual(names(got), want) {
|
||||||
|
t.Errorf("no filter: got %v, want %v", names(got), want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListDirMissing(t *testing.T) {
|
||||||
|
if _, err := ListDir(filepath.Join(t.TempDir(), "nope"), nil, false); err == nil {
|
||||||
|
t.Error("expected error for missing directory")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveStart(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
file := filepath.Join(dir, "schema.dbml")
|
||||||
|
touch(t, file)
|
||||||
|
wd, _ := os.Getwd()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
in string
|
||||||
|
wantDir string
|
||||||
|
wantFileName string
|
||||||
|
}{
|
||||||
|
{"existing file", file, dir, "schema.dbml"},
|
||||||
|
{"directory", dir, dir, ""},
|
||||||
|
{"new file in existing dir", filepath.Join(dir, "new.dbml"), dir, "new.dbml"},
|
||||||
|
{"empty", "", wd, ""},
|
||||||
|
{"nonexistent parent", filepath.Join(dir, "no", "such", "f.dbml"), wd, ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
d, n := ResolveStart(tt.in)
|
||||||
|
if d != tt.wantDir || n != tt.wantFileName {
|
||||||
|
t.Errorf("got (%q, %q), want (%q, %q)", d, n, tt.wantDir, tt.wantFileName)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatExtensions(t *testing.T) {
|
||||||
|
if got := FormatExtensions("yaml"); !reflect.DeepEqual(got, []string{".yaml", ".yml"}) {
|
||||||
|
t.Errorf("yaml: %v", got)
|
||||||
|
}
|
||||||
|
if FormatExtensions("unknown") != nil {
|
||||||
|
t.Error("unknown format should not filter")
|
||||||
|
}
|
||||||
|
if !IsDirectoryFormat("gorm") || IsDirectoryFormat("json") {
|
||||||
|
t.Error("directory format detection wrong")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newDialogTestEditor() *SchemaEditor {
|
||||||
|
se := &SchemaEditor{app: tview.NewApplication(), pages: tview.NewPages()}
|
||||||
|
se.pages.AddPage("origin", tview.NewBox(), true, true)
|
||||||
|
return se
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFileBrowserOpensOnEachMode(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
touch(t, filepath.Join(dir, "a.dbml"))
|
||||||
|
|
||||||
|
for _, mode := range []FileBrowserMode{FileBrowserLoad, FileBrowserSave} {
|
||||||
|
se := newDialogTestEditor()
|
||||||
|
se.showFileBrowser(FileBrowserConfig{
|
||||||
|
Mode: mode,
|
||||||
|
StartPath: filepath.Join(dir, "a.dbml"),
|
||||||
|
Extensions: FormatExtensions("dbml"),
|
||||||
|
ReturnPage: "origin",
|
||||||
|
OnSelect: func(string) { t.Error("OnSelect must not fire without a selection") },
|
||||||
|
})
|
||||||
|
if !se.pages.HasPage("file-browser") {
|
||||||
|
t.Errorf("mode %d: file-browser page missing", mode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConnStringBuilderOpensForEachKind(t *testing.T) {
|
||||||
|
for _, in := range []string{
|
||||||
|
"",
|
||||||
|
"postgres://u:pw@db:5432/app?sslmode=disable",
|
||||||
|
"sqlserver://sa:pw@sql:1433?database=shop&encrypt=disable",
|
||||||
|
"/tmp/app.db",
|
||||||
|
"postgres://u:p@host:badport/db", // parse error falls back to defaults
|
||||||
|
} {
|
||||||
|
se := newDialogTestEditor()
|
||||||
|
se.showConnStringBuilder(in, ConnPostgres, "origin", func(string) {
|
||||||
|
t.Error("onDone must not fire without Save")
|
||||||
|
})
|
||||||
|
if !se.pages.HasPage(connBuilderPage) {
|
||||||
|
t.Errorf("%q: builder page missing", in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateConnFields(t *testing.T) {
|
||||||
|
if validateConnFields(ConnFields{Kind: ConnSQLite}) == "" {
|
||||||
|
t.Error("sqlite without path should be invalid")
|
||||||
|
}
|
||||||
|
if validateConnFields(ConnFields{Kind: ConnPostgres}) == "" {
|
||||||
|
t.Error("postgres without host should be invalid")
|
||||||
|
}
|
||||||
|
if msg := validateConnFields(DefaultConnFields(ConnMSSQL)); msg != "" {
|
||||||
|
t.Errorf("defaults should be valid, got %q", msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -92,6 +92,9 @@ func (se *SchemaEditor) showLoadScreen() {
|
|||||||
connString = value
|
connString = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "load-database", FileBrowserLoad, func() string { return currentFormat })
|
||||||
|
se.attachConnStringBuilder(form, "Connection String", "load-database", func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddTextView("Help", getLoadHelpText(), 0, 5, true, false)
|
form.AddTextView("Help", getLoadHelpText(), 0, 5, true, false)
|
||||||
|
|
||||||
// Buttons
|
// Buttons
|
||||||
@@ -190,6 +193,8 @@ func (se *SchemaEditor) showSaveScreen() {
|
|||||||
filePath = value
|
filePath = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "save-database", FileBrowserSave, func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddTextView("Help", getSaveHelpText(), 0, 5, true, false)
|
form.AddTextView("Help", getSaveHelpText(), 0, 5, true, false)
|
||||||
|
|
||||||
// Buttons
|
// Buttons
|
||||||
@@ -469,6 +474,8 @@ func getLoadHelpText() string {
|
|||||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm
|
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm
|
||||||
Database formats: pgsql (requires connection string)
|
Database formats: pgsql (requires connection string)
|
||||||
|
|
||||||
|
Press Enter in File Path to browse files, or in Connection String to open the builder.
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
- File path: ~/schemas/mydb.dbml or /path/to/schema.json
|
- File path: ~/schemas/mydb.dbml or /path/to/schema.json
|
||||||
- Connection: postgres://user:pass@localhost/dbname`
|
- Connection: postgres://user:pass@localhost/dbname`
|
||||||
@@ -520,6 +527,8 @@ func (se *SchemaEditor) showUpdateExistingDatabaseConfirm() {
|
|||||||
func getSaveHelpText() string {
|
func getSaveHelpText() string {
|
||||||
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql (SQL export)
|
return `File-based formats: dbml, dctx, drawdb, graphql, json, yaml, gorm, bun, drizzle, prisma, typeorm, pgsql (SQL export)
|
||||||
|
|
||||||
|
Press Enter in File Path to browse for a target.
|
||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
- File: ~/schemas/mydb.dbml
|
- File: ~/schemas/mydb.dbml
|
||||||
- Directory (for code formats): ./models/`
|
- Directory (for code formats): ./models/`
|
||||||
@@ -570,6 +579,9 @@ func (se *SchemaEditor) showImportScreen() {
|
|||||||
connString = value
|
connString = value
|
||||||
})
|
})
|
||||||
|
|
||||||
|
se.attachFileBrowser(form, "File Path", "import-database", FileBrowserLoad, func() string { return currentFormat })
|
||||||
|
se.attachConnStringBuilder(form, "Connection String", "import-database", func() string { return currentFormat })
|
||||||
|
|
||||||
form.AddInputField("Skip Tables (comma-separated)", "", 50, nil, func(value string) {
|
form.AddInputField("Skip Tables (comma-separated)", "", 50, nil, func(value string) {
|
||||||
skipTables = value
|
skipTables = value
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -39,6 +39,18 @@ func (se *SchemaEditor) createMainMenu() tview.Primitive {
|
|||||||
AddItem("Manage Domains", "View, create, edit, and delete domains", 'd', func() {
|
AddItem("Manage Domains", "View, create, edit, and delete domains", 'd', func() {
|
||||||
se.showDomainList()
|
se.showDomainList()
|
||||||
}).
|
}).
|
||||||
|
AddItem("Manage Indexes", "View, create, edit, and delete table indexes", 'x', func() {
|
||||||
|
se.showObjectList(se.indexKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Views", "View, create, edit, and delete views", 'v', func() {
|
||||||
|
se.showObjectList(se.viewKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Sequences", "View, create, edit, and delete sequences", 'u', func() {
|
||||||
|
se.showObjectList(se.sequenceKind())
|
||||||
|
}).
|
||||||
|
AddItem("Manage Scripts", "View, create, edit, and delete SQL scripts", 'c', func() {
|
||||||
|
se.showObjectList(se.scriptKind())
|
||||||
|
}).
|
||||||
AddItem("Import & Merge", "Import and merge schema from another database", 'i', func() {
|
AddItem("Import & Merge", "Import and merge schema from another database", 'i', func() {
|
||||||
se.showImportScreen()
|
se.showImportScreen()
|
||||||
}).
|
}).
|
||||||
|
|||||||
@@ -0,0 +1,263 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Data operations for indexes, views, sequences, scripts and domain/table assignment.
|
||||||
|
|
||||||
|
func (se *SchemaEditor) schemaAt(schemaIndex int) (*models.Schema, error) {
|
||||||
|
if schemaIndex < 0 || schemaIndex >= len(se.db.Schemas) {
|
||||||
|
return nil, errors.New("schema not found")
|
||||||
|
}
|
||||||
|
return se.db.Schemas[schemaIndex], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) tableAt(schemaIndex, tableIndex int) (*models.Schema, *models.Table, error) {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
if tableIndex < 0 || tableIndex >= len(schema.Tables) {
|
||||||
|
return nil, nil, errors.New("table not found")
|
||||||
|
}
|
||||||
|
return schema, schema.Tables[tableIndex], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// splitList splits a comma separated list, trimming blanks and dropping empty entries.
|
||||||
|
func splitList(s string) []string {
|
||||||
|
parts := make([]string, 0)
|
||||||
|
for _, p := range strings.Split(s, ",") {
|
||||||
|
if p = strings.TrimSpace(p); p != "" {
|
||||||
|
parts = append(parts, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return parts
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveIndex adds an index to a table. When oldName is non-empty the index of that
|
||||||
|
// name is replaced (and renamed if needed).
|
||||||
|
func (se *SchemaEditor) SaveIndex(schemaIndex, tableIndex int, oldName string, idx *models.Index) error {
|
||||||
|
schema, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
idx.Name = strings.TrimSpace(idx.Name)
|
||||||
|
if idx.Name == "" {
|
||||||
|
return errors.New("index name is required")
|
||||||
|
}
|
||||||
|
if len(idx.Columns) == 0 {
|
||||||
|
return errors.New("index needs at least one column")
|
||||||
|
}
|
||||||
|
for _, c := range idx.Columns {
|
||||||
|
if _, ok := table.Columns[c]; !ok {
|
||||||
|
return fmt.Errorf("column %q not found in table %s", c, table.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, exists := table.Indexes[idx.Name]; exists && idx.Name != oldName {
|
||||||
|
return fmt.Errorf("index %q already exists", idx.Name)
|
||||||
|
}
|
||||||
|
if table.Indexes == nil {
|
||||||
|
table.Indexes = make(map[string]*models.Index)
|
||||||
|
}
|
||||||
|
if oldName != "" {
|
||||||
|
delete(table.Indexes, oldName)
|
||||||
|
}
|
||||||
|
idx.Table = table.Name
|
||||||
|
idx.Schema = schema.Name
|
||||||
|
table.Indexes[idx.Name] = idx
|
||||||
|
table.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteIndex removes an index from a table.
|
||||||
|
func (se *SchemaEditor) DeleteIndex(schemaIndex, tableIndex int, name string) bool {
|
||||||
|
_, table, err := se.tableAt(schemaIndex, tableIndex)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if _, ok := table.Indexes[name]; !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
delete(table.Indexes, name)
|
||||||
|
table.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveView adds a view to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveView(schemaIndex, at int, v *models.View) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
v.Name = strings.TrimSpace(v.Name)
|
||||||
|
if v.Name == "" {
|
||||||
|
return errors.New("view name is required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(v.Definition) == "" {
|
||||||
|
return errors.New("view definition is required")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Views {
|
||||||
|
if i != at && o.Name == v.Name {
|
||||||
|
return fmt.Errorf("view %q already exists", v.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
v.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Views) {
|
||||||
|
schema.Views[at] = v
|
||||||
|
} else {
|
||||||
|
schema.Views = append(schema.Views, v)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteView removes the view at position at.
|
||||||
|
func (se *SchemaEditor) DeleteView(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Views) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Views = append(schema.Views[:at], schema.Views[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveSequence adds a sequence to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveSequence(schemaIndex, at int, s *models.Sequence) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.Name = strings.TrimSpace(s.Name)
|
||||||
|
if s.Name == "" {
|
||||||
|
return errors.New("sequence name is required")
|
||||||
|
}
|
||||||
|
if s.IncrementBy == 0 {
|
||||||
|
return errors.New("increment must not be zero")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Sequences {
|
||||||
|
if i != at && o.Name == s.Name {
|
||||||
|
return fmt.Errorf("sequence %q already exists", s.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Sequences) {
|
||||||
|
schema.Sequences[at] = s
|
||||||
|
} else {
|
||||||
|
schema.Sequences = append(schema.Sequences, s)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteSequence removes the sequence at position at.
|
||||||
|
func (se *SchemaEditor) DeleteSequence(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Sequences) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Sequences = append(schema.Sequences[:at], schema.Sequences[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveScript adds a script to a schema, or replaces the one at position at (use -1 to add).
|
||||||
|
func (se *SchemaEditor) SaveScript(schemaIndex, at int, s *models.Script) error {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.Name = strings.TrimSpace(s.Name)
|
||||||
|
if s.Name == "" {
|
||||||
|
return errors.New("script name is required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(s.SQL) == "" {
|
||||||
|
return errors.New("script SQL is required")
|
||||||
|
}
|
||||||
|
for i, o := range schema.Scripts {
|
||||||
|
if i != at && o.Name == s.Name {
|
||||||
|
return fmt.Errorf("script %q already exists", s.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.Schema = schema.Name
|
||||||
|
if at >= 0 && at < len(schema.Scripts) {
|
||||||
|
schema.Scripts[at] = s
|
||||||
|
} else {
|
||||||
|
schema.Scripts = append(schema.Scripts, s)
|
||||||
|
}
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteScript removes the script at position at.
|
||||||
|
func (se *SchemaEditor) DeleteScript(schemaIndex, at int) bool {
|
||||||
|
schema, err := se.schemaAt(schemaIndex)
|
||||||
|
if err != nil || at < 0 || at >= len(schema.Scripts) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
schema.Scripts = append(schema.Scripts[:at], schema.Scripts[at+1:]...)
|
||||||
|
schema.UpdateDate()
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// AssignTableToDomain adds a reference to schemaName.tableName to the domain at domainIndex.
|
||||||
|
func (se *SchemaEditor) AssignTableToDomain(domainIndex int, schemaName, tableName string) error {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return errors.New("domain not found")
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
var table *models.Table
|
||||||
|
for _, s := range se.db.Schemas {
|
||||||
|
if s.Name != schemaName {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
if t.Name == tableName {
|
||||||
|
table = t
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if table == nil {
|
||||||
|
return fmt.Errorf("table %s.%s not found", schemaName, tableName)
|
||||||
|
}
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||||
|
return fmt.Errorf("table %s.%s is already in domain %s", schemaName, tableName, domain.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
dt := models.InitDomainTable(tableName, schemaName)
|
||||||
|
dt.RefTable = table
|
||||||
|
dt.Sequence = uint(len(domain.Tables))
|
||||||
|
domain.Tables = append(domain.Tables, dt)
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UnassignTableFromDomain removes the reference to schemaName.tableName from the domain.
|
||||||
|
func (se *SchemaEditor) UnassignTableFromDomain(domainIndex int, schemaName, tableName string) bool {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
for i, dt := range domain.Tables {
|
||||||
|
if dt.SchemaName == schemaName && dt.TableName == tableName {
|
||||||
|
domain.Tables = append(domain.Tables[:i], domain.Tables[i+1:]...)
|
||||||
|
se.db.UpdateDate()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestEditor() *SchemaEditor {
|
||||||
|
db := models.InitDatabase("test")
|
||||||
|
schema := models.InitSchema("public")
|
||||||
|
table := models.InitTable("users", "public")
|
||||||
|
table.Columns["id"] = models.InitColumn("id", "users", "public")
|
||||||
|
table.Columns["email"] = models.InitColumn("email", "users", "public")
|
||||||
|
schema.Tables = append(schema.Tables, table)
|
||||||
|
db.Schemas = append(db.Schemas, schema)
|
||||||
|
return &SchemaEditor{db: db}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSaveIndex(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
table := se.db.Schemas[0].Tables[0]
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
old string
|
||||||
|
idx *models.Index
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"valid", "", &models.Index{Name: "idx_email", Columns: []string{"email"}, Unique: true}, false},
|
||||||
|
{"duplicate", "", &models.Index{Name: "idx_email", Columns: []string{"email"}}, true},
|
||||||
|
{"missing name", "", &models.Index{Columns: []string{"email"}}, true},
|
||||||
|
{"no columns", "", &models.Index{Name: "idx_none"}, true},
|
||||||
|
{"unknown column", "", &models.Index{Name: "idx_bad", Columns: []string{"nope"}}, true},
|
||||||
|
{"rename", "idx_email", &models.Index{Name: "idx_email2", Columns: []string{"email", "id"}}, false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if err := se.SaveIndex(0, 0, tt.old, tt.idx); (err != nil) != tt.wantErr {
|
||||||
|
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if _, ok := table.Indexes["idx_email"]; ok {
|
||||||
|
t.Error("renamed index should be gone under old name")
|
||||||
|
}
|
||||||
|
if idx := table.Indexes["idx_email2"]; idx == nil || idx.Table != "users" || idx.Schema != "public" {
|
||||||
|
t.Errorf("unexpected renamed index: %+v", idx)
|
||||||
|
}
|
||||||
|
if !se.DeleteIndex(0, 0, "idx_email2") || se.DeleteIndex(0, 0, "idx_email2") {
|
||||||
|
t.Error("delete should succeed once")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSaveViewSequenceScript(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
schema := se.db.Schemas[0]
|
||||||
|
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "v", Definition: "select 2"}); err == nil {
|
||||||
|
t.Error("duplicate view accepted")
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, 0, &models.View{Name: "v", Definition: "select 3"}); err != nil {
|
||||||
|
t.Errorf("editing in place should not conflict: %v", err)
|
||||||
|
}
|
||||||
|
if err := se.SaveView(0, -1, &models.View{Name: "w"}); err == nil {
|
||||||
|
t.Error("view without definition accepted")
|
||||||
|
}
|
||||||
|
if len(schema.Views) != 1 || schema.Views[0].Definition != "select 3" || schema.Views[0].Schema != "public" {
|
||||||
|
t.Errorf("unexpected views: %+v", schema.Views)
|
||||||
|
}
|
||||||
|
if !se.DeleteView(0, 0) || se.DeleteView(0, 0) {
|
||||||
|
t.Error("view delete mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "s", IncrementBy: 1, StartValue: 1}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveSequence(0, -1, &models.Sequence{Name: "z"}); err == nil {
|
||||||
|
t.Error("zero increment accepted")
|
||||||
|
}
|
||||||
|
if !se.DeleteSequence(0, 0) || len(schema.Sequences) != 0 {
|
||||||
|
t.Error("sequence delete failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := se.SaveScript(0, -1, &models.Script{Name: "init", SQL: "select 1"}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.SaveScript(0, -1, &models.Script{Name: "empty"}); err == nil {
|
||||||
|
t.Error("script without SQL accepted")
|
||||||
|
}
|
||||||
|
if err := se.SaveScript(5, -1, &models.Script{Name: "x", SQL: "y"}); err == nil {
|
||||||
|
t.Error("bad schema index accepted")
|
||||||
|
}
|
||||||
|
if !se.DeleteScript(0, 0) || len(schema.Scripts) != 0 {
|
||||||
|
t.Error("script delete failed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDomainTableAssignment(t *testing.T) {
|
||||||
|
se := newTestEditor()
|
||||||
|
se.createDomainNoUI("core")
|
||||||
|
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "users"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "users"); err == nil {
|
||||||
|
t.Error("duplicate assignment accepted")
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(0, "public", "missing"); err == nil {
|
||||||
|
t.Error("unknown table accepted")
|
||||||
|
}
|
||||||
|
if err := se.AssignTableToDomain(3, "public", "users"); err == nil {
|
||||||
|
t.Error("bad domain index accepted")
|
||||||
|
}
|
||||||
|
dt := se.db.Domains[0].Tables[0]
|
||||||
|
if dt.RefTable != se.db.Schemas[0].Tables[0] {
|
||||||
|
t.Error("RefTable not linked")
|
||||||
|
}
|
||||||
|
if !se.UnassignTableFromDomain(0, "public", "users") || se.UnassignTableFromDomain(0, "public", "users") {
|
||||||
|
t.Error("unassign mismatch")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) createDomainNoUI(name string) {
|
||||||
|
se.db.Domains = append(se.db.Domains, models.InitDomain(name))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSplitList(t *testing.T) {
|
||||||
|
got := splitList(" a, b,, c ,")
|
||||||
|
if len(got) != 3 || got[0] != "a" || got[2] != "c" {
|
||||||
|
t.Errorf("got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,476 @@
|
|||||||
|
package ui
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gdamore/tcell/v2"
|
||||||
|
"github.com/rivo/tview"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
// objectLocation identifies where a new object is created: a schema, and for indexes also a table.
|
||||||
|
type objectLocation struct {
|
||||||
|
label string
|
||||||
|
schemaIndex int
|
||||||
|
tableIndex int
|
||||||
|
}
|
||||||
|
|
||||||
|
// objectRow is one existing object shown in an object list.
|
||||||
|
type objectRow struct {
|
||||||
|
cells []string
|
||||||
|
schemaIndex int
|
||||||
|
tableIndex int
|
||||||
|
at int // position within the schema slice (views, sequences, scripts)
|
||||||
|
name string // map key (indexes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// objectKind describes how a kind of schema object is listed and edited.
|
||||||
|
type objectKind struct {
|
||||||
|
page string
|
||||||
|
title string
|
||||||
|
singular string
|
||||||
|
headers []string
|
||||||
|
rows func() []objectRow
|
||||||
|
locations func() []objectLocation
|
||||||
|
// buildForm adds the editable fields to the form for row (nil when creating) and
|
||||||
|
// returns a function that validates and saves the values at the given location.
|
||||||
|
buildForm func(form *tview.Form, row *objectRow) func(loc objectLocation) error
|
||||||
|
remove func(row objectRow) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) schemaLocations() []objectLocation {
|
||||||
|
locs := make([]objectLocation, 0, len(se.db.Schemas))
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
locs = append(locs, objectLocation{label: s.Name, schemaIndex: si, tableIndex: -1})
|
||||||
|
}
|
||||||
|
return locs
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) tableLocations() []objectLocation {
|
||||||
|
locs := make([]objectLocation, 0)
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for ti, t := range s.Tables {
|
||||||
|
locs = append(locs, objectLocation{label: s.Name + "." + t.Name, schemaIndex: si, tableIndex: ti})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return locs
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) indexKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "indexes",
|
||||||
|
title: "Manage Indexes",
|
||||||
|
singular: "Index",
|
||||||
|
headers: []string{"Name", "Schema", "Table", "Type", "Unique", "Columns"},
|
||||||
|
locations: se.tableLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for ti, t := range s.Tables {
|
||||||
|
for _, name := range sortedKeys(t.Indexes) {
|
||||||
|
idx := t.Indexes[name]
|
||||||
|
rows = append(rows, objectRow{
|
||||||
|
cells: []string{idx.Name, s.Name, t.Name, idx.Type, strconv.FormatBool(idx.Unique), strings.Join(idx.Columns, ",")},
|
||||||
|
schemaIndex: si, tableIndex: ti, name: name,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
idx := models.InitIndex("", "", "")
|
||||||
|
idx.Type = "btree"
|
||||||
|
if row != nil {
|
||||||
|
idx = se.db.Schemas[row.schemaIndex].Tables[row.tableIndex].Indexes[row.name]
|
||||||
|
}
|
||||||
|
name, columns, typ, where := idx.Name, strings.Join(idx.Columns, ", "), idx.Type, idx.Where
|
||||||
|
unique := idx.Unique
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Columns (comma separated)", columns, 50, nil, func(v string) { columns = v })
|
||||||
|
form.AddInputField("Type", typ, 20, nil, func(v string) { typ = v })
|
||||||
|
form.AddCheckbox("Unique", unique, func(v bool) { unique = v })
|
||||||
|
form.AddInputField("Where", where, 50, nil, func(v string) { where = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
oldName := ""
|
||||||
|
if row != nil {
|
||||||
|
oldName = row.name
|
||||||
|
}
|
||||||
|
next := *idx
|
||||||
|
next.Name, next.Columns, next.Type, next.Unique, next.Where = name, splitList(columns), typ, unique, where
|
||||||
|
return se.SaveIndex(loc.schemaIndex, loc.tableIndex, oldName, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteIndex(r.schemaIndex, r.tableIndex, r.name) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) viewKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "views",
|
||||||
|
title: "Manage Views",
|
||||||
|
singular: "View",
|
||||||
|
headers: []string{"Name", "Schema", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, v := range s.Views {
|
||||||
|
rows = append(rows, objectRow{cells: []string{v.Name, s.Name, v.Description}, schemaIndex: si, at: i})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
view := models.InitView("", "")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
view, at = se.db.Schemas[row.schemaIndex].Views[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc, def := view.Name, view.Description, view.Definition
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddTextArea("Definition (SQL)", def, 60, 8, 0, func(v string) { def = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
next := *view
|
||||||
|
next.Name, next.Description, next.Definition = name, desc, def
|
||||||
|
return se.SaveView(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteView(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) sequenceKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "sequences",
|
||||||
|
title: "Manage Sequences",
|
||||||
|
singular: "Sequence",
|
||||||
|
headers: []string{"Name", "Schema", "Start", "Increment", "Cycle", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, q := range s.Sequences {
|
||||||
|
rows = append(rows, objectRow{
|
||||||
|
cells: []string{q.Name, s.Name, strconv.FormatInt(q.StartValue, 10), strconv.FormatInt(q.IncrementBy, 10), strconv.FormatBool(q.Cycle), q.Description}, schemaIndex: si, at: i,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
seq := models.InitSequence("", "")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
seq, at = se.db.Schemas[row.schemaIndex].Sequences[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc := seq.Name, seq.Description
|
||||||
|
start, incr := strconv.FormatInt(seq.StartValue, 10), strconv.FormatInt(seq.IncrementBy, 10)
|
||||||
|
minV, maxV := strconv.FormatInt(seq.MinValue, 10), strconv.FormatInt(seq.MaxValue, 10)
|
||||||
|
cycle := seq.Cycle
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddInputField("Start", start, 20, nil, func(v string) { start = v })
|
||||||
|
form.AddInputField("Increment", incr, 20, nil, func(v string) { incr = v })
|
||||||
|
form.AddInputField("Min (0 = none)", minV, 20, nil, func(v string) { minV = v })
|
||||||
|
form.AddInputField("Max (0 = none)", maxV, 20, nil, func(v string) { maxV = v })
|
||||||
|
form.AddCheckbox("Cycle", cycle, func(v bool) { cycle = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
next := *seq
|
||||||
|
next.Name, next.Description, next.Cycle = name, desc, cycle
|
||||||
|
for _, f := range []struct {
|
||||||
|
label string
|
||||||
|
text string
|
||||||
|
dst *int64
|
||||||
|
}{{"start", start, &next.StartValue}, {"increment", incr, &next.IncrementBy}, {"min", minV, &next.MinValue}, {"max", maxV, &next.MaxValue}} {
|
||||||
|
n, err := strconv.ParseInt(strings.TrimSpace(f.text), 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("%s must be an integer", f.label)
|
||||||
|
}
|
||||||
|
*f.dst = n
|
||||||
|
}
|
||||||
|
return se.SaveSequence(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteSequence(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (se *SchemaEditor) scriptKind() objectKind {
|
||||||
|
return objectKind{
|
||||||
|
page: "scripts",
|
||||||
|
title: "Manage Scripts",
|
||||||
|
singular: "Script",
|
||||||
|
headers: []string{"Name", "Schema", "Version", "Priority", "Description"},
|
||||||
|
locations: se.schemaLocations,
|
||||||
|
rows: func() []objectRow {
|
||||||
|
var rows []objectRow
|
||||||
|
for si, s := range se.db.Schemas {
|
||||||
|
for i, sc := range s.Scripts {
|
||||||
|
rows = append(rows, objectRow{cells: []string{sc.Name, s.Name, sc.Version, strconv.Itoa(sc.Priority), sc.Description}, schemaIndex: si, at: i})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return rows
|
||||||
|
},
|
||||||
|
buildForm: func(form *tview.Form, row *objectRow) func(objectLocation) error {
|
||||||
|
script := models.InitScript("")
|
||||||
|
at := -1
|
||||||
|
if row != nil {
|
||||||
|
script, at = se.db.Schemas[row.schemaIndex].Scripts[row.at], row.at
|
||||||
|
}
|
||||||
|
name, desc, version, sql, rollback := script.Name, script.Description, script.Version, script.SQL, script.Rollback
|
||||||
|
priority, runAfter := strconv.Itoa(script.Priority), strings.Join(script.RunAfter, ", ")
|
||||||
|
form.AddInputField("Name", name, 40, nil, func(v string) { name = v })
|
||||||
|
form.AddInputField("Description", desc, 50, nil, func(v string) { desc = v })
|
||||||
|
form.AddInputField("Version", version, 20, nil, func(v string) { version = v })
|
||||||
|
form.AddInputField("Priority", priority, 10, nil, func(v string) { priority = v })
|
||||||
|
form.AddInputField("Run after (comma separated)", runAfter, 50, nil, func(v string) { runAfter = v })
|
||||||
|
form.AddTextArea("SQL", sql, 60, 8, 0, func(v string) { sql = v })
|
||||||
|
form.AddTextArea("Rollback SQL", rollback, 60, 4, 0, func(v string) { rollback = v })
|
||||||
|
return func(loc objectLocation) error {
|
||||||
|
prio, err := strconv.Atoi(strings.TrimSpace(priority))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("priority must be an integer")
|
||||||
|
}
|
||||||
|
next := *script
|
||||||
|
next.Name, next.Description, next.Version, next.Priority = name, desc, version, prio
|
||||||
|
next.RunAfter, next.SQL, next.Rollback = splitList(runAfter), sql, rollback
|
||||||
|
return se.SaveScript(loc.schemaIndex, at, &next)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
remove: func(r objectRow) bool { return se.DeleteScript(r.schemaIndex, r.at) },
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedKeys[V any](m map[string]V) []string {
|
||||||
|
keys := make([]string, 0, len(m))
|
||||||
|
for k := range m {
|
||||||
|
keys = append(keys, k)
|
||||||
|
}
|
||||||
|
sort.Strings(keys)
|
||||||
|
return keys
|
||||||
|
}
|
||||||
|
|
||||||
|
// showObjectList displays all objects of a kind across schemas.
|
||||||
|
func (se *SchemaEditor) showObjectList(k objectKind) {
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow)
|
||||||
|
title := tview.NewTextView().SetText("[::b]" + k.title).SetDynamicColors(true).SetTextAlign(tview.AlignCenter)
|
||||||
|
|
||||||
|
table := tview.NewTable().SetBorders(true).SetSelectable(true, false).SetFixed(1, 0)
|
||||||
|
for i, h := range k.headers {
|
||||||
|
table.SetCell(0, i, tview.NewTableCell(h).SetTextColor(tcell.ColorYellow).SetSelectable(false).SetAlign(tview.AlignLeft))
|
||||||
|
}
|
||||||
|
rows := k.rows()
|
||||||
|
for r, row := range rows {
|
||||||
|
for c, text := range row.cells {
|
||||||
|
table.SetCell(r+1, c, tview.NewTableCell(text).SetSelectable(true))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
table.SetTitle(" " + k.title[len("Manage "):] + " ").SetBorder(true).SetTitleAlign(tview.AlignLeft)
|
||||||
|
|
||||||
|
back := func() {
|
||||||
|
se.pages.SwitchToPage("main")
|
||||||
|
se.pages.RemovePage(k.page)
|
||||||
|
}
|
||||||
|
btnNew := tview.NewButton("New " + k.singular + " [n]").SetSelectedFunc(func() { se.showObjectForm(k, nil) })
|
||||||
|
btnBack := tview.NewButton("Back [b]").SetSelectedFunc(back)
|
||||||
|
btnNew.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
se.app.SetFocus(table)
|
||||||
|
return nil
|
||||||
|
case tcell.KeyTab:
|
||||||
|
se.app.SetFocus(btnBack)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
btnBack.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch event.Key() {
|
||||||
|
case tcell.KeyBacktab:
|
||||||
|
se.app.SetFocus(btnNew)
|
||||||
|
return nil
|
||||||
|
case tcell.KeyTab:
|
||||||
|
se.app.SetFocus(table)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
btnFlex := tview.NewFlex().AddItem(btnNew, 0, 1, true).AddItem(btnBack, 0, 1, false)
|
||||||
|
|
||||||
|
table.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
switch {
|
||||||
|
case event.Key() == tcell.KeyEscape, event.Rune() == 'b':
|
||||||
|
back()
|
||||||
|
return nil
|
||||||
|
case event.Key() == tcell.KeyTab:
|
||||||
|
se.app.SetFocus(btnNew)
|
||||||
|
return nil
|
||||||
|
case event.Key() == tcell.KeyEnter:
|
||||||
|
if row, _ := table.GetSelection(); row > 0 && row <= len(rows) {
|
||||||
|
se.showObjectForm(k, &rows[row-1])
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
case event.Rune() == 'n':
|
||||||
|
se.showObjectForm(k, nil)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
flex.AddItem(title, 1, 0, false).AddItem(table, 0, 1, true).AddItem(btnFlex, 1, 0, false)
|
||||||
|
se.pages.AddPage(k.page, flex, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showObjectForm shows the create (row == nil) or edit form for an object.
|
||||||
|
func (se *SchemaEditor) showObjectForm(k objectKind, row *objectRow) {
|
||||||
|
formPage := k.page + "-form"
|
||||||
|
form := tview.NewForm()
|
||||||
|
errView := tview.NewTextView().SetDynamicColors(true)
|
||||||
|
|
||||||
|
locs := k.locations()
|
||||||
|
loc := objectLocation{schemaIndex: -1, tableIndex: -1}
|
||||||
|
switch {
|
||||||
|
case row != nil:
|
||||||
|
loc = objectLocation{schemaIndex: row.schemaIndex, tableIndex: row.tableIndex}
|
||||||
|
case len(locs) > 0:
|
||||||
|
loc = locs[0]
|
||||||
|
labels := make([]string, len(locs))
|
||||||
|
for i, l := range locs {
|
||||||
|
labels[i] = l.label
|
||||||
|
}
|
||||||
|
form.AddDropDown("Location", labels, 0, func(_ string, i int) { loc = locs[i] })
|
||||||
|
}
|
||||||
|
|
||||||
|
save := k.buildForm(form, row)
|
||||||
|
|
||||||
|
closeForm := func() {
|
||||||
|
se.pages.RemovePage(formPage)
|
||||||
|
se.pages.RemovePage(k.page)
|
||||||
|
se.showObjectList(k)
|
||||||
|
}
|
||||||
|
form.AddButton("Save", func() {
|
||||||
|
if err := save(loc); err != nil {
|
||||||
|
errView.SetText("[red]" + tview.Escape(err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
closeForm()
|
||||||
|
})
|
||||||
|
if row != nil {
|
||||||
|
form.AddButton("Delete", func() {
|
||||||
|
modal := tview.NewModal().
|
||||||
|
SetText(fmt.Sprintf("Delete %s '%s'? This action cannot be undone.", strings.ToLower(k.singular), row.cells[0])).
|
||||||
|
AddButtons([]string{"Cancel", "Delete"}).
|
||||||
|
SetDoneFunc(func(_ int, label string) {
|
||||||
|
se.pages.RemovePage(formPage + "-delete")
|
||||||
|
if label == "Delete" {
|
||||||
|
k.remove(*row)
|
||||||
|
closeForm()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
se.pages.AddAndSwitchToPage(formPage+"-delete", modal, true)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
form.AddButton("Back", closeForm)
|
||||||
|
|
||||||
|
verb := "New"
|
||||||
|
if row != nil {
|
||||||
|
verb = "Edit"
|
||||||
|
}
|
||||||
|
form.SetBorder(true).SetTitle(" " + verb + " " + k.singular + " ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.showExitConfirmation(formPage, k.page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(locs) == 0 && row == nil {
|
||||||
|
errView.SetText("[red]No schema/table available. Create one first.")
|
||||||
|
}
|
||||||
|
flex := tview.NewFlex().SetDirection(tview.FlexRow).AddItem(form, 0, 1, true).AddItem(errView, 1, 0, false)
|
||||||
|
se.pages.AddPage(formPage, flex, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showDomainTables lists the tables assigned to a domain and allows assigning/unassigning.
|
||||||
|
func (se *SchemaEditor) showDomainTables(domainIndex int) {
|
||||||
|
if domainIndex < 0 || domainIndex >= len(se.db.Domains) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
page := "domain-tables"
|
||||||
|
list := tview.NewList().ShowSecondaryText(true)
|
||||||
|
refresh := func() {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
se.showDomainTables(domainIndex)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
dt := dt
|
||||||
|
list.AddItem(dt.SchemaName+"."+dt.TableName, "Enter to remove from domain", 0, func() {
|
||||||
|
se.UnassignTableFromDomain(domainIndex, dt.SchemaName, dt.TableName)
|
||||||
|
refresh()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
list.AddItem("[Assign Table]", "Add a table to this domain", 'a', func() {
|
||||||
|
se.showAssignDomainTable(domainIndex, refresh)
|
||||||
|
})
|
||||||
|
list.AddItem("[Back]", "Return to domain", 'b', func() {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
})
|
||||||
|
list.SetBorder(true).SetTitle(" Domain " + domain.Name + " - Tables ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
list.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddPage(page, list, true, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// showAssignDomainTable shows a form to pick a table not yet in the domain.
|
||||||
|
func (se *SchemaEditor) showAssignDomainTable(domainIndex int, done func()) {
|
||||||
|
page := "assign-domain-table"
|
||||||
|
domain := se.db.Domains[domainIndex]
|
||||||
|
var options []string
|
||||||
|
var refs []models.DomainTable
|
||||||
|
for _, s := range se.db.Schemas {
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
taken := false
|
||||||
|
for _, dt := range domain.Tables {
|
||||||
|
taken = taken || (dt.SchemaName == s.Name && dt.TableName == t.Name)
|
||||||
|
}
|
||||||
|
if !taken {
|
||||||
|
options = append(options, s.Name+"."+t.Name)
|
||||||
|
refs = append(refs, models.DomainTable{SchemaName: s.Name, TableName: t.Name})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
form := tview.NewForm()
|
||||||
|
selected := 0
|
||||||
|
form.AddDropDown("Table", options, 0, func(_ string, i int) { selected = i })
|
||||||
|
form.AddButton("Assign", func() {
|
||||||
|
if len(refs) > 0 {
|
||||||
|
_ = se.AssignTableToDomain(domainIndex, refs[selected].SchemaName, refs[selected].TableName)
|
||||||
|
}
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
done()
|
||||||
|
})
|
||||||
|
form.AddButton("Back", func() { se.pages.RemovePage(page) })
|
||||||
|
form.SetBorder(true).SetTitle(" Assign Table ").SetTitleAlign(tview.AlignLeft)
|
||||||
|
form.SetInputCapture(func(event *tcell.EventKey) *tcell.EventKey {
|
||||||
|
if event.Key() == tcell.KeyEscape {
|
||||||
|
se.pages.RemovePage(page)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return event
|
||||||
|
})
|
||||||
|
se.pages.AddPage(page, form, true, true)
|
||||||
|
}
|
||||||
@@ -205,12 +205,24 @@ Organize UI code into these files:
|
|||||||
- **column_screens.go** - Column editor, new column dialog
|
- **column_screens.go** - Column editor, new column dialog
|
||||||
- **domain_screens.go** - Domain list, domain editor, new/edit domain dialogs
|
- **domain_screens.go** - Domain list, domain editor, new/edit domain dialogs
|
||||||
- **dialogs.go** - Confirmation dialogs (exit, delete)
|
- **dialogs.go** - Confirmation dialogs (exit, delete)
|
||||||
|
- **filebrowser_screens.go** - File browser dialog (`file-browser`), opened with Enter on File Path inputs
|
||||||
|
- **connstring_screens.go** - Connection string builder dialog (`conn-builder`), opened with Enter on Connection String inputs
|
||||||
|
|
||||||
### Data Operations Files (Business Logic)
|
### Data Operations Files (Business Logic)
|
||||||
|
|
||||||
- **schema_dataops.go** - Schema CRUD operations (Create, Read, Update, Delete)
|
- **schema_dataops.go** - Schema CRUD operations (Create, Read, Update, Delete)
|
||||||
- **table_dataops.go** - Table CRUD operations
|
- **table_dataops.go** - Table CRUD operations
|
||||||
- **column_dataops.go** - Column CRUD operations
|
- **column_dataops.go** - Column CRUD operations
|
||||||
|
- **filebrowser.go** - Directory listing, extension filtering and start-path resolution (no tview)
|
||||||
|
- **connstring.go**, **connstring_check.go** - Connection string build/parse/mask and connection test (no tview)
|
||||||
|
|
||||||
|
### Input Dialogs
|
||||||
|
|
||||||
|
- **File browser** - Enter on a File Path input opens it. Up/Down move, Enter opens a directory or selects a file,
|
||||||
|
Backspace/`[..]` goes to the parent, `h` toggles hidden files, `f` toggles the format extension filter,
|
||||||
|
`s` selects, Esc/`b` cancels (input unchanged). Save mode adds a file name field and asks before overwriting.
|
||||||
|
- **Connection string builder** - Enter on a Connection String input opens it. Fields per type (PostgreSQL, MSSQL,
|
||||||
|
SQLite), masked password and preview, F2 save, F3 test connection, Esc cancels (input unchanged).
|
||||||
|
|
||||||
## Code Separation Rules
|
## Code Separation Rules
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types for Bun
|
// TypeMapper handles type conversions between SQL and Go types for Bun
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
|
typeMappings map[string]string
|
||||||
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
arrayNullable string // writers.NullableArraysSlice | writers.NullableArraysPointerSlice
|
||||||
}
|
}
|
||||||
@@ -37,6 +38,10 @@ func NewTypeMapper(typeStyle, arrayNullable string) *TypeMapper {
|
|||||||
|
|
||||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
||||||
|
if goType, ok := tm.overrideGoType(sqlType, notNull); ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
|
||||||
// Array columns always use a native Go slice, regardless of typeStyle.
|
// Array columns always use a native Go slice, regardless of typeStyle.
|
||||||
if pgsql.IsArrayType(sqlType) {
|
if pgsql.IsArrayType(sqlType) {
|
||||||
goType := tm.arrayGoType(tm.extractBaseType(sqlType))
|
goType := tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||||
@@ -68,6 +73,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
|||||||
return tm.bunGoType(baseType)
|
return tm.bunGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetTypeMappings installs user-configured SQL-to-Go type overrides.
|
||||||
|
func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) {
|
||||||
|
tm.typeMappings = mappings
|
||||||
|
}
|
||||||
|
|
||||||
|
// overrideGoType applies a configured override, if any, for the column type.
|
||||||
|
func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) {
|
||||||
|
if len(tm.typeMappings) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType))
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if pgsql.IsArrayType(sqlType) {
|
||||||
|
return "[]" + goType, true
|
||||||
|
}
|
||||||
|
return writers.ApplyTypeMapping(goType, notNull), true
|
||||||
|
}
|
||||||
|
|
||||||
// extractBaseType extracts the base type from a SQL type string
|
// extractBaseType extracts the base type from a SQL type string
|
||||||
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
||||||
return pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
return pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package bun
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestTypeMapper_CustomTypeMappings(t *testing.T) {
|
||||||
|
mapper := NewTypeMapper("", "")
|
||||||
|
mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"uuid", true, "uuid.UUID"},
|
||||||
|
{"uuid", false, "*uuid.UUID"},
|
||||||
|
{"numeric(10,2)", true, "decimal.Decimal"},
|
||||||
|
{"UUID[]", false, "[]uuid.UUID"},
|
||||||
|
{"bigint", true, "int64"}, // unmapped types keep defaults
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want {
|
||||||
|
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer {
|
|||||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
w.typeMapper.SetTypeMappings(options.TypeMappings)
|
||||||
|
|
||||||
// Initialize templates
|
// Initialize templates
|
||||||
tmpl, err := NewTemplates()
|
tmpl, err := NewTemplates()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
// TypeMapper handles type conversions between SQL and Go types
|
// TypeMapper handles type conversions between SQL and Go types
|
||||||
type TypeMapper struct {
|
type TypeMapper struct {
|
||||||
sqlTypesAlias string
|
sqlTypesAlias string
|
||||||
|
typeMappings map[string]string
|
||||||
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
typeStyle string // writers.NullableTypeSqlTypes | writers.NullableTypeStdlib | writers.NullableTypeBaselib
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -30,6 +31,10 @@ func NewTypeMapper(typeStyle string) *TypeMapper {
|
|||||||
|
|
||||||
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
// SQLTypeToGoType converts a SQL type to its Go equivalent.
|
||||||
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
||||||
|
if goType, ok := tm.overrideGoType(sqlType, notNull); ok {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
|
||||||
// Array types are handled separately for both styles.
|
// Array types are handled separately for both styles.
|
||||||
if pgsql.IsArrayType(sqlType) {
|
if pgsql.IsArrayType(sqlType) {
|
||||||
return tm.arrayGoType(tm.extractBaseType(sqlType))
|
return tm.arrayGoType(tm.extractBaseType(sqlType))
|
||||||
@@ -57,6 +62,26 @@ func (tm *TypeMapper) SQLTypeToGoType(sqlType string, notNull bool) string {
|
|||||||
return tm.nullableGoType(baseType)
|
return tm.nullableGoType(baseType)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetTypeMappings installs user-configured SQL-to-Go type overrides.
|
||||||
|
func (tm *TypeMapper) SetTypeMappings(mappings map[string]string) {
|
||||||
|
tm.typeMappings = mappings
|
||||||
|
}
|
||||||
|
|
||||||
|
// overrideGoType applies a configured override, if any, for the column type.
|
||||||
|
func (tm *TypeMapper) overrideGoType(sqlType string, notNull bool) (string, bool) {
|
||||||
|
if len(tm.typeMappings) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
goType, ok := writers.LookupTypeMapping(tm.typeMappings, tm.extractBaseType(sqlType))
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if pgsql.IsArrayType(sqlType) {
|
||||||
|
return "[]" + goType, true
|
||||||
|
}
|
||||||
|
return writers.ApplyTypeMapping(goType, notNull), true
|
||||||
|
}
|
||||||
|
|
||||||
// extractBaseType extracts the base type from a SQL type string
|
// extractBaseType extracts the base type from a SQL type string
|
||||||
// Examples: varchar(100) → varchar, numeric(10,2) → numeric
|
// Examples: varchar(100) → varchar, numeric(10,2) → numeric
|
||||||
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
func (tm *TypeMapper) extractBaseType(sqlType string) string {
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package gorm
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestTypeMapper_CustomTypeMappings(t *testing.T) {
|
||||||
|
mapper := NewTypeMapper("")
|
||||||
|
mapper.SetTypeMappings(map[string]string{"uuid": "uuid.UUID", "numeric": "decimal.Decimal"})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
sqlType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"uuid", true, "uuid.UUID"},
|
||||||
|
{"uuid", false, "*uuid.UUID"},
|
||||||
|
{"numeric(10,2)", true, "decimal.Decimal"},
|
||||||
|
{"UUID[]", false, "[]uuid.UUID"},
|
||||||
|
{"bigint", true, "int64"}, // unmapped types keep defaults
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := mapper.SQLTypeToGoType(tt.sqlType, tt.notNull); got != tt.want {
|
||||||
|
t.Errorf("SQLTypeToGoType(%q, %v) = %q, want %q", tt.sqlType, tt.notNull, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -28,6 +28,8 @@ func NewWriter(options *writers.WriterOptions) *Writer {
|
|||||||
config: LoadMethodConfigFromMetadata(options.Metadata),
|
config: LoadMethodConfigFromMetadata(options.Metadata),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
w.typeMapper.SetTypeMappings(options.TypeMappings)
|
||||||
|
|
||||||
// Initialize templates
|
// Initialize templates
|
||||||
tmpl, err := NewTemplates()
|
tmpl, err := NewTemplates()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,176 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
_ "github.com/go-sql-driver/mysql"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/mariadb"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Writer struct {
|
||||||
|
options *writers.WriterOptions
|
||||||
|
writer io.Writer
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewWriter(options *writers.WriterOptions) *Writer { return &Writer{options: options} }
|
||||||
|
func (w *Writer) WriteDatabase(db *models.Database) error {
|
||||||
|
if w.options == nil {
|
||||||
|
return fmt.Errorf("writer options are required")
|
||||||
|
}
|
||||||
|
if conn, ok := w.options.Metadata["connection_string"].(string); ok && conn != "" {
|
||||||
|
return w.execute(db, conn)
|
||||||
|
}
|
||||||
|
if w.writer == nil {
|
||||||
|
if w.options.OutputPath != "" {
|
||||||
|
f, err := os.Create(w.options.OutputPath)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer f.Close()
|
||||||
|
w.writer = f
|
||||||
|
} else {
|
||||||
|
w.writer = os.Stdout
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return w.writeDatabaseDDL(db)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) writeDatabaseDDL(db *models.Database) error {
|
||||||
|
if _, err := fmt.Fprintf(w.writer, "-- MySQL Database Schema\n-- Database: %s\n-- Generated by RelSpec\n\n", db.Name); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, s := range db.Schemas {
|
||||||
|
if err := w.WriteSchema(s); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) WriteSchema(s *models.Schema) error {
|
||||||
|
for _, t := range s.Tables {
|
||||||
|
if err := w.writeTable(s, t); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) WriteTable(t *models.Table) error {
|
||||||
|
if w.writer == nil {
|
||||||
|
w.writer = os.Stdout
|
||||||
|
}
|
||||||
|
return w.writeTable(nil, t)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) writeTable(s *models.Schema, t *models.Table) error {
|
||||||
|
name := t.Name
|
||||||
|
if s != nil {
|
||||||
|
name = fmt.Sprintf("%s.%s", quote(s.Name), quote(t.Name))
|
||||||
|
} else {
|
||||||
|
name = quote(name)
|
||||||
|
}
|
||||||
|
cols := make([]*models.Column, 0, len(t.Columns))
|
||||||
|
for _, c := range t.Columns {
|
||||||
|
cols = append(cols, c)
|
||||||
|
}
|
||||||
|
sort.Slice(cols, func(i, j int) bool { return cols[i].Sequence < cols[j].Sequence })
|
||||||
|
defs := []string{}
|
||||||
|
pk := []string{}
|
||||||
|
for _, c := range cols {
|
||||||
|
d := fmt.Sprintf(" %s %s", quote(c.Name), mariadb.ConvertCanonicalToMariaDB(c.Type))
|
||||||
|
if c.Length > 0 && strings.EqualFold(c.Type, "string") {
|
||||||
|
d = fmt.Sprintf(" %s VARCHAR(%d)", quote(c.Name), c.Length)
|
||||||
|
}
|
||||||
|
if c.NotNull {
|
||||||
|
d += " NOT NULL"
|
||||||
|
}
|
||||||
|
if c.AutoIncrement {
|
||||||
|
d += " AUTO_INCREMENT"
|
||||||
|
}
|
||||||
|
if c.Default != nil {
|
||||||
|
d += fmt.Sprintf(" DEFAULT %s", writers.QuoteDefaultValue(fmt.Sprint(c.Default), c.Type))
|
||||||
|
}
|
||||||
|
defs = append(defs, d)
|
||||||
|
if c.IsPrimaryKey {
|
||||||
|
pk = append(pk, quote(c.Name))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, c := range t.Constraints {
|
||||||
|
if c.Type == models.PrimaryKeyConstraint {
|
||||||
|
pk = nil
|
||||||
|
for _, n := range c.Columns {
|
||||||
|
pk = append(pk, quote(n))
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(pk) > 0 {
|
||||||
|
defs = append(defs, " PRIMARY KEY ("+strings.Join(pk, ", ")+")")
|
||||||
|
}
|
||||||
|
for _, c := range t.Constraints {
|
||||||
|
if c.Type == models.UniqueConstraint {
|
||||||
|
defs = append(defs, fmt.Sprintf(" CONSTRAINT %s UNIQUE (%s)", quote(c.Name), quoted(c.Columns)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sql := fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (\n%s\n) ENGINE=InnoDB;\n\n", name, strings.Join(defs, ",\n"))
|
||||||
|
if _, err := io.WriteString(w.writer, sql); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Writer) execute(dbm *models.Database, conn string) error {
|
||||||
|
var b strings.Builder
|
||||||
|
old := w.writer
|
||||||
|
w.writer = &b
|
||||||
|
if err := w.writeDatabaseDDL(dbm); err != nil {
|
||||||
|
w.writer = old
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
w.writer = old
|
||||||
|
db, err := sql.Open("mysql", conn)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to connect: %w", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
for _, stmt := range strings.Split(b.String(), ";\n") {
|
||||||
|
stmt = stripComments(strings.TrimSpace(stmt))
|
||||||
|
if stmt == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := db.ExecContext(context.Background(), stmt); err != nil {
|
||||||
|
return fmt.Errorf("failed to execute SQL: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func stripComments(sqlText string) string {
|
||||||
|
lines := strings.Split(sqlText, "\n")
|
||||||
|
kept := lines[:0]
|
||||||
|
for _, line := range lines {
|
||||||
|
if !strings.HasPrefix(strings.TrimSpace(line), "--") {
|
||||||
|
kept = append(kept, line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(strings.Join(kept, "\n"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func quote(s string) string { return "`" + strings.ReplaceAll(s, "`", "``") + "`" }
|
||||||
|
func quoted(xs []string) string {
|
||||||
|
out := make([]string, len(xs))
|
||||||
|
for i, x := range xs {
|
||||||
|
out[i] = quote(x)
|
||||||
|
}
|
||||||
|
return strings.Join(out, ", ")
|
||||||
|
}
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestStripComments(t *testing.T) {
|
||||||
|
got := stripComments("-- header\nCREATE TABLE `users` (\n `id` INT\n);")
|
||||||
|
if strings.HasPrefix(got, "--") || !strings.HasPrefix(got, "CREATE TABLE") {
|
||||||
|
t.Fatalf("stripComments() = %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriterGeneratesMySQLDDL(t *testing.T) {
|
||||||
|
db := models.InitDatabase("app")
|
||||||
|
s := models.InitSchema("app")
|
||||||
|
table := models.InitTable("users", "app")
|
||||||
|
table.Columns["id"] = models.InitColumn("id", "users", "app")
|
||||||
|
table.Columns["id"].Type = "int"
|
||||||
|
table.Columns["id"].IsPrimaryKey = true
|
||||||
|
table.Columns["id"].NotNull = true
|
||||||
|
table.Columns["name"] = models.InitColumn("name", "users", "app")
|
||||||
|
table.Columns["name"].Type = "string"
|
||||||
|
table.Columns["name"].Length = 80
|
||||||
|
s.Tables = append(s.Tables, table)
|
||||||
|
db.Schemas = append(db.Schemas, s)
|
||||||
|
var out bytes.Buffer
|
||||||
|
w := NewWriter(&writers.WriterOptions{Metadata: map[string]interface{}{}})
|
||||||
|
w.writer = &out
|
||||||
|
if err := w.WriteDatabase(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := out.String()
|
||||||
|
for _, want := range []string{"CREATE TABLE IF NOT EXISTS `app`.`users`", "`id` INT NOT NULL", "`name` VARCHAR(80)", "PRIMARY KEY (`id`)"} {
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Errorf("DDL missing %q:\n%s", want, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package writers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.warky.dev/wdevs/relspecgo/pkg/pgsql"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ParseTypeMappings parses "sqltype=gotype" entries (as given to --type-map)
|
||||||
|
// into a map keyed by the canonical lower-case SQL base type, so that
|
||||||
|
// "VARCHAR", "character varying" and "varchar(50)" all address one entry.
|
||||||
|
// It returns nil for empty input.
|
||||||
|
func ParseTypeMappings(entries []string) (map[string]string, error) {
|
||||||
|
if len(entries) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
out := make(map[string]string, len(entries))
|
||||||
|
for _, entry := range entries {
|
||||||
|
sqlType, goType, ok := strings.Cut(entry, "=")
|
||||||
|
sqlType, goType = strings.TrimSpace(sqlType), strings.TrimSpace(goType)
|
||||||
|
if !ok || sqlType == "" || goType == "" {
|
||||||
|
return nil, fmt.Errorf("invalid type mapping %q: expected sqltype=gotype", entry)
|
||||||
|
}
|
||||||
|
out[pgsql.CanonicalizeBaseType(pgsql.ExtractBaseTypeLower(sqlType))] = goType
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// LookupTypeMapping returns the user-configured override for baseType, which
|
||||||
|
// the caller must already have canonicalized.
|
||||||
|
func LookupTypeMapping(mappings map[string]string, baseType string) (string, bool) {
|
||||||
|
goType, ok := mappings[baseType]
|
||||||
|
return goType, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// ApplyTypeMapping wraps an overridden Go type for nullability: NOT NULL uses
|
||||||
|
// the type verbatim; nullable columns get a pointer prefix unless the type is
|
||||||
|
// already a pointer, slice, map or interface.
|
||||||
|
func ApplyTypeMapping(goType string, notNull bool) string {
|
||||||
|
if notNull || strings.HasPrefix(goType, "*") || strings.HasPrefix(goType, "[]") ||
|
||||||
|
strings.HasPrefix(goType, "map[") || goType == "any" || goType == "interface{}" {
|
||||||
|
return goType
|
||||||
|
}
|
||||||
|
return "*" + goType
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package writers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"reflect"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestParseTypeMappings(t *testing.T) {
|
||||||
|
got, err := ParseTypeMappings([]string{"UUID=uuid.UUID", " varchar(50) = MyString ", "INT4=MyInt"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
want := map[string]string{"uuid": "uuid.UUID", "varchar": "MyString", "integer": "MyInt"}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Errorf("got %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
if m, err := ParseTypeMappings(nil); m != nil || err != nil {
|
||||||
|
t.Errorf("empty input: got %v, %v", m, err)
|
||||||
|
}
|
||||||
|
for _, bad := range []string{"uuid", "=string", "uuid="} {
|
||||||
|
if _, err := ParseTypeMappings([]string{bad}); err == nil {
|
||||||
|
t.Errorf("expected error for %q", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyTypeMapping(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
goType string
|
||||||
|
notNull bool
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"uuid.UUID", true, "uuid.UUID"},
|
||||||
|
{"uuid.UUID", false, "*uuid.UUID"},
|
||||||
|
{"*uuid.UUID", false, "*uuid.UUID"},
|
||||||
|
{"[]byte", false, "[]byte"},
|
||||||
|
{"map[string]any", false, "map[string]any"},
|
||||||
|
{"any", false, "any"},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
if got := ApplyTypeMapping(tt.goType, tt.notNull); got != tt.want {
|
||||||
|
t.Errorf("ApplyTypeMapping(%q, %v) = %q, want %q", tt.goType, tt.notNull, got, tt.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -83,6 +83,13 @@ type WriterOptions struct {
|
|||||||
// SqlXxxArray wrapper types.
|
// SqlXxxArray wrapper types.
|
||||||
NullableArrays string
|
NullableArrays string
|
||||||
|
|
||||||
|
// TypeMappings overrides the SQL-to-Go type mapping of the code-generation
|
||||||
|
// writers (bun, gorm). Keys are SQL base types (aliases are canonicalized,
|
||||||
|
// see ParseTypeMappings), values are Go type expressions. Array columns
|
||||||
|
// use the override for their element type. Unmapped types keep the
|
||||||
|
// built-in defaults.
|
||||||
|
TypeMappings map[string]string
|
||||||
|
|
||||||
// Prisma7 enables Prisma 7-specific output for Prisma writers.
|
// Prisma7 enables Prisma 7-specific output for Prisma writers.
|
||||||
Prisma7 bool
|
Prisma7 bool
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
# Test Coverage Plans
|
||||||
|
|
||||||
|
Baseline: 51.0% total statements (2026-10-03, after PRs #41-#43).
|
||||||
|
Scope: pgsql, sqlexec, template, plus non-reader/writer packages. Other readers/writers deferred.
|
||||||
|
|
||||||
|
## Order
|
||||||
|
|
||||||
|
| # | Plan | Package(s) | Now |
|
||||||
|
|---|------|-----------|-----|
|
||||||
|
| 1 | [pgsql.md](pgsql.md) | readers/pgsql, writers/pgsql, pkg/pgsql | 16.0 / 74.0 / 87.8 |
|
||||||
|
| 2 | [sqlexec.md](sqlexec.md) | writers/sqlexec | 19.4 |
|
||||||
|
| 3 | [template.md](template.md) | writers/template | 8.5 |
|
||||||
|
| 4 | [models.md](models.md) | pkg/models | 20.4 |
|
||||||
|
| 5 | [cmd.md](cmd.md) | cmd/relspec, pkg/jobs | 49.3 / 72.0 |
|
||||||
|
| 6 | [ui.md](ui.md) | pkg/ui | 3.8 |
|
||||||
|
| 7 | [diff-merge.md](diff-merge.md) | pkg/diff, pkg/merge | 65.5 / 75.1 |
|
||||||
|
| 8 | [sqltypes.md](sqltypes.md) | pkg/sqltypes | 67.0 |
|
||||||
|
|
||||||
|
## Conventions
|
||||||
|
|
||||||
|
- Same package as code under test; table-driven; must pass `-race`.
|
||||||
|
- Existing data first: `tests/assets/*`, `examples/*.dbml`, `tests/postgres/init.sql`, `tests/postgres/issue21`. Generate new data only where listed under "Data needed".
|
||||||
|
- New fixtures go in `tests/assets/<format>/` or package `testdata/`.
|
||||||
|
- Live-DB tests: skip unless the env var is set (pattern in `pkg/readers/pgsql/reader_test.go`). Use `tests/dbtest/dbtest.sh` (podman/docker; see `tests/dbtest/README.md`):
|
||||||
|
- `dbtest.sh up|down <postgres|mssql|mysql|all>`
|
||||||
|
- `eval "$(dbtest.sh env postgres)"` sets `RELSPEC_TEST_PG_CONN` (mssql: `RELSPEC_TEST_MSSQL_CONN`, mysql: `RELSPEC_TEST_MYSQL_CONN`)
|
||||||
|
- `dbtest.sh test <db> [pkgs]` runs up, `go test`, down
|
||||||
|
- Fixtures for live DBs: postgres `tests/postgres/init.sql`, mssql `test_data/mssql/test_schema.sql`, mysql `tests/dbtest/init/mysql.sql`.
|
||||||
|
- Prerequisite: working container networking (currently blocked until reboot into matching kernel; `tun` module).
|
||||||
|
- Prefer pure-function tests over DB tests wherever the logic can be isolated.
|
||||||
|
- Output assertions must not depend on map order (see memory: map iteration determinism).
|
||||||
|
|
||||||
|
## Targets
|
||||||
|
|
||||||
|
| Package | Target |
|
||||||
|
|---------|--------|
|
||||||
|
| pgsql (all three) | >= 85 |
|
||||||
|
| sqlexec | >= 80 |
|
||||||
|
| template | >= 85 |
|
||||||
|
| models | >= 80 |
|
||||||
|
| cmd/relspec | >= 65 |
|
||||||
|
| jobs | >= 85 |
|
||||||
|
| ui | >= 40 (data ops/pure helpers; screens via smoke tests) |
|
||||||
|
| diff, merge | >= 85 |
|
||||||
|
| sqltypes | >= 85 |
|
||||||
|
|
||||||
|
## Verify
|
||||||
|
|
||||||
|
- `go test -race -coverprofile=c.out ./pkg/<pkg>/` then `go tool cover -func=c.out`
|
||||||
|
- `make test` before commit
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
# Plan: cmd/relspec (49.3%) and pkg/jobs (72.0%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- cmd tests: convert_from_list, diff_sqldir, dry_run (#41), job, merge_from_list, templ_from_list
|
||||||
|
- jobs: `jobs_test.go`
|
||||||
|
|
||||||
|
## cmd/relspec
|
||||||
|
|
||||||
|
| Area | Gap | Approach |
|
||||||
|
|------|-----|----------|
|
||||||
|
| convert | `readDatabaseForConvert` 22%, `writeDatabase` 36%, `validateWriteTarget` 44%, `loadExtraFields`, `getSchemaNames`, `stderrWarn` | Table-driven per format using `tests/assets`; unsupported format, missing package, bad extra-fields JSON/empty/non-bun, schema filter not found, dctx multi-schema |
|
||||||
|
| merge | `readDatabaseForMerge` 18%, `writeDatabaseForMerge` 14%, `expandPath`, `parseSkipTables`, `isMergeOutputFormat` | Table-driven formats; globs; skip-list parsing |
|
||||||
|
| diff | `runDiff`, `readDatabase`, `maskPasswordInDiff` | File-based inputs; password masking cases |
|
||||||
|
| inspect | `runInspect`, `readDatabaseForInspect`, `filterDatabaseBySchema` | File-based input; schema filter |
|
||||||
|
| scripts | `runScriptsList` | Use `pkg/readers/sqldir` fixtures; execute path live via dbtest postgres |
|
||||||
|
| assets | `runAssetsList`, `runAssetsExecute` | List against temp dir; execute live via dbtest postgres |
|
||||||
|
| edit | `runEdit`, `readDatabaseForEdit`, `writeDatabaseForEdit` | Test read/write helpers only; skip TUI loop |
|
||||||
|
| report | state dir, load/save state, token, machine id, `submitReport` | Temp HOME; submit against `httptest` server; never hit real endpoint |
|
||||||
|
| root/main | `printVersionHeader`, `hasSilentFlag` | Pure |
|
||||||
|
| dry-run | merge/split paths | Add merge dry-run and split dry-run tests (convert covered) |
|
||||||
|
|
||||||
|
## pkg/jobs
|
||||||
|
- `ResolvedLogPolicy`, `Dir`, `validateTemplInput` (0%), `validateOutput` (44%): table-driven valid/invalid job definitions.
|
||||||
|
- Existing job files: `examples/jobs`.
|
||||||
|
|
||||||
|
## Live DB cases (dbtest)
|
||||||
|
- `runScriptsExecute`, `runAssetsExecute`, job script-exec, `readDatabaseForConvert/Merge/Inspect` for pgsql, and pgsql merge/convert output: `tests/dbtest/dbtest.sh test postgres ./cmd/relspec/`.
|
||||||
|
- mssql source reads (convert/inspect): `dbtest.sh up mssql`, env `RELSPEC_TEST_MSSQL_CONN`; fixture `test_data/mssql/test_schema.sql`.
|
||||||
|
- Skip when env var unset.
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
# Plan: pkg/diff (65.5%) and pkg/merge (75.1%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- diff: `diff_test.go`, `formatters_test.go`
|
||||||
|
- merge: `merge_test.go`
|
||||||
|
|
||||||
|
## pkg/diff
|
||||||
|
|
||||||
|
| Func | Now | Cases |
|
||||||
|
|------|-----|-------|
|
||||||
|
| `compareSchemaDetails` | 0% | Description/owner/options changed |
|
||||||
|
| `compareConstraintDetails`, `normalizeConstraintAction` | 0% | Columns, referenced table, on-update/on-delete variants and case/default normalisation |
|
||||||
|
| `compareRelationshipDetails` | 0% | Changed endpoints/type |
|
||||||
|
| `compareViews`, `compareViewDetails` | 0% | Added/removed/changed definition |
|
||||||
|
| `compareSequences`, `compareSequenceDetails` | 0% | Added/removed/changed increment/min/max/start |
|
||||||
|
|
||||||
|
Data: pair `examples/test_schema.dbml` and `test_schema_modified.dbml`; add view/sequence changes in code-built fixtures.
|
||||||
|
|
||||||
|
## pkg/merge
|
||||||
|
|
||||||
|
| Func | Now | Cases |
|
||||||
|
|------|-----|-------|
|
||||||
|
| `mergeSequences`, `cloneSequence` | 33% / 0% | New, existing, conflicting; clone is deep |
|
||||||
|
| `cloneSchema` | 48% | Views, sequences, scripts, indexes cloned independently |
|
||||||
|
| `extractTypeParts` | 48% | Precision/scale, arrays, schema-qualified, no modifiers |
|
||||||
|
| `GetColumnTypeConflictSummary`, `min` | 0% | Limit truncation, zero conflicts |
|
||||||
|
|
||||||
|
## Live DB cases (dbtest)
|
||||||
|
- pgsql live diff/merge against a real DB: `dbtest.sh test postgres`; reuse `tests/postgres/init.sql` as the live side and `examples/test_schema*.dbml` as the desired side.
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
# Plan: pkg/models (20.4%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- `directives_test.go`
|
||||||
|
|
||||||
|
## Gaps
|
||||||
|
|
||||||
|
| File | Funcs | Cases |
|
||||||
|
|------|-------|-------|
|
||||||
|
| `models.go` | All `SQLName` methods, `UpdateDate`, `GetPrimaryKey`, `columnLess`, `GetForeignKeys` | Case handling, empty/nil maps, composite PK ordering, FK filtering |
|
||||||
|
| `models.go` | `Init*` constructors (Database, Schema, Table, Column, Index, Relation, Relationship, Constraint, Script, View, Sequence, Domain, DomainTable, Enum) | Maps/slices non-nil, name set, defaults |
|
||||||
|
| `sorting.go` | 20 Sort* funcs | By name and by sequence; ties; map variants return sorted slice; input not mutated where documented |
|
||||||
|
| `flatview.go` | ToFlatColumns, ToFlatTables, ToFlatConstraints, ToFlatRelationships | Multi-schema, empty db, deterministic order |
|
||||||
|
| `summaryview.go` | ToSummary | Counts across object types |
|
||||||
|
| `directives.go` | `directiveFromAny` (22%) | Each input type branch, invalid type |
|
||||||
|
|
||||||
|
## Data needed
|
||||||
|
- One shared in-test builder for a multi-schema Database (reuse `tests/assets/dbml/complex.dbml` via reader only if no import cycle; otherwise build in code).
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
# Plan: PostgreSQL
|
||||||
|
|
||||||
|
## Tooling
|
||||||
|
- Live tests: `tests/dbtest/dbtest.sh test postgres` (defaults to readers/pgsql, writers/pgsql, writers/sqlexec) or `up postgres` + `eval "$(tests/dbtest/dbtest.sh env postgres)"`.
|
||||||
|
- Isolation: each live test creates and drops its own schema; shared fixture DB comes from `init.sql`.
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- Reader tests: `pkg/readers/pgsql/reader_test.go` (live tests skipped without `RELSPEC_TEST_PG_CONN`; pure tests: MapDataType, ParseIndexDefinition, DeriveRelationship, composite FK)
|
||||||
|
- Writer tests: diff_statements, directives, extensions, generated_column, migration_writer, serial_sequence
|
||||||
|
- Data: `tests/postgres/init.sql`, `tests/postgres/issue21`, `tests/assets/dbml/*`, `examples/test_schema*.dbml`
|
||||||
|
|
||||||
|
## readers/pgsql (16.0%)
|
||||||
|
|
||||||
|
| Item | Gap | Approach |
|
||||||
|
|------|-----|----------|
|
||||||
|
| `normalizePostgresDefault` (queries.go) | 0% | Pure; table-driven: casts, nextval, functions, quoted literals, NULL |
|
||||||
|
| `countColumns/Constraints/Indexes` | 0% | Pure; build Database fixtures |
|
||||||
|
| `ReadDatabase/ReadSchema/ReadTable` | ~0% | Live; run against `init.sql` DB; assert counts, PK/FK/unique/check/index, views, sequences, extensions |
|
||||||
|
| `query*` (11 funcs) | 0% | Covered via live ReadDatabase; add one live case per object type |
|
||||||
|
| `close` | 0% | Live; connection released after read and on error |
|
||||||
|
|
||||||
|
Data needed: extend `tests/postgres/init.sql` (loaded by dbtest on `up`; apply changes with `dbtest.sh restart postgres`) with a view, sequence, check constraint, partial index, extension, composite FK (verify what already exists first).
|
||||||
|
|
||||||
|
## writers/pgsql (74.0%)
|
||||||
|
|
||||||
|
| Item | Gap | Approach |
|
||||||
|
|------|-----|----------|
|
||||||
|
| `extractTableNameFromCreate`, `extractStatementContext`, `extractSQLStringValue`, `parseQualifiedIdent`, `firstBareIdent`, `firstIdentAfterKeyword`, `stripQuotes`, `buildStmtContext`, `detectStatementType`, `truncateStatement` | 0% | Pure; table-driven; quoted/qualified/unquoted idents, each statement type, long statements |
|
||||||
|
| `getCurrentTimestamp`, `finishReport`, `writeReport` | 0% | Report written to temp file; JSON shape, counts, failed statements |
|
||||||
|
| `executeStatements`, `executeDatabaseSQL` | 0% | Live; success, failure with continue-on-error, failure stop, report output |
|
||||||
|
| `generateLiveDiffStatements` | 28.6% | Live; empty DB, drifted DB, identical DB |
|
||||||
|
| `currentColumnHasDescription`, `ExecuteCommentColumn` | 0% | Migration writer fixtures with comments added/removed/changed |
|
||||||
|
| `template_functions.go` `filter`, `mapFunc` | 0% | Pure |
|
||||||
|
|
||||||
|
Reuse `tests/integration/failed_statements_example.txt` for failed-statement report cases. Ad-hoc SQL setup: `dbtest.sh exec postgres <file>`.
|
||||||
|
|
||||||
|
## pkg/pgsql (87.8%)
|
||||||
|
- Spot-check uncovered funcs after the above; add keyword/datatype edge cases only.
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
# Plan: writers/sqlexec (19.4%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- `writer_test.go`: constructor, nil DB, missing conn string, empty scripts, script sorting, embed directives
|
||||||
|
|
||||||
|
## Gaps
|
||||||
|
|
||||||
|
| Item | Now | Approach |
|
||||||
|
|------|-----|----------|
|
||||||
|
| `Options` | 0% | Trivial getter |
|
||||||
|
| `WriteDatabase` | 31.2% | Multi-schema; error from one schema aborts; context/connect failure |
|
||||||
|
| `executeScripts` | 0% | Live via dbtest postgres (`RELSPEC_TEST_PG_CONN`); ordering by priority/sequence, failing script reports script name, empty SQL skipped, transaction/partial-apply behaviour as implemented |
|
||||||
|
| `WriteSchema` | partial | Connection error path, success path live |
|
||||||
|
|
||||||
|
## Data needed
|
||||||
|
- Small script set (3-4 scripts, mixed priority, one failing) as fixtures; check `tests/assets` and `pkg/readers/sqldir` testdata first.
|
||||||
|
- Cleanup: each live test uses a throwaway schema and drops it.
|
||||||
|
|
||||||
|
## Decision
|
||||||
|
- Live-only via `tests/dbtest/dbtest.sh test postgres ./pkg/writers/sqlexec/`; no connection interface or mock.
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
# Plan: pkg/sqltypes (67.0%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- array types, fromstring, sql types, yaml/xml, struct json, uuid integration tests
|
||||||
|
|
||||||
|
## Gaps
|
||||||
|
|
||||||
|
| Area | Funcs | Approach |
|
||||||
|
|------|-------|----------|
|
||||||
|
| Array types | MarshalYAML/UnmarshalYAML/MarshalXML/UnmarshalXML across each array type; some `UnmarshalJSON/MarshalJSON` | One round-trip test per array type (reuse helper from `sql_types_yaml_xml_test.go`) |
|
||||||
|
| Scalar types | `Value` (3 types), `MarshalJSON/UnmarshalJSON` for date, `Int64` (31%), `Float64` (40%) | Valid, null, invalid string, overflow |
|
||||||
|
| Constructors | `SqlTimeStampNow`, `SqlDateNow`, `SqlTimeNow`, `NewSql`, `NewSqlFloat32`, `ToJSONDT` | Assert non-zero/valid and approximately now |
|
||||||
|
|
||||||
|
No data needed.
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
# Plan: writers/template (8.5%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- `writer_test.go`: deterministic table index values only
|
||||||
|
|
||||||
|
## Approach
|
||||||
|
Pure helper functions; one test file per source file, table-driven. Then render tests through the writer.
|
||||||
|
|
||||||
|
| File | Funcs | Cases |
|
||||||
|
|------|-------|-------|
|
||||||
|
| `filters.go` | FilterTables, FilterTablesByPattern, FilterColumns, FilterColumnsByType, FilterPrimaryKeys, FilterForeignKeys, FilterUniqueConstraints, FilterCheckConstraints, FilterNullable, FilterNotNull, matchPattern | Empty input, no match, glob patterns, nil maps |
|
||||||
|
| `formatters.go` | ToJSON, ToJSONPretty, ToYAML, Indent, IndentWith, Escape, EscapeQuotes, Comment, QuoteString, UnquoteString | Empty string, multiline, special chars, marshal failure |
|
||||||
|
| `loop_helpers.go` | Enumerate, Batch, Chunk, Reverse, First, Last, Skip, Take, Concat, Unique, SortBy, GroupBy, CountIf, getFieldValue, compareValues | Empty, n > len, n <= 0, non-slice input, missing field |
|
||||||
|
| `safe_access.go` | Get, GetOr, GetPath, GetPathOr, SafeIndex, SafeIndexOr, Has, HasPath, Keys, Merge, Pick, Omit, SliceContains, IndexOf, Pluck | nil, missing key, nested path, out-of-range |
|
||||||
|
| `string_helpers.go` | ToUpper, ToLower, ToCamelCase, and rest | Empty, snake/kebab/space input, unicode |
|
||||||
|
| `errors.go` | Error, Unwrap, NewTemplate{Load,Parse,Execute}Error | errors.Is/As, message contents |
|
||||||
|
| `funcmap.go` | BuildFuncMap | Every registered name resolves and is callable |
|
||||||
|
| `type_mappers.go`, `template_data.go` | check after above | |
|
||||||
|
| `writer.go` | WriteDatabase/Schema/Table, modes | Template load/parse/execute error paths; per-table, per-schema, whole-db modes; output to file vs stdout |
|
||||||
|
|
||||||
|
## Data needed
|
||||||
|
- 2-3 small template fixtures in `pkg/writers/template/testdata/` (valid, parse error, execute error).
|
||||||
|
- Schema input: reuse `tests/assets/dbml/simple.dbml` / `complex.dbml`.
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
# Plan: pkg/ui (3.8%)
|
||||||
|
|
||||||
|
## Existing
|
||||||
|
- `object_dataops_test.go` (indexes, views, sequences, scripts, domain assignment)
|
||||||
|
- Rules: `pkg/ui/ui_rules.md`
|
||||||
|
|
||||||
|
## Layers
|
||||||
|
|
||||||
|
| Layer | Files | Testable? | Approach |
|
||||||
|
|-------|-------|-----------|----------|
|
||||||
|
| Data ops | column_, relation_, domain_, schema_, table_, database_dataops.go | Yes, pure | CRUD tests per file: create, duplicate, update/rename, delete, not-found, bounds, UpdateDate side effects |
|
||||||
|
| Pure helpers | `sortedKeys`, `schemaLocations`, `tableLocations`, `getColumnNames`, `parseSkipTablesUI`, help-text getters | Yes | Table-driven |
|
||||||
|
| Kind definitions | `indexKind/viewKind/sequenceKind/scriptKind` | Yes | Assert row builders and form-to-model mapping without rendering |
|
||||||
|
| Load/save | `loadDatabase`, `saveDatabase`, `createNewDatabase`, `importAndMergeDatabase`, `performMerge` | Partly | Temp files from `tests/assets`; verify format dispatch and error paths; avoid UI dialogs |
|
||||||
|
| Screens | *_screens.go, dialogs.go, main_menu.go | Yes, via simulation | tview app on tcell SimulationScreen; inject key events; assert navigation, form submit mutates model, cancel leaves it unchanged, delete confirm paths |
|
||||||
|
|
||||||
|
## Order
|
||||||
|
1. Data ops (largest gain, no tview)
|
||||||
|
2. Pure helpers and kinds
|
||||||
|
3. Load/save logic
|
||||||
|
4. Screen tests on simulation screen (menu, lists, forms, confirm dialogs, load/save)
|
||||||
|
|
||||||
|
## Decision
|
||||||
|
- Screen smoke tests via tview simulation screen are in scope (tcell `SimulationScreen`); drive keys/events, assert no panic and expected state.
|
||||||
|
|
||||||
|
## Live DB cases (dbtest)
|
||||||
|
- Load/save and import-merge from a live pgsql source: `dbtest.sh up postgres`; skip when `RELSPEC_TEST_PG_CONN` unset.
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
# dbtest: container test databases
|
||||||
|
|
||||||
|
Podman (preferred) or Docker. One tool for postgres, mssql, mysql.
|
||||||
|
|
||||||
|
## Commands
|
||||||
|
|
||||||
|
| Command | Action |
|
||||||
|
|---|---|
|
||||||
|
| `./tests/dbtest/dbtest.sh up <db\|all>` | Start, wait for ready, run init |
|
||||||
|
| `down <db\|all>` | Stop and remove |
|
||||||
|
| `restart <db>` | Fresh container |
|
||||||
|
| `status <db\|all>` | State and connection string |
|
||||||
|
| `env <db\|all>` | Print `export` line for the test env var |
|
||||||
|
| `test <db> [pkgs]` | up, `go test`, down |
|
||||||
|
| `shell <db>` | Interactive client |
|
||||||
|
| `exec <db> <file>` | Run SQL file |
|
||||||
|
| `logs <db>` | Container logs |
|
||||||
|
|
||||||
|
## Databases
|
||||||
|
|
||||||
|
| db | Port | Env var | Init | Default test pkgs |
|
||||||
|
|---|---|---|---|---|
|
||||||
|
| postgres | 5439 | `RELSPEC_TEST_PG_CONN` | `tests/postgres/init.sql` | readers/pgsql, writers/pgsql, writers/sqlexec |
|
||||||
|
| mssql | 1439 | `RELSPEC_TEST_MSSQL_CONN` | `test_data/mssql/test_schema.sql` (creates `RelSpecTest`) | readers/mssql, writers/mssql |
|
||||||
|
| mysql | 3309 | `RELSPEC_TEST_MYSQL_CONN` | `tests/dbtest/init/mysql.sql` | none (no Go driver/reader yet) |
|
||||||
|
|
||||||
|
## Env
|
||||||
|
|
||||||
|
| Var | Effect |
|
||||||
|
|---|---|
|
||||||
|
| `DBTEST_RUNTIME` | Force `podman` or `docker` |
|
||||||
|
| `DBTEST_TIMEOUT` | Ready wait, seconds (default 120) |
|
||||||
|
| `DBTEST_KEEP=1` | Keep container after `test` |
|
||||||
|
| `DBTEST_GOFLAGS` | Extra `go test` flags |
|
||||||
|
|
||||||
|
## Add a database
|
||||||
|
|
||||||
|
1. Add `dbs/<name>.sh` defining: `DB_NAME DB_IMAGE DB_CONTAINER DB_PORT DB_INTERNAL_PORT DB_ENV DB_INIT_MOUNT DB_CONN_VAR DB_CONN DB_DEFAULT_PKGS` and functions `db_ready db_post_init db_shell db_exec_file`.
|
||||||
|
2. Add the name to `ALL_DBS` in `dbtest.sh`.
|
||||||
|
|
||||||
|
## Notes
|
||||||
|
|
||||||
|
- Containers are named `relspec-test-<db>`; existing `tests/postgres/*.sh` and `make docker-*` use the same postgres name/port 5439 and remain independent.
|
||||||
|
- mssql needs ~2GB RAM and takes longer to become ready.
|
||||||
|
- Tests skip when the env var is unset.
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
# SQL Server container definition (sourced by dbtest.sh)
|
||||||
|
DB_NAME=mssql
|
||||||
|
DB_IMAGE=mcr.microsoft.com/mssql/server:2022-latest
|
||||||
|
DB_CONTAINER=relspec-test-mssql
|
||||||
|
DB_PORT=1439
|
||||||
|
DB_INTERNAL_PORT=1433
|
||||||
|
DB_ENV=(-e ACCEPT_EULA=Y -e "MSSQL_SA_PASSWORD=StrongPassword123!" -e MSSQL_PID=Express)
|
||||||
|
DB_INIT_MOUNT="$ROOT/test_data/mssql/test_schema.sql:/init/test_schema.sql"
|
||||||
|
DB_CONN_VAR=RELSPEC_TEST_MSSQL_CONN
|
||||||
|
DB_CONN="sqlserver://sa:StrongPassword123!@localhost:1439?database=RelSpecTest"
|
||||||
|
DB_DEFAULT_PKGS="./pkg/readers/mssql/ ./pkg/writers/mssql/"
|
||||||
|
|
||||||
|
_sqlcmd() { rt exec -i "$DB_CONTAINER" /opt/mssql-tools18/bin/sqlcmd -C -S localhost -U sa -P 'StrongPassword123!' "$@"; }
|
||||||
|
db_ready() { _sqlcmd -Q "SELECT 1" >/dev/null 2>&1; }
|
||||||
|
db_post_init() { _sqlcmd -b -i /init/test_schema.sql >/dev/null; }
|
||||||
|
db_shell() { rt exec -it "$DB_CONTAINER" /opt/mssql-tools18/bin/sqlcmd -C -S localhost -U sa -P 'StrongPassword123!' -d RelSpecTest "$@"; }
|
||||||
|
db_exec_file() { _sqlcmd -b -d RelSpecTest < "$1"; }
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
# MySQL container definition (sourced by dbtest.sh)
|
||||||
|
DB_NAME=mysql
|
||||||
|
DB_IMAGE=docker.io/library/mysql:8.4
|
||||||
|
DB_CONTAINER=relspec-test-mysql
|
||||||
|
DB_PORT=3309
|
||||||
|
DB_INTERNAL_PORT=3306
|
||||||
|
DB_ENV=(-e MYSQL_ROOT_PASSWORD=relspec_root_password -e MYSQL_DATABASE=relspec_test -e MYSQL_USER=relspec -e MYSQL_PASSWORD=relspec_test_password)
|
||||||
|
DB_INIT_MOUNT="$ROOT/tests/dbtest/init/mysql.sql:/docker-entrypoint-initdb.d/init.sql"
|
||||||
|
DB_CONN_VAR=RELSPEC_TEST_MYSQL_CONN
|
||||||
|
DB_CONN="relspec:relspec_test_password@tcp(localhost:3309)/relspec_test"
|
||||||
|
DB_DEFAULT_PKGS=""
|
||||||
|
|
||||||
|
db_ready() { rt exec "$DB_CONTAINER" mysqladmin ping -h 127.0.0.1 -urelspec -prelspec_test_password --silent >/dev/null 2>&1; }
|
||||||
|
db_post_init() { :; }
|
||||||
|
db_shell() { rt exec -it "$DB_CONTAINER" mysql -urelspec -prelspec_test_password relspec_test "$@"; }
|
||||||
|
db_exec_file() { rt exec -i "$DB_CONTAINER" mysql -urelspec -prelspec_test_password relspec_test < "$1"; }
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
# Postgres container definition (sourced by dbtest.sh)
|
||||||
|
DB_NAME=postgres
|
||||||
|
DB_IMAGE=docker.io/library/postgres:16-alpine
|
||||||
|
DB_CONTAINER=relspec-test-postgres
|
||||||
|
DB_PORT=5439
|
||||||
|
DB_INTERNAL_PORT=5432
|
||||||
|
DB_ENV=(-e POSTGRES_USER=relspec -e POSTGRES_PASSWORD=relspec_test_password -e POSTGRES_DB=relspec_test)
|
||||||
|
DB_INIT_MOUNT="$ROOT/tests/postgres/init.sql:/docker-entrypoint-initdb.d/init.sql"
|
||||||
|
DB_CONN_VAR=RELSPEC_TEST_PG_CONN
|
||||||
|
DB_CONN="postgres://relspec:relspec_test_password@localhost:5439/relspec_test"
|
||||||
|
DB_DEFAULT_PKGS="./pkg/readers/pgsql/ ./pkg/writers/pgsql/ ./pkg/writers/sqlexec/"
|
||||||
|
|
||||||
|
db_ready() { rt exec "$DB_CONTAINER" pg_isready -U relspec -d relspec_test >/dev/null 2>&1; }
|
||||||
|
db_post_init() { :; }
|
||||||
|
db_shell() { rt exec -it "$DB_CONTAINER" psql -U relspec -d relspec_test "$@"; }
|
||||||
|
db_exec_file() { rt exec -i "$DB_CONTAINER" psql -v ON_ERROR_STOP=1 -U relspec -d relspec_test < "$1"; }
|
||||||
Executable
+130
@@ -0,0 +1,130 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
# Reusable podman/docker test database tool for postgres, mssql and mysql.
|
||||||
|
# Usage: dbtest.sh <command> <db> [args]
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
HERE="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||||
|
ROOT="$(cd "$HERE/../.." && pwd)"
|
||||||
|
TIMEOUT="${DBTEST_TIMEOUT:-120}"
|
||||||
|
ALL_DBS=(postgres mssql mysql)
|
||||||
|
|
||||||
|
usage() {
|
||||||
|
cat <<USAGE
|
||||||
|
dbtest.sh <command> <db|all> [args]
|
||||||
|
|
||||||
|
commands:
|
||||||
|
up <db> start container, wait until ready, run init
|
||||||
|
down <db> stop and remove container
|
||||||
|
restart <db> down + up
|
||||||
|
status <db> running state and connection string
|
||||||
|
env <db> print 'export VAR=conn' (use: eval "\$(dbtest.sh env postgres)")
|
||||||
|
logs <db> container logs
|
||||||
|
shell <db> interactive client
|
||||||
|
exec <db> <file> run a SQL file against the database
|
||||||
|
test <db> [pkgs] up, run go tests with conn env set, down (keep with DBTEST_KEEP=1)
|
||||||
|
list supported databases
|
||||||
|
|
||||||
|
dbs: ${ALL_DBS[*]}
|
||||||
|
env: DBTEST_RUNTIME=podman|docker DBTEST_TIMEOUT=secs DBTEST_KEEP=1 DBTEST_GOFLAGS=...
|
||||||
|
USAGE
|
||||||
|
}
|
||||||
|
|
||||||
|
die() { echo "error: $*" >&2; exit 1; }
|
||||||
|
log() { echo "[dbtest] $*" >&2; }
|
||||||
|
|
||||||
|
detect_runtime() {
|
||||||
|
if [ -n "${DBTEST_RUNTIME:-}" ]; then echo "$DBTEST_RUNTIME"; return; fi
|
||||||
|
if command -v podman >/dev/null 2>&1; then echo podman
|
||||||
|
elif command -v docker >/dev/null 2>&1; then echo docker
|
||||||
|
else die "neither podman nor docker is installed"; fi
|
||||||
|
}
|
||||||
|
RUNTIME="$(detect_runtime)"
|
||||||
|
rt() { "$RUNTIME" "$@"; }
|
||||||
|
|
||||||
|
load_db() {
|
||||||
|
local f="$HERE/dbs/${1:-}.sh"
|
||||||
|
[ -f "$f" ] || die "unknown db '${1:-}' (supported: ${ALL_DBS[*]})"
|
||||||
|
# shellcheck disable=SC1090
|
||||||
|
source "$f"
|
||||||
|
}
|
||||||
|
|
||||||
|
is_running() { [ "$(rt inspect -f '{{.State.Running}}' "$DB_CONTAINER" 2>/dev/null || true)" = "true" ]; }
|
||||||
|
|
||||||
|
cmd_up() {
|
||||||
|
if is_running; then
|
||||||
|
log "$DB_NAME already running"
|
||||||
|
else
|
||||||
|
rt rm -f "$DB_CONTAINER" >/dev/null 2>&1 || true
|
||||||
|
log "starting $DB_NAME ($DB_IMAGE) on port $DB_PORT using $RUNTIME"
|
||||||
|
rt run -d --name "$DB_CONTAINER" "${DB_ENV[@]}" \
|
||||||
|
-p "$DB_PORT:$DB_INTERNAL_PORT" \
|
||||||
|
-v "$DB_INIT_MOUNT:ro,Z" "$DB_IMAGE" >/dev/null
|
||||||
|
fi
|
||||||
|
log "waiting for $DB_NAME (max ${TIMEOUT}s)"
|
||||||
|
local i=0
|
||||||
|
until db_ready; do
|
||||||
|
i=$((i + 1))
|
||||||
|
if [ "$i" -ge "$TIMEOUT" ]; then
|
||||||
|
rt logs --tail 50 "$DB_CONTAINER" >&2 || true
|
||||||
|
die "$DB_NAME did not become ready"
|
||||||
|
fi
|
||||||
|
sleep 1
|
||||||
|
done
|
||||||
|
db_post_init
|
||||||
|
log "$DB_NAME ready: $DB_CONN_VAR=$DB_CONN"
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd_down() {
|
||||||
|
rt rm -f "$DB_CONTAINER" >/dev/null 2>&1 || true
|
||||||
|
log "$DB_NAME removed"
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd_status() {
|
||||||
|
if is_running; then echo "$DB_NAME: running ($DB_CONTAINER, port $DB_PORT)"; else echo "$DB_NAME: stopped"; fi
|
||||||
|
echo "$DB_CONN_VAR=$DB_CONN"
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd_test() {
|
||||||
|
local pkgs="${*:-$DB_DEFAULT_PKGS}"
|
||||||
|
[ -n "$pkgs" ] || die "no test packages for $DB_NAME; pass packages as arguments"
|
||||||
|
cmd_up
|
||||||
|
[ -n "${DBTEST_KEEP:-}" ] || trap cmd_down EXIT
|
||||||
|
export "$DB_CONN_VAR=$DB_CONN"
|
||||||
|
cd "$ROOT"
|
||||||
|
# shellcheck disable=SC2086
|
||||||
|
go test -count=1 ${DBTEST_GOFLAGS:-} $pkgs
|
||||||
|
}
|
||||||
|
|
||||||
|
cmd="${1:-}"
|
||||||
|
case "$cmd" in
|
||||||
|
""|-h|--help|help) usage; exit 0 ;;
|
||||||
|
list) printf '%s\n' "${ALL_DBS[@]}"; exit 0 ;;
|
||||||
|
esac
|
||||||
|
shift
|
||||||
|
target="${1:-}"
|
||||||
|
[ -n "$target" ] || die "missing <db>"
|
||||||
|
shift || true
|
||||||
|
|
||||||
|
run_one() {
|
||||||
|
load_db "$1"
|
||||||
|
shift
|
||||||
|
case "$cmd" in
|
||||||
|
up) cmd_up ;;
|
||||||
|
down) cmd_down ;;
|
||||||
|
restart) cmd_down; cmd_up ;;
|
||||||
|
status) cmd_status ;;
|
||||||
|
env) echo "export $DB_CONN_VAR='$DB_CONN'" ;;
|
||||||
|
logs) rt logs "$DB_CONTAINER" ;;
|
||||||
|
shell) db_shell "$@" ;;
|
||||||
|
exec) [ -f "${1:-}" ] || die "usage: exec <db> <file>"; db_exec_file "$1" ;;
|
||||||
|
test) cmd_test "$@" ;;
|
||||||
|
*) usage; die "unknown command '$cmd'" ;;
|
||||||
|
esac
|
||||||
|
}
|
||||||
|
|
||||||
|
if [ "$target" = "all" ]; then
|
||||||
|
case "$cmd" in up|down|restart|status|env) ;; *) die "'all' is only valid for up/down/restart/status/env" ;; esac
|
||||||
|
for d in "${ALL_DBS[@]}"; do (run_one "$d"); done
|
||||||
|
else
|
||||||
|
run_one "$target" "$@"
|
||||||
|
fi
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
-- Minimal MySQL fixture for relspec container tests.
|
||||||
|
-- Extend when a MySQL reader/writer is added.
|
||||||
|
CREATE TABLE users (
|
||||||
|
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY,
|
||||||
|
email VARCHAR(255) NOT NULL,
|
||||||
|
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
UNIQUE KEY uq_users_email (email)
|
||||||
|
) ENGINE=InnoDB;
|
||||||
|
|
||||||
|
CREATE TABLE posts (
|
||||||
|
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY,
|
||||||
|
user_id BIGINT UNSIGNED NOT NULL,
|
||||||
|
title VARCHAR(255) NOT NULL,
|
||||||
|
body TEXT NULL,
|
||||||
|
KEY idx_posts_user (user_id),
|
||||||
|
CONSTRAINT fk_posts_user FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE
|
||||||
|
) ENGINE=InnoDB;
|
||||||
|
|
||||||
|
CREATE VIEW v_user_posts AS
|
||||||
|
SELECT u.id AS user_id, u.email, p.id AS post_id, p.title
|
||||||
|
FROM users u JOIN posts p ON p.user_id = u.id;
|
||||||
+27
@@ -0,0 +1,27 @@
|
|||||||
|
Copyright (c) 2009 The Go Authors. All rights reserved.
|
||||||
|
|
||||||
|
Redistribution and use in source and binary forms, with or without
|
||||||
|
modification, are permitted provided that the following conditions are
|
||||||
|
met:
|
||||||
|
|
||||||
|
* Redistributions of source code must retain the above copyright
|
||||||
|
notice, this list of conditions and the following disclaimer.
|
||||||
|
* Redistributions in binary form must reproduce the above
|
||||||
|
copyright notice, this list of conditions and the following disclaimer
|
||||||
|
in the documentation and/or other materials provided with the
|
||||||
|
distribution.
|
||||||
|
* Neither the name of Google Inc. nor the names of its
|
||||||
|
contributors may be used to endorse or promote products derived from
|
||||||
|
this software without specific prior written permission.
|
||||||
|
|
||||||
|
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
|
||||||
|
"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
|
||||||
|
LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
|
||||||
|
A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
|
||||||
|
OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
|
||||||
|
SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
|
||||||
|
LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
|
||||||
|
DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
|
||||||
|
THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
|
||||||
|
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||||
|
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||||
+14
@@ -0,0 +1,14 @@
|
|||||||
|
# filippo.io/edwards25519
|
||||||
|
|
||||||
|
```
|
||||||
|
import "filippo.io/edwards25519"
|
||||||
|
```
|
||||||
|
|
||||||
|
This library implements the edwards25519 elliptic curve, exposing the necessary APIs to build a wide array of higher-level primitives.
|
||||||
|
Read the docs at [pkg.go.dev/filippo.io/edwards25519](https://pkg.go.dev/filippo.io/edwards25519).
|
||||||
|
|
||||||
|
The code is originally derived from Adam Langley's internal implementation in the Go standard library, and includes George Tankersley's [performance improvements](https://golang.org/cl/71950). It was then further developed by Henry de Valence for use in ristretto255, and was finally [merged back into the Go standard library](https://golang.org/cl/276272) as of Go 1.17. It now tracks the upstream codebase and extends it with additional functionality.
|
||||||
|
|
||||||
|
Most users don't need this package, and should instead use `crypto/ed25519` for signatures, `golang.org/x/crypto/curve25519` for Diffie-Hellman, or `github.com/gtank/ristretto255` for prime order group logic. However, for anyone currently using a fork of `crypto/internal/edwards25519`/`crypto/ed25519/internal/edwards25519` or `github.com/agl/edwards25519`, this package should be a safer, faster, and more powerful alternative.
|
||||||
|
|
||||||
|
Since this package is meant to curb proliferation of edwards25519 implementations in the Go ecosystem, it welcomes requests for new APIs or reviewable performance improvements.
|
||||||
+20
@@ -0,0 +1,20 @@
|
|||||||
|
// Copyright (c) 2021 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
// Package edwards25519 implements group logic for the twisted Edwards curve
|
||||||
|
//
|
||||||
|
// -x^2 + y^2 = 1 + -(121665/121666)*x^2*y^2
|
||||||
|
//
|
||||||
|
// This is better known as the Edwards curve equivalent to Curve25519, and is
|
||||||
|
// the curve used by the Ed25519 signature scheme.
|
||||||
|
//
|
||||||
|
// Most users don't need this package, and should instead use crypto/ed25519 for
|
||||||
|
// signatures, golang.org/x/crypto/curve25519 for Diffie-Hellman, or
|
||||||
|
// github.com/gtank/ristretto255 for prime order group logic.
|
||||||
|
//
|
||||||
|
// However, developers who do need to interact with low-level edwards25519
|
||||||
|
// operations can use this package, which is an extended version of
|
||||||
|
// crypto/internal/edwards25519 from the standard library repackaged as
|
||||||
|
// an importable module.
|
||||||
|
package edwards25519
|
||||||
+427
@@ -0,0 +1,427 @@
|
|||||||
|
// Copyright (c) 2017 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package edwards25519
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"filippo.io/edwards25519/field"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Point types.
|
||||||
|
|
||||||
|
type projP1xP1 struct {
|
||||||
|
X, Y, Z, T field.Element
|
||||||
|
}
|
||||||
|
|
||||||
|
type projP2 struct {
|
||||||
|
X, Y, Z field.Element
|
||||||
|
}
|
||||||
|
|
||||||
|
// Point represents a point on the edwards25519 curve.
|
||||||
|
//
|
||||||
|
// This type works similarly to math/big.Int, and all arguments and receivers
|
||||||
|
// are allowed to alias.
|
||||||
|
//
|
||||||
|
// The zero value is NOT valid, and it may be used only as a receiver.
|
||||||
|
type Point struct {
|
||||||
|
// Make the type not comparable (i.e. used with == or as a map key), as
|
||||||
|
// equivalent points can be represented by different Go values.
|
||||||
|
_ incomparable
|
||||||
|
|
||||||
|
// The point is internally represented in extended coordinates (X, Y, Z, T)
|
||||||
|
// where x = X/Z, y = Y/Z, and xy = T/Z per https://eprint.iacr.org/2008/522.
|
||||||
|
x, y, z, t field.Element
|
||||||
|
}
|
||||||
|
|
||||||
|
type incomparable [0]func()
|
||||||
|
|
||||||
|
func checkInitialized(points ...*Point) {
|
||||||
|
for _, p := range points {
|
||||||
|
if p.x == (field.Element{}) && p.y == (field.Element{}) {
|
||||||
|
panic("edwards25519: use of uninitialized Point")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type projCached struct {
|
||||||
|
YplusX, YminusX, Z, T2d field.Element
|
||||||
|
}
|
||||||
|
|
||||||
|
type affineCached struct {
|
||||||
|
YplusX, YminusX, T2d field.Element
|
||||||
|
}
|
||||||
|
|
||||||
|
// Constructors.
|
||||||
|
|
||||||
|
func (v *projP2) Zero() *projP2 {
|
||||||
|
v.X.Zero()
|
||||||
|
v.Y.One()
|
||||||
|
v.Z.One()
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// identity is the point at infinity.
|
||||||
|
var identity, _ = new(Point).SetBytes([]byte{
|
||||||
|
1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
|
||||||
|
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0})
|
||||||
|
|
||||||
|
// NewIdentityPoint returns a new Point set to the identity.
|
||||||
|
func NewIdentityPoint() *Point {
|
||||||
|
return new(Point).Set(identity)
|
||||||
|
}
|
||||||
|
|
||||||
|
// generator is the canonical curve basepoint. See TestGenerator for the
|
||||||
|
// correspondence of this encoding with the values in RFC 8032.
|
||||||
|
var generator, _ = new(Point).SetBytes([]byte{
|
||||||
|
0x58, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66,
|
||||||
|
0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66,
|
||||||
|
0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66,
|
||||||
|
0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66})
|
||||||
|
|
||||||
|
// NewGeneratorPoint returns a new Point set to the canonical generator.
|
||||||
|
func NewGeneratorPoint() *Point {
|
||||||
|
return new(Point).Set(generator)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projCached) Zero() *projCached {
|
||||||
|
v.YplusX.One()
|
||||||
|
v.YminusX.One()
|
||||||
|
v.Z.One()
|
||||||
|
v.T2d.Zero()
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *affineCached) Zero() *affineCached {
|
||||||
|
v.YplusX.One()
|
||||||
|
v.YminusX.One()
|
||||||
|
v.T2d.Zero()
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Assignments.
|
||||||
|
|
||||||
|
// Set sets v = u, and returns v.
|
||||||
|
func (v *Point) Set(u *Point) *Point {
|
||||||
|
*v = *u
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Encoding.
|
||||||
|
|
||||||
|
// Bytes returns the canonical 32-byte encoding of v, according to RFC 8032,
|
||||||
|
// Section 5.1.2.
|
||||||
|
func (v *Point) Bytes() []byte {
|
||||||
|
// This function is outlined to make the allocations inline in the caller
|
||||||
|
// rather than happen on the heap.
|
||||||
|
var buf [32]byte
|
||||||
|
return v.bytes(&buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Point) bytes(buf *[32]byte) []byte {
|
||||||
|
checkInitialized(v)
|
||||||
|
|
||||||
|
var zInv, x, y field.Element
|
||||||
|
zInv.Invert(&v.z) // zInv = 1 / Z
|
||||||
|
x.Multiply(&v.x, &zInv) // x = X / Z
|
||||||
|
y.Multiply(&v.y, &zInv) // y = Y / Z
|
||||||
|
|
||||||
|
out := copyFieldElement(buf, &y)
|
||||||
|
out[31] |= byte(x.IsNegative() << 7)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
var feOne = new(field.Element).One()
|
||||||
|
|
||||||
|
// SetBytes sets v = x, where x is a 32-byte encoding of v. If x does not
|
||||||
|
// represent a valid point on the curve, SetBytes returns nil and an error and
|
||||||
|
// the receiver is unchanged. Otherwise, SetBytes returns v.
|
||||||
|
//
|
||||||
|
// Note that SetBytes accepts all non-canonical encodings of valid points.
|
||||||
|
// That is, it follows decoding rules that match most implementations in
|
||||||
|
// the ecosystem rather than RFC 8032.
|
||||||
|
func (v *Point) SetBytes(x []byte) (*Point, error) {
|
||||||
|
// Specifically, the non-canonical encodings that are accepted are
|
||||||
|
// 1) the ones where the field element is not reduced (see the
|
||||||
|
// (*field.Element).SetBytes docs) and
|
||||||
|
// 2) the ones where the x-coordinate is zero and the sign bit is set.
|
||||||
|
//
|
||||||
|
// Read more at https://hdevalence.ca/blog/2020-10-04-its-25519am,
|
||||||
|
// specifically the "Canonical A, R" section.
|
||||||
|
|
||||||
|
y, err := new(field.Element).SetBytes(x)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("edwards25519: invalid point encoding length")
|
||||||
|
}
|
||||||
|
|
||||||
|
// -x² + y² = 1 + dx²y²
|
||||||
|
// x² + dx²y² = x²(dy² + 1) = y² - 1
|
||||||
|
// x² = (y² - 1) / (dy² + 1)
|
||||||
|
|
||||||
|
// u = y² - 1
|
||||||
|
y2 := new(field.Element).Square(y)
|
||||||
|
u := new(field.Element).Subtract(y2, feOne)
|
||||||
|
|
||||||
|
// v = dy² + 1
|
||||||
|
vv := new(field.Element).Multiply(y2, d)
|
||||||
|
vv = vv.Add(vv, feOne)
|
||||||
|
|
||||||
|
// x = +√(u/v)
|
||||||
|
xx, wasSquare := new(field.Element).SqrtRatio(u, vv)
|
||||||
|
if wasSquare == 0 {
|
||||||
|
return nil, errors.New("edwards25519: invalid point encoding")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Select the negative square root if the sign bit is set.
|
||||||
|
xxNeg := new(field.Element).Negate(xx)
|
||||||
|
xx = xx.Select(xxNeg, xx, int(x[31]>>7))
|
||||||
|
|
||||||
|
v.x.Set(xx)
|
||||||
|
v.y.Set(y)
|
||||||
|
v.z.One()
|
||||||
|
v.t.Multiply(xx, y) // xy = T / Z
|
||||||
|
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func copyFieldElement(buf *[32]byte, v *field.Element) []byte {
|
||||||
|
copy(buf[:], v.Bytes())
|
||||||
|
return buf[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Conversions.
|
||||||
|
|
||||||
|
func (v *projP2) FromP1xP1(p *projP1xP1) *projP2 {
|
||||||
|
v.X.Multiply(&p.X, &p.T)
|
||||||
|
v.Y.Multiply(&p.Y, &p.Z)
|
||||||
|
v.Z.Multiply(&p.Z, &p.T)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projP2) FromP3(p *Point) *projP2 {
|
||||||
|
v.X.Set(&p.x)
|
||||||
|
v.Y.Set(&p.y)
|
||||||
|
v.Z.Set(&p.z)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Point) fromP1xP1(p *projP1xP1) *Point {
|
||||||
|
v.x.Multiply(&p.X, &p.T)
|
||||||
|
v.y.Multiply(&p.Y, &p.Z)
|
||||||
|
v.z.Multiply(&p.Z, &p.T)
|
||||||
|
v.t.Multiply(&p.X, &p.Y)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Point) fromP2(p *projP2) *Point {
|
||||||
|
v.x.Multiply(&p.X, &p.Z)
|
||||||
|
v.y.Multiply(&p.Y, &p.Z)
|
||||||
|
v.z.Square(&p.Z)
|
||||||
|
v.t.Multiply(&p.X, &p.Y)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// d is a constant in the curve equation.
|
||||||
|
var d, _ = new(field.Element).SetBytes([]byte{
|
||||||
|
0xa3, 0x78, 0x59, 0x13, 0xca, 0x4d, 0xeb, 0x75,
|
||||||
|
0xab, 0xd8, 0x41, 0x41, 0x4d, 0x0a, 0x70, 0x00,
|
||||||
|
0x98, 0xe8, 0x79, 0x77, 0x79, 0x40, 0xc7, 0x8c,
|
||||||
|
0x73, 0xfe, 0x6f, 0x2b, 0xee, 0x6c, 0x03, 0x52})
|
||||||
|
var d2 = new(field.Element).Add(d, d)
|
||||||
|
|
||||||
|
func (v *projCached) FromP3(p *Point) *projCached {
|
||||||
|
v.YplusX.Add(&p.y, &p.x)
|
||||||
|
v.YminusX.Subtract(&p.y, &p.x)
|
||||||
|
v.Z.Set(&p.z)
|
||||||
|
v.T2d.Multiply(&p.t, d2)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *affineCached) FromP3(p *Point) *affineCached {
|
||||||
|
v.YplusX.Add(&p.y, &p.x)
|
||||||
|
v.YminusX.Subtract(&p.y, &p.x)
|
||||||
|
v.T2d.Multiply(&p.t, d2)
|
||||||
|
|
||||||
|
var invZ field.Element
|
||||||
|
invZ.Invert(&p.z)
|
||||||
|
v.YplusX.Multiply(&v.YplusX, &invZ)
|
||||||
|
v.YminusX.Multiply(&v.YminusX, &invZ)
|
||||||
|
v.T2d.Multiply(&v.T2d, &invZ)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// (Re)addition and subtraction.
|
||||||
|
|
||||||
|
// Add sets v = p + q, and returns v.
|
||||||
|
func (v *Point) Add(p, q *Point) *Point {
|
||||||
|
checkInitialized(p, q)
|
||||||
|
qCached := new(projCached).FromP3(q)
|
||||||
|
result := new(projP1xP1).Add(p, qCached)
|
||||||
|
return v.fromP1xP1(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Subtract sets v = p - q, and returns v.
|
||||||
|
func (v *Point) Subtract(p, q *Point) *Point {
|
||||||
|
checkInitialized(p, q)
|
||||||
|
qCached := new(projCached).FromP3(q)
|
||||||
|
result := new(projP1xP1).Sub(p, qCached)
|
||||||
|
return v.fromP1xP1(result)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projP1xP1) Add(p *Point, q *projCached) *projP1xP1 {
|
||||||
|
var YplusX, YminusX, PP, MM, TT2d, ZZ2 field.Element
|
||||||
|
|
||||||
|
YplusX.Add(&p.y, &p.x)
|
||||||
|
YminusX.Subtract(&p.y, &p.x)
|
||||||
|
|
||||||
|
PP.Multiply(&YplusX, &q.YplusX)
|
||||||
|
MM.Multiply(&YminusX, &q.YminusX)
|
||||||
|
TT2d.Multiply(&p.t, &q.T2d)
|
||||||
|
ZZ2.Multiply(&p.z, &q.Z)
|
||||||
|
|
||||||
|
ZZ2.Add(&ZZ2, &ZZ2)
|
||||||
|
|
||||||
|
v.X.Subtract(&PP, &MM)
|
||||||
|
v.Y.Add(&PP, &MM)
|
||||||
|
v.Z.Add(&ZZ2, &TT2d)
|
||||||
|
v.T.Subtract(&ZZ2, &TT2d)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projP1xP1) Sub(p *Point, q *projCached) *projP1xP1 {
|
||||||
|
var YplusX, YminusX, PP, MM, TT2d, ZZ2 field.Element
|
||||||
|
|
||||||
|
YplusX.Add(&p.y, &p.x)
|
||||||
|
YminusX.Subtract(&p.y, &p.x)
|
||||||
|
|
||||||
|
PP.Multiply(&YplusX, &q.YminusX) // flipped sign
|
||||||
|
MM.Multiply(&YminusX, &q.YplusX) // flipped sign
|
||||||
|
TT2d.Multiply(&p.t, &q.T2d)
|
||||||
|
ZZ2.Multiply(&p.z, &q.Z)
|
||||||
|
|
||||||
|
ZZ2.Add(&ZZ2, &ZZ2)
|
||||||
|
|
||||||
|
v.X.Subtract(&PP, &MM)
|
||||||
|
v.Y.Add(&PP, &MM)
|
||||||
|
v.Z.Subtract(&ZZ2, &TT2d) // flipped sign
|
||||||
|
v.T.Add(&ZZ2, &TT2d) // flipped sign
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projP1xP1) AddAffine(p *Point, q *affineCached) *projP1xP1 {
|
||||||
|
var YplusX, YminusX, PP, MM, TT2d, Z2 field.Element
|
||||||
|
|
||||||
|
YplusX.Add(&p.y, &p.x)
|
||||||
|
YminusX.Subtract(&p.y, &p.x)
|
||||||
|
|
||||||
|
PP.Multiply(&YplusX, &q.YplusX)
|
||||||
|
MM.Multiply(&YminusX, &q.YminusX)
|
||||||
|
TT2d.Multiply(&p.t, &q.T2d)
|
||||||
|
|
||||||
|
Z2.Add(&p.z, &p.z)
|
||||||
|
|
||||||
|
v.X.Subtract(&PP, &MM)
|
||||||
|
v.Y.Add(&PP, &MM)
|
||||||
|
v.Z.Add(&Z2, &TT2d)
|
||||||
|
v.T.Subtract(&Z2, &TT2d)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *projP1xP1) SubAffine(p *Point, q *affineCached) *projP1xP1 {
|
||||||
|
var YplusX, YminusX, PP, MM, TT2d, Z2 field.Element
|
||||||
|
|
||||||
|
YplusX.Add(&p.y, &p.x)
|
||||||
|
YminusX.Subtract(&p.y, &p.x)
|
||||||
|
|
||||||
|
PP.Multiply(&YplusX, &q.YminusX) // flipped sign
|
||||||
|
MM.Multiply(&YminusX, &q.YplusX) // flipped sign
|
||||||
|
TT2d.Multiply(&p.t, &q.T2d)
|
||||||
|
|
||||||
|
Z2.Add(&p.z, &p.z)
|
||||||
|
|
||||||
|
v.X.Subtract(&PP, &MM)
|
||||||
|
v.Y.Add(&PP, &MM)
|
||||||
|
v.Z.Subtract(&Z2, &TT2d) // flipped sign
|
||||||
|
v.T.Add(&Z2, &TT2d) // flipped sign
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Doubling.
|
||||||
|
|
||||||
|
func (v *projP1xP1) Double(p *projP2) *projP1xP1 {
|
||||||
|
var XX, YY, ZZ2, XplusYsq field.Element
|
||||||
|
|
||||||
|
XX.Square(&p.X)
|
||||||
|
YY.Square(&p.Y)
|
||||||
|
ZZ2.Square(&p.Z)
|
||||||
|
ZZ2.Add(&ZZ2, &ZZ2)
|
||||||
|
XplusYsq.Add(&p.X, &p.Y)
|
||||||
|
XplusYsq.Square(&XplusYsq)
|
||||||
|
|
||||||
|
v.Y.Add(&YY, &XX)
|
||||||
|
v.Z.Subtract(&YY, &XX)
|
||||||
|
|
||||||
|
v.X.Subtract(&XplusYsq, &v.Y)
|
||||||
|
v.T.Subtract(&ZZ2, &v.Z)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Negation.
|
||||||
|
|
||||||
|
// Negate sets v = -p, and returns v.
|
||||||
|
func (v *Point) Negate(p *Point) *Point {
|
||||||
|
checkInitialized(p)
|
||||||
|
v.x.Negate(&p.x)
|
||||||
|
v.y.Set(&p.y)
|
||||||
|
v.z.Set(&p.z)
|
||||||
|
v.t.Negate(&p.t)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Equal returns 1 if v is equivalent to u, and 0 otherwise.
|
||||||
|
func (v *Point) Equal(u *Point) int {
|
||||||
|
checkInitialized(v, u)
|
||||||
|
|
||||||
|
var t1, t2, t3, t4 field.Element
|
||||||
|
t1.Multiply(&v.x, &u.z)
|
||||||
|
t2.Multiply(&u.x, &v.z)
|
||||||
|
t3.Multiply(&v.y, &u.z)
|
||||||
|
t4.Multiply(&u.y, &v.z)
|
||||||
|
|
||||||
|
return t1.Equal(&t2) & t3.Equal(&t4)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Constant-time operations
|
||||||
|
|
||||||
|
// Select sets v to a if cond == 1 and to b if cond == 0.
|
||||||
|
func (v *projCached) Select(a, b *projCached, cond int) *projCached {
|
||||||
|
v.YplusX.Select(&a.YplusX, &b.YplusX, cond)
|
||||||
|
v.YminusX.Select(&a.YminusX, &b.YminusX, cond)
|
||||||
|
v.Z.Select(&a.Z, &b.Z, cond)
|
||||||
|
v.T2d.Select(&a.T2d, &b.T2d, cond)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Select sets v to a if cond == 1 and to b if cond == 0.
|
||||||
|
func (v *affineCached) Select(a, b *affineCached, cond int) *affineCached {
|
||||||
|
v.YplusX.Select(&a.YplusX, &b.YplusX, cond)
|
||||||
|
v.YminusX.Select(&a.YminusX, &b.YminusX, cond)
|
||||||
|
v.T2d.Select(&a.T2d, &b.T2d, cond)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// CondNeg negates v if cond == 1 and leaves it unchanged if cond == 0.
|
||||||
|
func (v *projCached) CondNeg(cond int) *projCached {
|
||||||
|
v.YplusX.Swap(&v.YminusX, cond)
|
||||||
|
v.T2d.Select(new(field.Element).Negate(&v.T2d), &v.T2d, cond)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// CondNeg negates v if cond == 1 and leaves it unchanged if cond == 0.
|
||||||
|
func (v *affineCached) CondNeg(cond int) *affineCached {
|
||||||
|
v.YplusX.Swap(&v.YminusX, cond)
|
||||||
|
v.T2d.Select(new(field.Element).Negate(&v.T2d), &v.T2d, cond)
|
||||||
|
return v
|
||||||
|
}
|
||||||
+349
@@ -0,0 +1,349 @@
|
|||||||
|
// Copyright (c) 2021 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package edwards25519
|
||||||
|
|
||||||
|
// This file contains additional functionality that is not included in the
|
||||||
|
// upstream crypto/internal/edwards25519 package.
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"filippo.io/edwards25519/field"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ExtendedCoordinates returns v in extended coordinates (X:Y:Z:T) where
|
||||||
|
// x = X/Z, y = Y/Z, and xy = T/Z as in https://eprint.iacr.org/2008/522.
|
||||||
|
func (v *Point) ExtendedCoordinates() (X, Y, Z, T *field.Element) {
|
||||||
|
// This function is outlined to make the allocations inline in the caller
|
||||||
|
// rather than happen on the heap. Don't change the style without making
|
||||||
|
// sure it doesn't increase the inliner cost.
|
||||||
|
var e [4]field.Element
|
||||||
|
X, Y, Z, T = v.extendedCoordinates(&e)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Point) extendedCoordinates(e *[4]field.Element) (X, Y, Z, T *field.Element) {
|
||||||
|
checkInitialized(v)
|
||||||
|
X = e[0].Set(&v.x)
|
||||||
|
Y = e[1].Set(&v.y)
|
||||||
|
Z = e[2].Set(&v.z)
|
||||||
|
T = e[3].Set(&v.t)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetExtendedCoordinates sets v = (X:Y:Z:T) in extended coordinates where
|
||||||
|
// x = X/Z, y = Y/Z, and xy = T/Z as in https://eprint.iacr.org/2008/522.
|
||||||
|
//
|
||||||
|
// If the coordinates are invalid or don't represent a valid point on the curve,
|
||||||
|
// SetExtendedCoordinates returns nil and an error and the receiver is
|
||||||
|
// unchanged. Otherwise, SetExtendedCoordinates returns v.
|
||||||
|
func (v *Point) SetExtendedCoordinates(X, Y, Z, T *field.Element) (*Point, error) {
|
||||||
|
if !isOnCurve(X, Y, Z, T) {
|
||||||
|
return nil, errors.New("edwards25519: invalid point coordinates")
|
||||||
|
}
|
||||||
|
v.x.Set(X)
|
||||||
|
v.y.Set(Y)
|
||||||
|
v.z.Set(Z)
|
||||||
|
v.t.Set(T)
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func isOnCurve(X, Y, Z, T *field.Element) bool {
|
||||||
|
var lhs, rhs field.Element
|
||||||
|
XX := new(field.Element).Square(X)
|
||||||
|
YY := new(field.Element).Square(Y)
|
||||||
|
ZZ := new(field.Element).Square(Z)
|
||||||
|
TT := new(field.Element).Square(T)
|
||||||
|
// -x² + y² = 1 + dx²y²
|
||||||
|
// -(X/Z)² + (Y/Z)² = 1 + d(T/Z)²
|
||||||
|
// -X² + Y² = Z² + dT²
|
||||||
|
lhs.Subtract(YY, XX)
|
||||||
|
rhs.Multiply(d, TT).Add(&rhs, ZZ)
|
||||||
|
if lhs.Equal(&rhs) != 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// xy = T/Z
|
||||||
|
// XY/Z² = T/Z
|
||||||
|
// XY = TZ
|
||||||
|
lhs.Multiply(X, Y)
|
||||||
|
rhs.Multiply(T, Z)
|
||||||
|
return lhs.Equal(&rhs) == 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// BytesMontgomery converts v to a point on the birationally-equivalent
|
||||||
|
// Curve25519 Montgomery curve, and returns its canonical 32 bytes encoding
|
||||||
|
// according to RFC 7748.
|
||||||
|
//
|
||||||
|
// Note that BytesMontgomery only encodes the u-coordinate, so v and -v encode
|
||||||
|
// to the same value. If v is the identity point, BytesMontgomery returns 32
|
||||||
|
// zero bytes, analogously to the X25519 function.
|
||||||
|
//
|
||||||
|
// The lack of an inverse operation (such as SetMontgomeryBytes) is deliberate:
|
||||||
|
// while every valid edwards25519 point has a unique u-coordinate Montgomery
|
||||||
|
// encoding, X25519 accepts inputs on the quadratic twist, which don't correspond
|
||||||
|
// to any edwards25519 point, and every other X25519 input corresponds to two
|
||||||
|
// edwards25519 points.
|
||||||
|
func (v *Point) BytesMontgomery() []byte {
|
||||||
|
// This function is outlined to make the allocations inline in the caller
|
||||||
|
// rather than happen on the heap.
|
||||||
|
var buf [32]byte
|
||||||
|
return v.bytesMontgomery(&buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Point) bytesMontgomery(buf *[32]byte) []byte {
|
||||||
|
checkInitialized(v)
|
||||||
|
|
||||||
|
// RFC 7748, Section 4.1 provides the bilinear map to calculate the
|
||||||
|
// Montgomery u-coordinate
|
||||||
|
//
|
||||||
|
// u = (1 + y) / (1 - y)
|
||||||
|
//
|
||||||
|
// where y = Y / Z.
|
||||||
|
|
||||||
|
var y, recip, u field.Element
|
||||||
|
|
||||||
|
y.Multiply(&v.y, y.Invert(&v.z)) // y = Y / Z
|
||||||
|
recip.Invert(recip.Subtract(feOne, &y)) // r = 1/(1 - y)
|
||||||
|
u.Multiply(u.Add(feOne, &y), &recip) // u = (1 + y)*r
|
||||||
|
|
||||||
|
return copyFieldElement(buf, &u)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MultByCofactor sets v = 8 * p, and returns v.
|
||||||
|
func (v *Point) MultByCofactor(p *Point) *Point {
|
||||||
|
checkInitialized(p)
|
||||||
|
result := projP1xP1{}
|
||||||
|
pp := (&projP2{}).FromP3(p)
|
||||||
|
result.Double(pp)
|
||||||
|
pp.FromP1xP1(&result)
|
||||||
|
result.Double(pp)
|
||||||
|
pp.FromP1xP1(&result)
|
||||||
|
result.Double(pp)
|
||||||
|
return v.fromP1xP1(&result)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Given k > 0, set s = s**(2*i).
|
||||||
|
func (s *Scalar) pow2k(k int) {
|
||||||
|
for i := 0; i < k; i++ {
|
||||||
|
s.Multiply(s, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Invert sets s to the inverse of a nonzero scalar v, and returns s.
|
||||||
|
//
|
||||||
|
// If t is zero, Invert returns zero.
|
||||||
|
func (s *Scalar) Invert(t *Scalar) *Scalar {
|
||||||
|
// Uses a hardcoded sliding window of width 4.
|
||||||
|
var table [8]Scalar
|
||||||
|
var tt Scalar
|
||||||
|
tt.Multiply(t, t)
|
||||||
|
table[0] = *t
|
||||||
|
for i := 0; i < 7; i++ {
|
||||||
|
table[i+1].Multiply(&table[i], &tt)
|
||||||
|
}
|
||||||
|
// Now table = [t**1, t**3, t**5, t**7, t**9, t**11, t**13, t**15]
|
||||||
|
// so t**k = t[k/2] for odd k
|
||||||
|
|
||||||
|
// To compute the sliding window digits, use the following Sage script:
|
||||||
|
|
||||||
|
// sage: import itertools
|
||||||
|
// sage: def sliding_window(w,k):
|
||||||
|
// ....: digits = []
|
||||||
|
// ....: while k > 0:
|
||||||
|
// ....: if k % 2 == 1:
|
||||||
|
// ....: kmod = k % (2**w)
|
||||||
|
// ....: digits.append(kmod)
|
||||||
|
// ....: k = k - kmod
|
||||||
|
// ....: else:
|
||||||
|
// ....: digits.append(0)
|
||||||
|
// ....: k = k // 2
|
||||||
|
// ....: return digits
|
||||||
|
|
||||||
|
// Now we can compute s roughly as follows:
|
||||||
|
|
||||||
|
// sage: s = 1
|
||||||
|
// sage: for coeff in reversed(sliding_window(4,l-2)):
|
||||||
|
// ....: s = s*s
|
||||||
|
// ....: if coeff > 0 :
|
||||||
|
// ....: s = s*t**coeff
|
||||||
|
|
||||||
|
// This works on one bit at a time, with many runs of zeros.
|
||||||
|
// The digits can be collapsed into [(count, coeff)] as follows:
|
||||||
|
|
||||||
|
// sage: [(len(list(group)),d) for d,group in itertools.groupby(sliding_window(4,l-2))]
|
||||||
|
|
||||||
|
// Entries of the form (k, 0) turn into pow2k(k)
|
||||||
|
// Entries of the form (1, coeff) turn into a squaring and then a table lookup.
|
||||||
|
// We can fold the squaring into the previous pow2k(k) as pow2k(k+1).
|
||||||
|
|
||||||
|
*s = table[1/2]
|
||||||
|
s.pow2k(127 + 1)
|
||||||
|
s.Multiply(s, &table[1/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[9/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[11/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[13/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[15/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[7/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[15/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[5/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[1/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[15/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[15/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[7/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[3/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[11/2])
|
||||||
|
s.pow2k(5 + 1)
|
||||||
|
s.Multiply(s, &table[11/2])
|
||||||
|
s.pow2k(9 + 1)
|
||||||
|
s.Multiply(s, &table[9/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[3/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[3/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[3/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[9/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[7/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[3/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[13/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[7/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[9/2])
|
||||||
|
s.pow2k(3 + 1)
|
||||||
|
s.Multiply(s, &table[15/2])
|
||||||
|
s.pow2k(4 + 1)
|
||||||
|
s.Multiply(s, &table[11/2])
|
||||||
|
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// MultiScalarMult sets v = sum(scalars[i] * points[i]), and returns v.
|
||||||
|
//
|
||||||
|
// Execution time depends only on the lengths of the two slices, which must match.
|
||||||
|
func (v *Point) MultiScalarMult(scalars []*Scalar, points []*Point) *Point {
|
||||||
|
if len(scalars) != len(points) {
|
||||||
|
panic("edwards25519: called MultiScalarMult with different size inputs")
|
||||||
|
}
|
||||||
|
checkInitialized(points...)
|
||||||
|
|
||||||
|
// Proceed as in the single-base case, but share doublings
|
||||||
|
// between each point in the multiscalar equation.
|
||||||
|
|
||||||
|
// Build lookup tables for each point
|
||||||
|
tables := make([]projLookupTable, len(points))
|
||||||
|
for i := range tables {
|
||||||
|
tables[i].FromP3(points[i])
|
||||||
|
}
|
||||||
|
// Compute signed radix-16 digits for each scalar
|
||||||
|
digits := make([][64]int8, len(scalars))
|
||||||
|
for i := range digits {
|
||||||
|
digits[i] = scalars[i].signedRadix16()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unwrap first loop iteration to save computing 16*identity
|
||||||
|
multiple := &projCached{}
|
||||||
|
tmp1 := &projP1xP1{}
|
||||||
|
tmp2 := &projP2{}
|
||||||
|
// Lookup-and-add the appropriate multiple of each input point
|
||||||
|
for j := range tables {
|
||||||
|
tables[j].SelectInto(multiple, digits[j][63])
|
||||||
|
tmp1.Add(v, multiple) // tmp1 = v + x_(j,63)*Q in P1xP1 coords
|
||||||
|
v.fromP1xP1(tmp1) // update v
|
||||||
|
}
|
||||||
|
tmp2.FromP3(v) // set up tmp2 = v in P2 coords for next iteration
|
||||||
|
for i := 62; i >= 0; i-- {
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 2*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 2*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 4*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 4*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 8*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 8*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 16*(prev) in P1xP1 coords
|
||||||
|
v.fromP1xP1(tmp1) // v = 16*(prev) in P3 coords
|
||||||
|
// Lookup-and-add the appropriate multiple of each input point
|
||||||
|
for j := range tables {
|
||||||
|
tables[j].SelectInto(multiple, digits[j][i])
|
||||||
|
tmp1.Add(v, multiple) // tmp1 = v + x_(j,i)*Q in P1xP1 coords
|
||||||
|
v.fromP1xP1(tmp1) // update v
|
||||||
|
}
|
||||||
|
tmp2.FromP3(v) // set up tmp2 = v in P2 coords for next iteration
|
||||||
|
}
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// VarTimeMultiScalarMult sets v = sum(scalars[i] * points[i]), and returns v.
|
||||||
|
//
|
||||||
|
// Execution time depends on the inputs.
|
||||||
|
func (v *Point) VarTimeMultiScalarMult(scalars []*Scalar, points []*Point) *Point {
|
||||||
|
if len(scalars) != len(points) {
|
||||||
|
panic("edwards25519: called VarTimeMultiScalarMult with different size inputs")
|
||||||
|
}
|
||||||
|
checkInitialized(points...)
|
||||||
|
|
||||||
|
// Generalize double-base NAF computation to arbitrary sizes.
|
||||||
|
// Here all the points are dynamic, so we only use the smaller
|
||||||
|
// tables.
|
||||||
|
|
||||||
|
// Build lookup tables for each point
|
||||||
|
tables := make([]nafLookupTable5, len(points))
|
||||||
|
for i := range tables {
|
||||||
|
tables[i].FromP3(points[i])
|
||||||
|
}
|
||||||
|
// Compute a NAF for each scalar
|
||||||
|
nafs := make([][256]int8, len(scalars))
|
||||||
|
for i := range nafs {
|
||||||
|
nafs[i] = scalars[i].nonAdjacentForm(5)
|
||||||
|
}
|
||||||
|
|
||||||
|
multiple := &projCached{}
|
||||||
|
tmp1 := &projP1xP1{}
|
||||||
|
tmp2 := &projP2{}
|
||||||
|
tmp2.Zero()
|
||||||
|
|
||||||
|
// Move from high to low bits, doubling the accumulator
|
||||||
|
// at each iteration and checking whether there is a nonzero
|
||||||
|
// coefficient to look up a multiple of.
|
||||||
|
//
|
||||||
|
// Skip trying to find the first nonzero coefficent, because
|
||||||
|
// searching might be more work than a few extra doublings.
|
||||||
|
for i := 255; i >= 0; i-- {
|
||||||
|
tmp1.Double(tmp2)
|
||||||
|
|
||||||
|
for j := range nafs {
|
||||||
|
if nafs[j][i] > 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
tables[j].SelectInto(multiple, nafs[j][i])
|
||||||
|
tmp1.Add(v, multiple)
|
||||||
|
} else if nafs[j][i] < 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
tables[j].SelectInto(multiple, -nafs[j][i])
|
||||||
|
tmp1.Sub(v, multiple)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
tmp2.FromP1xP1(tmp1)
|
||||||
|
}
|
||||||
|
|
||||||
|
v.fromP2(tmp2)
|
||||||
|
return v
|
||||||
|
}
|
||||||
+420
@@ -0,0 +1,420 @@
|
|||||||
|
// Copyright (c) 2017 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
// Package field implements fast arithmetic modulo 2^255-19.
|
||||||
|
package field
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"math/bits"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Element represents an element of the field GF(2^255-19). Note that this
|
||||||
|
// is not a cryptographically secure group, and should only be used to interact
|
||||||
|
// with edwards25519.Point coordinates.
|
||||||
|
//
|
||||||
|
// This type works similarly to math/big.Int, and all arguments and receivers
|
||||||
|
// are allowed to alias.
|
||||||
|
//
|
||||||
|
// The zero value is a valid zero element.
|
||||||
|
type Element struct {
|
||||||
|
// An element t represents the integer
|
||||||
|
// t.l0 + t.l1*2^51 + t.l2*2^102 + t.l3*2^153 + t.l4*2^204
|
||||||
|
//
|
||||||
|
// Between operations, all limbs are expected to be lower than 2^52.
|
||||||
|
l0 uint64
|
||||||
|
l1 uint64
|
||||||
|
l2 uint64
|
||||||
|
l3 uint64
|
||||||
|
l4 uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
const maskLow51Bits uint64 = (1 << 51) - 1
|
||||||
|
|
||||||
|
var feZero = &Element{0, 0, 0, 0, 0}
|
||||||
|
|
||||||
|
// Zero sets v = 0, and returns v.
|
||||||
|
func (v *Element) Zero() *Element {
|
||||||
|
*v = *feZero
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
var feOne = &Element{1, 0, 0, 0, 0}
|
||||||
|
|
||||||
|
// One sets v = 1, and returns v.
|
||||||
|
func (v *Element) One() *Element {
|
||||||
|
*v = *feOne
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// reduce reduces v modulo 2^255 - 19 and returns it.
|
||||||
|
func (v *Element) reduce() *Element {
|
||||||
|
v.carryPropagate()
|
||||||
|
|
||||||
|
// After the light reduction we now have a field element representation
|
||||||
|
// v < 2^255 + 2^13 * 19, but need v < 2^255 - 19.
|
||||||
|
|
||||||
|
// If v >= 2^255 - 19, then v + 19 >= 2^255, which would overflow 2^255 - 1,
|
||||||
|
// generating a carry. That is, c will be 0 if v < 2^255 - 19, and 1 otherwise.
|
||||||
|
c := (v.l0 + 19) >> 51
|
||||||
|
c = (v.l1 + c) >> 51
|
||||||
|
c = (v.l2 + c) >> 51
|
||||||
|
c = (v.l3 + c) >> 51
|
||||||
|
c = (v.l4 + c) >> 51
|
||||||
|
|
||||||
|
// If v < 2^255 - 19 and c = 0, this will be a no-op. Otherwise, it's
|
||||||
|
// effectively applying the reduction identity to the carry.
|
||||||
|
v.l0 += 19 * c
|
||||||
|
|
||||||
|
v.l1 += v.l0 >> 51
|
||||||
|
v.l0 = v.l0 & maskLow51Bits
|
||||||
|
v.l2 += v.l1 >> 51
|
||||||
|
v.l1 = v.l1 & maskLow51Bits
|
||||||
|
v.l3 += v.l2 >> 51
|
||||||
|
v.l2 = v.l2 & maskLow51Bits
|
||||||
|
v.l4 += v.l3 >> 51
|
||||||
|
v.l3 = v.l3 & maskLow51Bits
|
||||||
|
// no additional carry
|
||||||
|
v.l4 = v.l4 & maskLow51Bits
|
||||||
|
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add sets v = a + b, and returns v.
|
||||||
|
func (v *Element) Add(a, b *Element) *Element {
|
||||||
|
v.l0 = a.l0 + b.l0
|
||||||
|
v.l1 = a.l1 + b.l1
|
||||||
|
v.l2 = a.l2 + b.l2
|
||||||
|
v.l3 = a.l3 + b.l3
|
||||||
|
v.l4 = a.l4 + b.l4
|
||||||
|
// Using the generic implementation here is actually faster than the
|
||||||
|
// assembly. Probably because the body of this function is so simple that
|
||||||
|
// the compiler can figure out better optimizations by inlining the carry
|
||||||
|
// propagation.
|
||||||
|
return v.carryPropagateGeneric()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Subtract sets v = a - b, and returns v.
|
||||||
|
func (v *Element) Subtract(a, b *Element) *Element {
|
||||||
|
// We first add 2 * p, to guarantee the subtraction won't underflow, and
|
||||||
|
// then subtract b (which can be up to 2^255 + 2^13 * 19).
|
||||||
|
v.l0 = (a.l0 + 0xFFFFFFFFFFFDA) - b.l0
|
||||||
|
v.l1 = (a.l1 + 0xFFFFFFFFFFFFE) - b.l1
|
||||||
|
v.l2 = (a.l2 + 0xFFFFFFFFFFFFE) - b.l2
|
||||||
|
v.l3 = (a.l3 + 0xFFFFFFFFFFFFE) - b.l3
|
||||||
|
v.l4 = (a.l4 + 0xFFFFFFFFFFFFE) - b.l4
|
||||||
|
return v.carryPropagate()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Negate sets v = -a, and returns v.
|
||||||
|
func (v *Element) Negate(a *Element) *Element {
|
||||||
|
return v.Subtract(feZero, a)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Invert sets v = 1/z mod p, and returns v.
|
||||||
|
//
|
||||||
|
// If z == 0, Invert returns v = 0.
|
||||||
|
func (v *Element) Invert(z *Element) *Element {
|
||||||
|
// Inversion is implemented as exponentiation with exponent p − 2. It uses the
|
||||||
|
// same sequence of 255 squarings and 11 multiplications as [Curve25519].
|
||||||
|
var z2, z9, z11, z2_5_0, z2_10_0, z2_20_0, z2_50_0, z2_100_0, t Element
|
||||||
|
|
||||||
|
z2.Square(z) // 2
|
||||||
|
t.Square(&z2) // 4
|
||||||
|
t.Square(&t) // 8
|
||||||
|
z9.Multiply(&t, z) // 9
|
||||||
|
z11.Multiply(&z9, &z2) // 11
|
||||||
|
t.Square(&z11) // 22
|
||||||
|
z2_5_0.Multiply(&t, &z9) // 31 = 2^5 - 2^0
|
||||||
|
|
||||||
|
t.Square(&z2_5_0) // 2^6 - 2^1
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
t.Square(&t) // 2^10 - 2^5
|
||||||
|
}
|
||||||
|
z2_10_0.Multiply(&t, &z2_5_0) // 2^10 - 2^0
|
||||||
|
|
||||||
|
t.Square(&z2_10_0) // 2^11 - 2^1
|
||||||
|
for i := 0; i < 9; i++ {
|
||||||
|
t.Square(&t) // 2^20 - 2^10
|
||||||
|
}
|
||||||
|
z2_20_0.Multiply(&t, &z2_10_0) // 2^20 - 2^0
|
||||||
|
|
||||||
|
t.Square(&z2_20_0) // 2^21 - 2^1
|
||||||
|
for i := 0; i < 19; i++ {
|
||||||
|
t.Square(&t) // 2^40 - 2^20
|
||||||
|
}
|
||||||
|
t.Multiply(&t, &z2_20_0) // 2^40 - 2^0
|
||||||
|
|
||||||
|
t.Square(&t) // 2^41 - 2^1
|
||||||
|
for i := 0; i < 9; i++ {
|
||||||
|
t.Square(&t) // 2^50 - 2^10
|
||||||
|
}
|
||||||
|
z2_50_0.Multiply(&t, &z2_10_0) // 2^50 - 2^0
|
||||||
|
|
||||||
|
t.Square(&z2_50_0) // 2^51 - 2^1
|
||||||
|
for i := 0; i < 49; i++ {
|
||||||
|
t.Square(&t) // 2^100 - 2^50
|
||||||
|
}
|
||||||
|
z2_100_0.Multiply(&t, &z2_50_0) // 2^100 - 2^0
|
||||||
|
|
||||||
|
t.Square(&z2_100_0) // 2^101 - 2^1
|
||||||
|
for i := 0; i < 99; i++ {
|
||||||
|
t.Square(&t) // 2^200 - 2^100
|
||||||
|
}
|
||||||
|
t.Multiply(&t, &z2_100_0) // 2^200 - 2^0
|
||||||
|
|
||||||
|
t.Square(&t) // 2^201 - 2^1
|
||||||
|
for i := 0; i < 49; i++ {
|
||||||
|
t.Square(&t) // 2^250 - 2^50
|
||||||
|
}
|
||||||
|
t.Multiply(&t, &z2_50_0) // 2^250 - 2^0
|
||||||
|
|
||||||
|
t.Square(&t) // 2^251 - 2^1
|
||||||
|
t.Square(&t) // 2^252 - 2^2
|
||||||
|
t.Square(&t) // 2^253 - 2^3
|
||||||
|
t.Square(&t) // 2^254 - 2^4
|
||||||
|
t.Square(&t) // 2^255 - 2^5
|
||||||
|
|
||||||
|
return v.Multiply(&t, &z11) // 2^255 - 21
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set sets v = a, and returns v.
|
||||||
|
func (v *Element) Set(a *Element) *Element {
|
||||||
|
*v = *a
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetBytes sets v to x, where x is a 32-byte little-endian encoding. If x is
|
||||||
|
// not of the right length, SetBytes returns nil and an error, and the
|
||||||
|
// receiver is unchanged.
|
||||||
|
//
|
||||||
|
// Consistent with RFC 7748, the most significant bit (the high bit of the
|
||||||
|
// last byte) is ignored, and non-canonical values (2^255-19 through 2^255-1)
|
||||||
|
// are accepted. Note that this is laxer than specified by RFC 8032, but
|
||||||
|
// consistent with most Ed25519 implementations.
|
||||||
|
func (v *Element) SetBytes(x []byte) (*Element, error) {
|
||||||
|
if len(x) != 32 {
|
||||||
|
return nil, errors.New("edwards25519: invalid field element input size")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bits 0:51 (bytes 0:8, bits 0:64, shift 0, mask 51).
|
||||||
|
v.l0 = binary.LittleEndian.Uint64(x[0:8])
|
||||||
|
v.l0 &= maskLow51Bits
|
||||||
|
// Bits 51:102 (bytes 6:14, bits 48:112, shift 3, mask 51).
|
||||||
|
v.l1 = binary.LittleEndian.Uint64(x[6:14]) >> 3
|
||||||
|
v.l1 &= maskLow51Bits
|
||||||
|
// Bits 102:153 (bytes 12:20, bits 96:160, shift 6, mask 51).
|
||||||
|
v.l2 = binary.LittleEndian.Uint64(x[12:20]) >> 6
|
||||||
|
v.l2 &= maskLow51Bits
|
||||||
|
// Bits 153:204 (bytes 19:27, bits 152:216, shift 1, mask 51).
|
||||||
|
v.l3 = binary.LittleEndian.Uint64(x[19:27]) >> 1
|
||||||
|
v.l3 &= maskLow51Bits
|
||||||
|
// Bits 204:255 (bytes 24:32, bits 192:256, shift 12, mask 51).
|
||||||
|
// Note: not bytes 25:33, shift 4, to avoid overread.
|
||||||
|
v.l4 = binary.LittleEndian.Uint64(x[24:32]) >> 12
|
||||||
|
v.l4 &= maskLow51Bits
|
||||||
|
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bytes returns the canonical 32-byte little-endian encoding of v.
|
||||||
|
func (v *Element) Bytes() []byte {
|
||||||
|
// This function is outlined to make the allocations inline in the caller
|
||||||
|
// rather than happen on the heap.
|
||||||
|
var out [32]byte
|
||||||
|
return v.bytes(&out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *Element) bytes(out *[32]byte) []byte {
|
||||||
|
t := *v
|
||||||
|
t.reduce()
|
||||||
|
|
||||||
|
var buf [8]byte
|
||||||
|
for i, l := range [5]uint64{t.l0, t.l1, t.l2, t.l3, t.l4} {
|
||||||
|
bitsOffset := i * 51
|
||||||
|
binary.LittleEndian.PutUint64(buf[:], l<<uint(bitsOffset%8))
|
||||||
|
for i, bb := range buf {
|
||||||
|
off := bitsOffset/8 + i
|
||||||
|
if off >= len(out) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
out[off] |= bb
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return out[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Equal returns 1 if v and u are equal, and 0 otherwise.
|
||||||
|
func (v *Element) Equal(u *Element) int {
|
||||||
|
sa, sv := u.Bytes(), v.Bytes()
|
||||||
|
return subtle.ConstantTimeCompare(sa, sv)
|
||||||
|
}
|
||||||
|
|
||||||
|
// mask64Bits returns 0xffffffff if cond is 1, and 0 otherwise.
|
||||||
|
func mask64Bits(cond int) uint64 { return ^(uint64(cond) - 1) }
|
||||||
|
|
||||||
|
// Select sets v to a if cond == 1, and to b if cond == 0.
|
||||||
|
func (v *Element) Select(a, b *Element, cond int) *Element {
|
||||||
|
m := mask64Bits(cond)
|
||||||
|
v.l0 = (m & a.l0) | (^m & b.l0)
|
||||||
|
v.l1 = (m & a.l1) | (^m & b.l1)
|
||||||
|
v.l2 = (m & a.l2) | (^m & b.l2)
|
||||||
|
v.l3 = (m & a.l3) | (^m & b.l3)
|
||||||
|
v.l4 = (m & a.l4) | (^m & b.l4)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Swap swaps v and u if cond == 1 or leaves them unchanged if cond == 0, and returns v.
|
||||||
|
func (v *Element) Swap(u *Element, cond int) {
|
||||||
|
m := mask64Bits(cond)
|
||||||
|
t := m & (v.l0 ^ u.l0)
|
||||||
|
v.l0 ^= t
|
||||||
|
u.l0 ^= t
|
||||||
|
t = m & (v.l1 ^ u.l1)
|
||||||
|
v.l1 ^= t
|
||||||
|
u.l1 ^= t
|
||||||
|
t = m & (v.l2 ^ u.l2)
|
||||||
|
v.l2 ^= t
|
||||||
|
u.l2 ^= t
|
||||||
|
t = m & (v.l3 ^ u.l3)
|
||||||
|
v.l3 ^= t
|
||||||
|
u.l3 ^= t
|
||||||
|
t = m & (v.l4 ^ u.l4)
|
||||||
|
v.l4 ^= t
|
||||||
|
u.l4 ^= t
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsNegative returns 1 if v is negative, and 0 otherwise.
|
||||||
|
func (v *Element) IsNegative() int {
|
||||||
|
return int(v.Bytes()[0] & 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Absolute sets v to |u|, and returns v.
|
||||||
|
func (v *Element) Absolute(u *Element) *Element {
|
||||||
|
return v.Select(new(Element).Negate(u), u, u.IsNegative())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiply sets v = x * y, and returns v.
|
||||||
|
func (v *Element) Multiply(x, y *Element) *Element {
|
||||||
|
feMul(v, x, y)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Square sets v = x * x, and returns v.
|
||||||
|
func (v *Element) Square(x *Element) *Element {
|
||||||
|
feSquare(v, x)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mult32 sets v = x * y, and returns v.
|
||||||
|
func (v *Element) Mult32(x *Element, y uint32) *Element {
|
||||||
|
x0lo, x0hi := mul51(x.l0, y)
|
||||||
|
x1lo, x1hi := mul51(x.l1, y)
|
||||||
|
x2lo, x2hi := mul51(x.l2, y)
|
||||||
|
x3lo, x3hi := mul51(x.l3, y)
|
||||||
|
x4lo, x4hi := mul51(x.l4, y)
|
||||||
|
v.l0 = x0lo + 19*x4hi // carried over per the reduction identity
|
||||||
|
v.l1 = x1lo + x0hi
|
||||||
|
v.l2 = x2lo + x1hi
|
||||||
|
v.l3 = x3lo + x2hi
|
||||||
|
v.l4 = x4lo + x3hi
|
||||||
|
// The hi portions are going to be only 32 bits, plus any previous excess,
|
||||||
|
// so we can skip the carry propagation.
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// mul51 returns lo + hi * 2⁵¹ = a * b.
|
||||||
|
func mul51(a uint64, b uint32) (lo uint64, hi uint64) {
|
||||||
|
mh, ml := bits.Mul64(a, uint64(b))
|
||||||
|
lo = ml & maskLow51Bits
|
||||||
|
hi = (mh << 13) | (ml >> 51)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pow22523 set v = x^((p-5)/8), and returns v. (p-5)/8 is 2^252-3.
|
||||||
|
func (v *Element) Pow22523(x *Element) *Element {
|
||||||
|
var t0, t1, t2 Element
|
||||||
|
|
||||||
|
t0.Square(x) // x^2
|
||||||
|
t1.Square(&t0) // x^4
|
||||||
|
t1.Square(&t1) // x^8
|
||||||
|
t1.Multiply(x, &t1) // x^9
|
||||||
|
t0.Multiply(&t0, &t1) // x^11
|
||||||
|
t0.Square(&t0) // x^22
|
||||||
|
t0.Multiply(&t1, &t0) // x^31
|
||||||
|
t1.Square(&t0) // x^62
|
||||||
|
for i := 1; i < 5; i++ { // x^992
|
||||||
|
t1.Square(&t1)
|
||||||
|
}
|
||||||
|
t0.Multiply(&t1, &t0) // x^1023 -> 1023 = 2^10 - 1
|
||||||
|
t1.Square(&t0) // 2^11 - 2
|
||||||
|
for i := 1; i < 10; i++ { // 2^20 - 2^10
|
||||||
|
t1.Square(&t1)
|
||||||
|
}
|
||||||
|
t1.Multiply(&t1, &t0) // 2^20 - 1
|
||||||
|
t2.Square(&t1) // 2^21 - 2
|
||||||
|
for i := 1; i < 20; i++ { // 2^40 - 2^20
|
||||||
|
t2.Square(&t2)
|
||||||
|
}
|
||||||
|
t1.Multiply(&t2, &t1) // 2^40 - 1
|
||||||
|
t1.Square(&t1) // 2^41 - 2
|
||||||
|
for i := 1; i < 10; i++ { // 2^50 - 2^10
|
||||||
|
t1.Square(&t1)
|
||||||
|
}
|
||||||
|
t0.Multiply(&t1, &t0) // 2^50 - 1
|
||||||
|
t1.Square(&t0) // 2^51 - 2
|
||||||
|
for i := 1; i < 50; i++ { // 2^100 - 2^50
|
||||||
|
t1.Square(&t1)
|
||||||
|
}
|
||||||
|
t1.Multiply(&t1, &t0) // 2^100 - 1
|
||||||
|
t2.Square(&t1) // 2^101 - 2
|
||||||
|
for i := 1; i < 100; i++ { // 2^200 - 2^100
|
||||||
|
t2.Square(&t2)
|
||||||
|
}
|
||||||
|
t1.Multiply(&t2, &t1) // 2^200 - 1
|
||||||
|
t1.Square(&t1) // 2^201 - 2
|
||||||
|
for i := 1; i < 50; i++ { // 2^250 - 2^50
|
||||||
|
t1.Square(&t1)
|
||||||
|
}
|
||||||
|
t0.Multiply(&t1, &t0) // 2^250 - 1
|
||||||
|
t0.Square(&t0) // 2^251 - 2
|
||||||
|
t0.Square(&t0) // 2^252 - 4
|
||||||
|
return v.Multiply(&t0, x) // 2^252 - 3 -> x^(2^252-3)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sqrtM1 is 2^((p-1)/4), which squared is equal to -1 by Euler's Criterion.
|
||||||
|
var sqrtM1 = &Element{1718705420411056, 234908883556509,
|
||||||
|
2233514472574048, 2117202627021982, 765476049583133}
|
||||||
|
|
||||||
|
// SqrtRatio sets r to the non-negative square root of the ratio of u and v.
|
||||||
|
//
|
||||||
|
// If u/v is square, SqrtRatio returns r and 1. If u/v is not square, SqrtRatio
|
||||||
|
// sets r according to Section 4.3 of draft-irtf-cfrg-ristretto255-decaf448-00,
|
||||||
|
// and returns r and 0.
|
||||||
|
func (r *Element) SqrtRatio(u, v *Element) (R *Element, wasSquare int) {
|
||||||
|
t0 := new(Element)
|
||||||
|
|
||||||
|
// r = (u * v3) * (u * v7)^((p-5)/8)
|
||||||
|
v2 := new(Element).Square(v)
|
||||||
|
uv3 := new(Element).Multiply(u, t0.Multiply(v2, v))
|
||||||
|
uv7 := new(Element).Multiply(uv3, t0.Square(v2))
|
||||||
|
rr := new(Element).Multiply(uv3, t0.Pow22523(uv7))
|
||||||
|
|
||||||
|
check := new(Element).Multiply(v, t0.Square(rr)) // check = v * r^2
|
||||||
|
|
||||||
|
uNeg := new(Element).Negate(u)
|
||||||
|
correctSignSqrt := check.Equal(u)
|
||||||
|
flippedSignSqrt := check.Equal(uNeg)
|
||||||
|
flippedSignSqrtI := check.Equal(t0.Multiply(uNeg, sqrtM1))
|
||||||
|
|
||||||
|
rPrime := new(Element).Multiply(rr, sqrtM1) // r_prime = SQRT_M1 * r
|
||||||
|
// r = CT_SELECT(r_prime IF flipped_sign_sqrt | flipped_sign_sqrt_i ELSE r)
|
||||||
|
rr.Select(rPrime, rr, flippedSignSqrt|flippedSignSqrtI)
|
||||||
|
|
||||||
|
r.Absolute(rr) // Choose the nonnegative square root.
|
||||||
|
return r, correctSignSqrt | flippedSignSqrt
|
||||||
|
}
|
||||||
+16
@@ -0,0 +1,16 @@
|
|||||||
|
// Code generated by command: go run fe_amd64_asm.go -out ../fe_amd64.s -stubs ../fe_amd64.go -pkg field. DO NOT EDIT.
|
||||||
|
|
||||||
|
//go:build amd64 && gc && !purego
|
||||||
|
// +build amd64,gc,!purego
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
// feMul sets out = a * b. It works like feMulGeneric.
|
||||||
|
//
|
||||||
|
//go:noescape
|
||||||
|
func feMul(out *Element, a *Element, b *Element)
|
||||||
|
|
||||||
|
// feSquare sets out = a * a. It works like feSquareGeneric.
|
||||||
|
//
|
||||||
|
//go:noescape
|
||||||
|
func feSquare(out *Element, a *Element)
|
||||||
+379
@@ -0,0 +1,379 @@
|
|||||||
|
// Code generated by command: go run fe_amd64_asm.go -out ../fe_amd64.s -stubs ../fe_amd64.go -pkg field. DO NOT EDIT.
|
||||||
|
|
||||||
|
//go:build amd64 && gc && !purego
|
||||||
|
// +build amd64,gc,!purego
|
||||||
|
|
||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// func feMul(out *Element, a *Element, b *Element)
|
||||||
|
TEXT ·feMul(SB), NOSPLIT, $0-24
|
||||||
|
MOVQ a+8(FP), CX
|
||||||
|
MOVQ b+16(FP), BX
|
||||||
|
|
||||||
|
// r0 = a0×b0
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ (BX)
|
||||||
|
MOVQ AX, DI
|
||||||
|
MOVQ DX, SI
|
||||||
|
|
||||||
|
// r0 += 19×a1×b4
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 32(BX)
|
||||||
|
ADDQ AX, DI
|
||||||
|
ADCQ DX, SI
|
||||||
|
|
||||||
|
// r0 += 19×a2×b3
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 24(BX)
|
||||||
|
ADDQ AX, DI
|
||||||
|
ADCQ DX, SI
|
||||||
|
|
||||||
|
// r0 += 19×a3×b2
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 16(BX)
|
||||||
|
ADDQ AX, DI
|
||||||
|
ADCQ DX, SI
|
||||||
|
|
||||||
|
// r0 += 19×a4×b1
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 8(BX)
|
||||||
|
ADDQ AX, DI
|
||||||
|
ADCQ DX, SI
|
||||||
|
|
||||||
|
// r1 = a0×b1
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ 8(BX)
|
||||||
|
MOVQ AX, R9
|
||||||
|
MOVQ DX, R8
|
||||||
|
|
||||||
|
// r1 += a1×b0
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
MULQ (BX)
|
||||||
|
ADDQ AX, R9
|
||||||
|
ADCQ DX, R8
|
||||||
|
|
||||||
|
// r1 += 19×a2×b4
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 32(BX)
|
||||||
|
ADDQ AX, R9
|
||||||
|
ADCQ DX, R8
|
||||||
|
|
||||||
|
// r1 += 19×a3×b3
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 24(BX)
|
||||||
|
ADDQ AX, R9
|
||||||
|
ADCQ DX, R8
|
||||||
|
|
||||||
|
// r1 += 19×a4×b2
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 16(BX)
|
||||||
|
ADDQ AX, R9
|
||||||
|
ADCQ DX, R8
|
||||||
|
|
||||||
|
// r2 = a0×b2
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ 16(BX)
|
||||||
|
MOVQ AX, R11
|
||||||
|
MOVQ DX, R10
|
||||||
|
|
||||||
|
// r2 += a1×b1
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
MULQ 8(BX)
|
||||||
|
ADDQ AX, R11
|
||||||
|
ADCQ DX, R10
|
||||||
|
|
||||||
|
// r2 += a2×b0
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
MULQ (BX)
|
||||||
|
ADDQ AX, R11
|
||||||
|
ADCQ DX, R10
|
||||||
|
|
||||||
|
// r2 += 19×a3×b4
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 32(BX)
|
||||||
|
ADDQ AX, R11
|
||||||
|
ADCQ DX, R10
|
||||||
|
|
||||||
|
// r2 += 19×a4×b3
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 24(BX)
|
||||||
|
ADDQ AX, R11
|
||||||
|
ADCQ DX, R10
|
||||||
|
|
||||||
|
// r3 = a0×b3
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ 24(BX)
|
||||||
|
MOVQ AX, R13
|
||||||
|
MOVQ DX, R12
|
||||||
|
|
||||||
|
// r3 += a1×b2
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
MULQ 16(BX)
|
||||||
|
ADDQ AX, R13
|
||||||
|
ADCQ DX, R12
|
||||||
|
|
||||||
|
// r3 += a2×b1
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
MULQ 8(BX)
|
||||||
|
ADDQ AX, R13
|
||||||
|
ADCQ DX, R12
|
||||||
|
|
||||||
|
// r3 += a3×b0
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
MULQ (BX)
|
||||||
|
ADDQ AX, R13
|
||||||
|
ADCQ DX, R12
|
||||||
|
|
||||||
|
// r3 += 19×a4×b4
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 32(BX)
|
||||||
|
ADDQ AX, R13
|
||||||
|
ADCQ DX, R12
|
||||||
|
|
||||||
|
// r4 = a0×b4
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ 32(BX)
|
||||||
|
MOVQ AX, R15
|
||||||
|
MOVQ DX, R14
|
||||||
|
|
||||||
|
// r4 += a1×b3
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
MULQ 24(BX)
|
||||||
|
ADDQ AX, R15
|
||||||
|
ADCQ DX, R14
|
||||||
|
|
||||||
|
// r4 += a2×b2
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
MULQ 16(BX)
|
||||||
|
ADDQ AX, R15
|
||||||
|
ADCQ DX, R14
|
||||||
|
|
||||||
|
// r4 += a3×b1
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
MULQ 8(BX)
|
||||||
|
ADDQ AX, R15
|
||||||
|
ADCQ DX, R14
|
||||||
|
|
||||||
|
// r4 += a4×b0
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
MULQ (BX)
|
||||||
|
ADDQ AX, R15
|
||||||
|
ADCQ DX, R14
|
||||||
|
|
||||||
|
// First reduction chain
|
||||||
|
MOVQ $0x0007ffffffffffff, AX
|
||||||
|
SHLQ $0x0d, DI, SI
|
||||||
|
SHLQ $0x0d, R9, R8
|
||||||
|
SHLQ $0x0d, R11, R10
|
||||||
|
SHLQ $0x0d, R13, R12
|
||||||
|
SHLQ $0x0d, R15, R14
|
||||||
|
ANDQ AX, DI
|
||||||
|
IMUL3Q $0x13, R14, R14
|
||||||
|
ADDQ R14, DI
|
||||||
|
ANDQ AX, R9
|
||||||
|
ADDQ SI, R9
|
||||||
|
ANDQ AX, R11
|
||||||
|
ADDQ R8, R11
|
||||||
|
ANDQ AX, R13
|
||||||
|
ADDQ R10, R13
|
||||||
|
ANDQ AX, R15
|
||||||
|
ADDQ R12, R15
|
||||||
|
|
||||||
|
// Second reduction chain (carryPropagate)
|
||||||
|
MOVQ DI, SI
|
||||||
|
SHRQ $0x33, SI
|
||||||
|
MOVQ R9, R8
|
||||||
|
SHRQ $0x33, R8
|
||||||
|
MOVQ R11, R10
|
||||||
|
SHRQ $0x33, R10
|
||||||
|
MOVQ R13, R12
|
||||||
|
SHRQ $0x33, R12
|
||||||
|
MOVQ R15, R14
|
||||||
|
SHRQ $0x33, R14
|
||||||
|
ANDQ AX, DI
|
||||||
|
IMUL3Q $0x13, R14, R14
|
||||||
|
ADDQ R14, DI
|
||||||
|
ANDQ AX, R9
|
||||||
|
ADDQ SI, R9
|
||||||
|
ANDQ AX, R11
|
||||||
|
ADDQ R8, R11
|
||||||
|
ANDQ AX, R13
|
||||||
|
ADDQ R10, R13
|
||||||
|
ANDQ AX, R15
|
||||||
|
ADDQ R12, R15
|
||||||
|
|
||||||
|
// Store output
|
||||||
|
MOVQ out+0(FP), AX
|
||||||
|
MOVQ DI, (AX)
|
||||||
|
MOVQ R9, 8(AX)
|
||||||
|
MOVQ R11, 16(AX)
|
||||||
|
MOVQ R13, 24(AX)
|
||||||
|
MOVQ R15, 32(AX)
|
||||||
|
RET
|
||||||
|
|
||||||
|
// func feSquare(out *Element, a *Element)
|
||||||
|
TEXT ·feSquare(SB), NOSPLIT, $0-16
|
||||||
|
MOVQ a+8(FP), CX
|
||||||
|
|
||||||
|
// r0 = l0×l0
|
||||||
|
MOVQ (CX), AX
|
||||||
|
MULQ (CX)
|
||||||
|
MOVQ AX, SI
|
||||||
|
MOVQ DX, BX
|
||||||
|
|
||||||
|
// r0 += 38×l1×l4
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
IMUL3Q $0x26, AX, AX
|
||||||
|
MULQ 32(CX)
|
||||||
|
ADDQ AX, SI
|
||||||
|
ADCQ DX, BX
|
||||||
|
|
||||||
|
// r0 += 38×l2×l3
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
IMUL3Q $0x26, AX, AX
|
||||||
|
MULQ 24(CX)
|
||||||
|
ADDQ AX, SI
|
||||||
|
ADCQ DX, BX
|
||||||
|
|
||||||
|
// r1 = 2×l0×l1
|
||||||
|
MOVQ (CX), AX
|
||||||
|
SHLQ $0x01, AX
|
||||||
|
MULQ 8(CX)
|
||||||
|
MOVQ AX, R8
|
||||||
|
MOVQ DX, DI
|
||||||
|
|
||||||
|
// r1 += 38×l2×l4
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
IMUL3Q $0x26, AX, AX
|
||||||
|
MULQ 32(CX)
|
||||||
|
ADDQ AX, R8
|
||||||
|
ADCQ DX, DI
|
||||||
|
|
||||||
|
// r1 += 19×l3×l3
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 24(CX)
|
||||||
|
ADDQ AX, R8
|
||||||
|
ADCQ DX, DI
|
||||||
|
|
||||||
|
// r2 = 2×l0×l2
|
||||||
|
MOVQ (CX), AX
|
||||||
|
SHLQ $0x01, AX
|
||||||
|
MULQ 16(CX)
|
||||||
|
MOVQ AX, R10
|
||||||
|
MOVQ DX, R9
|
||||||
|
|
||||||
|
// r2 += l1×l1
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
MULQ 8(CX)
|
||||||
|
ADDQ AX, R10
|
||||||
|
ADCQ DX, R9
|
||||||
|
|
||||||
|
// r2 += 38×l3×l4
|
||||||
|
MOVQ 24(CX), AX
|
||||||
|
IMUL3Q $0x26, AX, AX
|
||||||
|
MULQ 32(CX)
|
||||||
|
ADDQ AX, R10
|
||||||
|
ADCQ DX, R9
|
||||||
|
|
||||||
|
// r3 = 2×l0×l3
|
||||||
|
MOVQ (CX), AX
|
||||||
|
SHLQ $0x01, AX
|
||||||
|
MULQ 24(CX)
|
||||||
|
MOVQ AX, R12
|
||||||
|
MOVQ DX, R11
|
||||||
|
|
||||||
|
// r3 += 2×l1×l2
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
IMUL3Q $0x02, AX, AX
|
||||||
|
MULQ 16(CX)
|
||||||
|
ADDQ AX, R12
|
||||||
|
ADCQ DX, R11
|
||||||
|
|
||||||
|
// r3 += 19×l4×l4
|
||||||
|
MOVQ 32(CX), AX
|
||||||
|
IMUL3Q $0x13, AX, AX
|
||||||
|
MULQ 32(CX)
|
||||||
|
ADDQ AX, R12
|
||||||
|
ADCQ DX, R11
|
||||||
|
|
||||||
|
// r4 = 2×l0×l4
|
||||||
|
MOVQ (CX), AX
|
||||||
|
SHLQ $0x01, AX
|
||||||
|
MULQ 32(CX)
|
||||||
|
MOVQ AX, R14
|
||||||
|
MOVQ DX, R13
|
||||||
|
|
||||||
|
// r4 += 2×l1×l3
|
||||||
|
MOVQ 8(CX), AX
|
||||||
|
IMUL3Q $0x02, AX, AX
|
||||||
|
MULQ 24(CX)
|
||||||
|
ADDQ AX, R14
|
||||||
|
ADCQ DX, R13
|
||||||
|
|
||||||
|
// r4 += l2×l2
|
||||||
|
MOVQ 16(CX), AX
|
||||||
|
MULQ 16(CX)
|
||||||
|
ADDQ AX, R14
|
||||||
|
ADCQ DX, R13
|
||||||
|
|
||||||
|
// First reduction chain
|
||||||
|
MOVQ $0x0007ffffffffffff, AX
|
||||||
|
SHLQ $0x0d, SI, BX
|
||||||
|
SHLQ $0x0d, R8, DI
|
||||||
|
SHLQ $0x0d, R10, R9
|
||||||
|
SHLQ $0x0d, R12, R11
|
||||||
|
SHLQ $0x0d, R14, R13
|
||||||
|
ANDQ AX, SI
|
||||||
|
IMUL3Q $0x13, R13, R13
|
||||||
|
ADDQ R13, SI
|
||||||
|
ANDQ AX, R8
|
||||||
|
ADDQ BX, R8
|
||||||
|
ANDQ AX, R10
|
||||||
|
ADDQ DI, R10
|
||||||
|
ANDQ AX, R12
|
||||||
|
ADDQ R9, R12
|
||||||
|
ANDQ AX, R14
|
||||||
|
ADDQ R11, R14
|
||||||
|
|
||||||
|
// Second reduction chain (carryPropagate)
|
||||||
|
MOVQ SI, BX
|
||||||
|
SHRQ $0x33, BX
|
||||||
|
MOVQ R8, DI
|
||||||
|
SHRQ $0x33, DI
|
||||||
|
MOVQ R10, R9
|
||||||
|
SHRQ $0x33, R9
|
||||||
|
MOVQ R12, R11
|
||||||
|
SHRQ $0x33, R11
|
||||||
|
MOVQ R14, R13
|
||||||
|
SHRQ $0x33, R13
|
||||||
|
ANDQ AX, SI
|
||||||
|
IMUL3Q $0x13, R13, R13
|
||||||
|
ADDQ R13, SI
|
||||||
|
ANDQ AX, R8
|
||||||
|
ADDQ BX, R8
|
||||||
|
ANDQ AX, R10
|
||||||
|
ADDQ DI, R10
|
||||||
|
ANDQ AX, R12
|
||||||
|
ADDQ R9, R12
|
||||||
|
ANDQ AX, R14
|
||||||
|
ADDQ R11, R14
|
||||||
|
|
||||||
|
// Store output
|
||||||
|
MOVQ out+0(FP), AX
|
||||||
|
MOVQ SI, (AX)
|
||||||
|
MOVQ R8, 8(AX)
|
||||||
|
MOVQ R10, 16(AX)
|
||||||
|
MOVQ R12, 24(AX)
|
||||||
|
MOVQ R14, 32(AX)
|
||||||
|
RET
|
||||||
+12
@@ -0,0 +1,12 @@
|
|||||||
|
// Copyright (c) 2019 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
//go:build !amd64 || !gc || purego
|
||||||
|
// +build !amd64 !gc purego
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
func feMul(v, x, y *Element) { feMulGeneric(v, x, y) }
|
||||||
|
|
||||||
|
func feSquare(v, x *Element) { feSquareGeneric(v, x) }
|
||||||
+16
@@ -0,0 +1,16 @@
|
|||||||
|
// Copyright (c) 2020 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
//go:build arm64 && gc && !purego
|
||||||
|
// +build arm64,gc,!purego
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
func carryPropagate(v *Element)
|
||||||
|
|
||||||
|
func (v *Element) carryPropagate() *Element {
|
||||||
|
carryPropagate(v)
|
||||||
|
return v
|
||||||
|
}
|
||||||
+42
@@ -0,0 +1,42 @@
|
|||||||
|
// Copyright (c) 2020 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
//go:build arm64 && gc && !purego
|
||||||
|
|
||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// carryPropagate works exactly like carryPropagateGeneric and uses the
|
||||||
|
// same AND, ADD, and LSR+MADD instructions emitted by the compiler, but
|
||||||
|
// avoids loading R0-R4 twice and uses LDP and STP.
|
||||||
|
//
|
||||||
|
// See https://golang.org/issues/43145 for the main compiler issue.
|
||||||
|
//
|
||||||
|
// func carryPropagate(v *Element)
|
||||||
|
TEXT ·carryPropagate(SB),NOFRAME|NOSPLIT,$0-8
|
||||||
|
MOVD v+0(FP), R20
|
||||||
|
|
||||||
|
LDP 0(R20), (R0, R1)
|
||||||
|
LDP 16(R20), (R2, R3)
|
||||||
|
MOVD 32(R20), R4
|
||||||
|
|
||||||
|
AND $0x7ffffffffffff, R0, R10
|
||||||
|
AND $0x7ffffffffffff, R1, R11
|
||||||
|
AND $0x7ffffffffffff, R2, R12
|
||||||
|
AND $0x7ffffffffffff, R3, R13
|
||||||
|
AND $0x7ffffffffffff, R4, R14
|
||||||
|
|
||||||
|
ADD R0>>51, R11, R11
|
||||||
|
ADD R1>>51, R12, R12
|
||||||
|
ADD R2>>51, R13, R13
|
||||||
|
ADD R3>>51, R14, R14
|
||||||
|
// R4>>51 * 19 + R10 -> R10
|
||||||
|
LSR $51, R4, R21
|
||||||
|
MOVD $19, R22
|
||||||
|
MADD R22, R10, R21, R10
|
||||||
|
|
||||||
|
STP (R10, R11), 0(R20)
|
||||||
|
STP (R12, R13), 16(R20)
|
||||||
|
MOVD R14, 32(R20)
|
||||||
|
|
||||||
|
RET
|
||||||
+12
@@ -0,0 +1,12 @@
|
|||||||
|
// Copyright (c) 2021 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
//go:build !arm64 || !gc || purego
|
||||||
|
// +build !arm64 !gc purego
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
func (v *Element) carryPropagate() *Element {
|
||||||
|
return v.carryPropagateGeneric()
|
||||||
|
}
|
||||||
+50
@@ -0,0 +1,50 @@
|
|||||||
|
// Copyright (c) 2021 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
import "errors"
|
||||||
|
|
||||||
|
// This file contains additional functionality that is not included in the
|
||||||
|
// upstream crypto/ed25519/edwards25519/field package.
|
||||||
|
|
||||||
|
// SetWideBytes sets v to x, where x is a 64-byte little-endian encoding, which
|
||||||
|
// is reduced modulo the field order. If x is not of the right length,
|
||||||
|
// SetWideBytes returns nil and an error, and the receiver is unchanged.
|
||||||
|
//
|
||||||
|
// SetWideBytes is not necessary to select a uniformly distributed value, and is
|
||||||
|
// only provided for compatibility: SetBytes can be used instead as the chance
|
||||||
|
// of bias is less than 2⁻²⁵⁰.
|
||||||
|
func (v *Element) SetWideBytes(x []byte) (*Element, error) {
|
||||||
|
if len(x) != 64 {
|
||||||
|
return nil, errors.New("edwards25519: invalid SetWideBytes input size")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Split the 64 bytes into two elements, and extract the most significant
|
||||||
|
// bit of each, which is ignored by SetBytes.
|
||||||
|
lo, _ := new(Element).SetBytes(x[:32])
|
||||||
|
loMSB := uint64(x[31] >> 7)
|
||||||
|
hi, _ := new(Element).SetBytes(x[32:])
|
||||||
|
hiMSB := uint64(x[63] >> 7)
|
||||||
|
|
||||||
|
// The output we want is
|
||||||
|
//
|
||||||
|
// v = lo + loMSB * 2²⁵⁵ + hi * 2²⁵⁶ + hiMSB * 2⁵¹¹
|
||||||
|
//
|
||||||
|
// which applying the reduction identity comes out to
|
||||||
|
//
|
||||||
|
// v = lo + loMSB * 19 + hi * 2 * 19 + hiMSB * 2 * 19²
|
||||||
|
//
|
||||||
|
// l0 will be the sum of a 52 bits value (lo.l0), plus a 5 bits value
|
||||||
|
// (loMSB * 19), a 6 bits value (hi.l0 * 2 * 19), and a 10 bits value
|
||||||
|
// (hiMSB * 2 * 19²), so it fits in a uint64.
|
||||||
|
|
||||||
|
v.l0 = lo.l0 + loMSB*19 + hi.l0*2*19 + hiMSB*2*19*19
|
||||||
|
v.l1 = lo.l1 + hi.l1*2*19
|
||||||
|
v.l2 = lo.l2 + hi.l2*2*19
|
||||||
|
v.l3 = lo.l3 + hi.l3*2*19
|
||||||
|
v.l4 = lo.l4 + hi.l4*2*19
|
||||||
|
|
||||||
|
return v.carryPropagate(), nil
|
||||||
|
}
|
||||||
+266
@@ -0,0 +1,266 @@
|
|||||||
|
// Copyright (c) 2017 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package field
|
||||||
|
|
||||||
|
import "math/bits"
|
||||||
|
|
||||||
|
// uint128 holds a 128-bit number as two 64-bit limbs, for use with the
|
||||||
|
// bits.Mul64 and bits.Add64 intrinsics.
|
||||||
|
type uint128 struct {
|
||||||
|
lo, hi uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
// mul64 returns a * b.
|
||||||
|
func mul64(a, b uint64) uint128 {
|
||||||
|
hi, lo := bits.Mul64(a, b)
|
||||||
|
return uint128{lo, hi}
|
||||||
|
}
|
||||||
|
|
||||||
|
// addMul64 returns v + a * b.
|
||||||
|
func addMul64(v uint128, a, b uint64) uint128 {
|
||||||
|
hi, lo := bits.Mul64(a, b)
|
||||||
|
lo, c := bits.Add64(lo, v.lo, 0)
|
||||||
|
hi, _ = bits.Add64(hi, v.hi, c)
|
||||||
|
return uint128{lo, hi}
|
||||||
|
}
|
||||||
|
|
||||||
|
// shiftRightBy51 returns a >> 51. a is assumed to be at most 115 bits.
|
||||||
|
func shiftRightBy51(a uint128) uint64 {
|
||||||
|
return (a.hi << (64 - 51)) | (a.lo >> 51)
|
||||||
|
}
|
||||||
|
|
||||||
|
func feMulGeneric(v, a, b *Element) {
|
||||||
|
a0 := a.l0
|
||||||
|
a1 := a.l1
|
||||||
|
a2 := a.l2
|
||||||
|
a3 := a.l3
|
||||||
|
a4 := a.l4
|
||||||
|
|
||||||
|
b0 := b.l0
|
||||||
|
b1 := b.l1
|
||||||
|
b2 := b.l2
|
||||||
|
b3 := b.l3
|
||||||
|
b4 := b.l4
|
||||||
|
|
||||||
|
// Limb multiplication works like pen-and-paper columnar multiplication, but
|
||||||
|
// with 51-bit limbs instead of digits.
|
||||||
|
//
|
||||||
|
// a4 a3 a2 a1 a0 x
|
||||||
|
// b4 b3 b2 b1 b0 =
|
||||||
|
// ------------------------
|
||||||
|
// a4b0 a3b0 a2b0 a1b0 a0b0 +
|
||||||
|
// a4b1 a3b1 a2b1 a1b1 a0b1 +
|
||||||
|
// a4b2 a3b2 a2b2 a1b2 a0b2 +
|
||||||
|
// a4b3 a3b3 a2b3 a1b3 a0b3 +
|
||||||
|
// a4b4 a3b4 a2b4 a1b4 a0b4 =
|
||||||
|
// ----------------------------------------------
|
||||||
|
// r8 r7 r6 r5 r4 r3 r2 r1 r0
|
||||||
|
//
|
||||||
|
// We can then use the reduction identity (a * 2²⁵⁵ + b = a * 19 + b) to
|
||||||
|
// reduce the limbs that would overflow 255 bits. r5 * 2²⁵⁵ becomes 19 * r5,
|
||||||
|
// r6 * 2³⁰⁶ becomes 19 * r6 * 2⁵¹, etc.
|
||||||
|
//
|
||||||
|
// Reduction can be carried out simultaneously to multiplication. For
|
||||||
|
// example, we do not compute r5: whenever the result of a multiplication
|
||||||
|
// belongs to r5, like a1b4, we multiply it by 19 and add the result to r0.
|
||||||
|
//
|
||||||
|
// a4b0 a3b0 a2b0 a1b0 a0b0 +
|
||||||
|
// a3b1 a2b1 a1b1 a0b1 19×a4b1 +
|
||||||
|
// a2b2 a1b2 a0b2 19×a4b2 19×a3b2 +
|
||||||
|
// a1b3 a0b3 19×a4b3 19×a3b3 19×a2b3 +
|
||||||
|
// a0b4 19×a4b4 19×a3b4 19×a2b4 19×a1b4 =
|
||||||
|
// --------------------------------------
|
||||||
|
// r4 r3 r2 r1 r0
|
||||||
|
//
|
||||||
|
// Finally we add up the columns into wide, overlapping limbs.
|
||||||
|
|
||||||
|
a1_19 := a1 * 19
|
||||||
|
a2_19 := a2 * 19
|
||||||
|
a3_19 := a3 * 19
|
||||||
|
a4_19 := a4 * 19
|
||||||
|
|
||||||
|
// r0 = a0×b0 + 19×(a1×b4 + a2×b3 + a3×b2 + a4×b1)
|
||||||
|
r0 := mul64(a0, b0)
|
||||||
|
r0 = addMul64(r0, a1_19, b4)
|
||||||
|
r0 = addMul64(r0, a2_19, b3)
|
||||||
|
r0 = addMul64(r0, a3_19, b2)
|
||||||
|
r0 = addMul64(r0, a4_19, b1)
|
||||||
|
|
||||||
|
// r1 = a0×b1 + a1×b0 + 19×(a2×b4 + a3×b3 + a4×b2)
|
||||||
|
r1 := mul64(a0, b1)
|
||||||
|
r1 = addMul64(r1, a1, b0)
|
||||||
|
r1 = addMul64(r1, a2_19, b4)
|
||||||
|
r1 = addMul64(r1, a3_19, b3)
|
||||||
|
r1 = addMul64(r1, a4_19, b2)
|
||||||
|
|
||||||
|
// r2 = a0×b2 + a1×b1 + a2×b0 + 19×(a3×b4 + a4×b3)
|
||||||
|
r2 := mul64(a0, b2)
|
||||||
|
r2 = addMul64(r2, a1, b1)
|
||||||
|
r2 = addMul64(r2, a2, b0)
|
||||||
|
r2 = addMul64(r2, a3_19, b4)
|
||||||
|
r2 = addMul64(r2, a4_19, b3)
|
||||||
|
|
||||||
|
// r3 = a0×b3 + a1×b2 + a2×b1 + a3×b0 + 19×a4×b4
|
||||||
|
r3 := mul64(a0, b3)
|
||||||
|
r3 = addMul64(r3, a1, b2)
|
||||||
|
r3 = addMul64(r3, a2, b1)
|
||||||
|
r3 = addMul64(r3, a3, b0)
|
||||||
|
r3 = addMul64(r3, a4_19, b4)
|
||||||
|
|
||||||
|
// r4 = a0×b4 + a1×b3 + a2×b2 + a3×b1 + a4×b0
|
||||||
|
r4 := mul64(a0, b4)
|
||||||
|
r4 = addMul64(r4, a1, b3)
|
||||||
|
r4 = addMul64(r4, a2, b2)
|
||||||
|
r4 = addMul64(r4, a3, b1)
|
||||||
|
r4 = addMul64(r4, a4, b0)
|
||||||
|
|
||||||
|
// After the multiplication, we need to reduce (carry) the five coefficients
|
||||||
|
// to obtain a result with limbs that are at most slightly larger than 2⁵¹,
|
||||||
|
// to respect the Element invariant.
|
||||||
|
//
|
||||||
|
// Overall, the reduction works the same as carryPropagate, except with
|
||||||
|
// wider inputs: we take the carry for each coefficient by shifting it right
|
||||||
|
// by 51, and add it to the limb above it. The top carry is multiplied by 19
|
||||||
|
// according to the reduction identity and added to the lowest limb.
|
||||||
|
//
|
||||||
|
// The largest coefficient (r0) will be at most 111 bits, which guarantees
|
||||||
|
// that all carries are at most 111 - 51 = 60 bits, which fits in a uint64.
|
||||||
|
//
|
||||||
|
// r0 = a0×b0 + 19×(a1×b4 + a2×b3 + a3×b2 + a4×b1)
|
||||||
|
// r0 < 2⁵²×2⁵² + 19×(2⁵²×2⁵² + 2⁵²×2⁵² + 2⁵²×2⁵² + 2⁵²×2⁵²)
|
||||||
|
// r0 < (1 + 19 × 4) × 2⁵² × 2⁵²
|
||||||
|
// r0 < 2⁷ × 2⁵² × 2⁵²
|
||||||
|
// r0 < 2¹¹¹
|
||||||
|
//
|
||||||
|
// Moreover, the top coefficient (r4) is at most 107 bits, so c4 is at most
|
||||||
|
// 56 bits, and c4 * 19 is at most 61 bits, which again fits in a uint64 and
|
||||||
|
// allows us to easily apply the reduction identity.
|
||||||
|
//
|
||||||
|
// r4 = a0×b4 + a1×b3 + a2×b2 + a3×b1 + a4×b0
|
||||||
|
// r4 < 5 × 2⁵² × 2⁵²
|
||||||
|
// r4 < 2¹⁰⁷
|
||||||
|
//
|
||||||
|
|
||||||
|
c0 := shiftRightBy51(r0)
|
||||||
|
c1 := shiftRightBy51(r1)
|
||||||
|
c2 := shiftRightBy51(r2)
|
||||||
|
c3 := shiftRightBy51(r3)
|
||||||
|
c4 := shiftRightBy51(r4)
|
||||||
|
|
||||||
|
rr0 := r0.lo&maskLow51Bits + c4*19
|
||||||
|
rr1 := r1.lo&maskLow51Bits + c0
|
||||||
|
rr2 := r2.lo&maskLow51Bits + c1
|
||||||
|
rr3 := r3.lo&maskLow51Bits + c2
|
||||||
|
rr4 := r4.lo&maskLow51Bits + c3
|
||||||
|
|
||||||
|
// Now all coefficients fit into 64-bit registers but are still too large to
|
||||||
|
// be passed around as an Element. We therefore do one last carry chain,
|
||||||
|
// where the carries will be small enough to fit in the wiggle room above 2⁵¹.
|
||||||
|
*v = Element{rr0, rr1, rr2, rr3, rr4}
|
||||||
|
v.carryPropagate()
|
||||||
|
}
|
||||||
|
|
||||||
|
func feSquareGeneric(v, a *Element) {
|
||||||
|
l0 := a.l0
|
||||||
|
l1 := a.l1
|
||||||
|
l2 := a.l2
|
||||||
|
l3 := a.l3
|
||||||
|
l4 := a.l4
|
||||||
|
|
||||||
|
// Squaring works precisely like multiplication above, but thanks to its
|
||||||
|
// symmetry we get to group a few terms together.
|
||||||
|
//
|
||||||
|
// l4 l3 l2 l1 l0 x
|
||||||
|
// l4 l3 l2 l1 l0 =
|
||||||
|
// ------------------------
|
||||||
|
// l4l0 l3l0 l2l0 l1l0 l0l0 +
|
||||||
|
// l4l1 l3l1 l2l1 l1l1 l0l1 +
|
||||||
|
// l4l2 l3l2 l2l2 l1l2 l0l2 +
|
||||||
|
// l4l3 l3l3 l2l3 l1l3 l0l3 +
|
||||||
|
// l4l4 l3l4 l2l4 l1l4 l0l4 =
|
||||||
|
// ----------------------------------------------
|
||||||
|
// r8 r7 r6 r5 r4 r3 r2 r1 r0
|
||||||
|
//
|
||||||
|
// l4l0 l3l0 l2l0 l1l0 l0l0 +
|
||||||
|
// l3l1 l2l1 l1l1 l0l1 19×l4l1 +
|
||||||
|
// l2l2 l1l2 l0l2 19×l4l2 19×l3l2 +
|
||||||
|
// l1l3 l0l3 19×l4l3 19×l3l3 19×l2l3 +
|
||||||
|
// l0l4 19×l4l4 19×l3l4 19×l2l4 19×l1l4 =
|
||||||
|
// --------------------------------------
|
||||||
|
// r4 r3 r2 r1 r0
|
||||||
|
//
|
||||||
|
// With precomputed 2×, 19×, and 2×19× terms, we can compute each limb with
|
||||||
|
// only three Mul64 and four Add64, instead of five and eight.
|
||||||
|
|
||||||
|
l0_2 := l0 * 2
|
||||||
|
l1_2 := l1 * 2
|
||||||
|
|
||||||
|
l1_38 := l1 * 38
|
||||||
|
l2_38 := l2 * 38
|
||||||
|
l3_38 := l3 * 38
|
||||||
|
|
||||||
|
l3_19 := l3 * 19
|
||||||
|
l4_19 := l4 * 19
|
||||||
|
|
||||||
|
// r0 = l0×l0 + 19×(l1×l4 + l2×l3 + l3×l2 + l4×l1) = l0×l0 + 19×2×(l1×l4 + l2×l3)
|
||||||
|
r0 := mul64(l0, l0)
|
||||||
|
r0 = addMul64(r0, l1_38, l4)
|
||||||
|
r0 = addMul64(r0, l2_38, l3)
|
||||||
|
|
||||||
|
// r1 = l0×l1 + l1×l0 + 19×(l2×l4 + l3×l3 + l4×l2) = 2×l0×l1 + 19×2×l2×l4 + 19×l3×l3
|
||||||
|
r1 := mul64(l0_2, l1)
|
||||||
|
r1 = addMul64(r1, l2_38, l4)
|
||||||
|
r1 = addMul64(r1, l3_19, l3)
|
||||||
|
|
||||||
|
// r2 = l0×l2 + l1×l1 + l2×l0 + 19×(l3×l4 + l4×l3) = 2×l0×l2 + l1×l1 + 19×2×l3×l4
|
||||||
|
r2 := mul64(l0_2, l2)
|
||||||
|
r2 = addMul64(r2, l1, l1)
|
||||||
|
r2 = addMul64(r2, l3_38, l4)
|
||||||
|
|
||||||
|
// r3 = l0×l3 + l1×l2 + l2×l1 + l3×l0 + 19×l4×l4 = 2×l0×l3 + 2×l1×l2 + 19×l4×l4
|
||||||
|
r3 := mul64(l0_2, l3)
|
||||||
|
r3 = addMul64(r3, l1_2, l2)
|
||||||
|
r3 = addMul64(r3, l4_19, l4)
|
||||||
|
|
||||||
|
// r4 = l0×l4 + l1×l3 + l2×l2 + l3×l1 + l4×l0 = 2×l0×l4 + 2×l1×l3 + l2×l2
|
||||||
|
r4 := mul64(l0_2, l4)
|
||||||
|
r4 = addMul64(r4, l1_2, l3)
|
||||||
|
r4 = addMul64(r4, l2, l2)
|
||||||
|
|
||||||
|
c0 := shiftRightBy51(r0)
|
||||||
|
c1 := shiftRightBy51(r1)
|
||||||
|
c2 := shiftRightBy51(r2)
|
||||||
|
c3 := shiftRightBy51(r3)
|
||||||
|
c4 := shiftRightBy51(r4)
|
||||||
|
|
||||||
|
rr0 := r0.lo&maskLow51Bits + c4*19
|
||||||
|
rr1 := r1.lo&maskLow51Bits + c0
|
||||||
|
rr2 := r2.lo&maskLow51Bits + c1
|
||||||
|
rr3 := r3.lo&maskLow51Bits + c2
|
||||||
|
rr4 := r4.lo&maskLow51Bits + c3
|
||||||
|
|
||||||
|
*v = Element{rr0, rr1, rr2, rr3, rr4}
|
||||||
|
v.carryPropagate()
|
||||||
|
}
|
||||||
|
|
||||||
|
// carryPropagateGeneric brings the limbs below 52 bits by applying the reduction
|
||||||
|
// identity (a * 2²⁵⁵ + b = a * 19 + b) to the l4 carry.
|
||||||
|
func (v *Element) carryPropagateGeneric() *Element {
|
||||||
|
c0 := v.l0 >> 51
|
||||||
|
c1 := v.l1 >> 51
|
||||||
|
c2 := v.l2 >> 51
|
||||||
|
c3 := v.l3 >> 51
|
||||||
|
c4 := v.l4 >> 51
|
||||||
|
|
||||||
|
// c4 is at most 64 - 51 = 13 bits, so c4*19 is at most 18 bits, and
|
||||||
|
// the final l0 will be at most 52 bits. Similarly for the rest.
|
||||||
|
v.l0 = v.l0&maskLow51Bits + c4*19
|
||||||
|
v.l1 = v.l1&maskLow51Bits + c0
|
||||||
|
v.l2 = v.l2&maskLow51Bits + c1
|
||||||
|
v.l3 = v.l3&maskLow51Bits + c2
|
||||||
|
v.l4 = v.l4&maskLow51Bits + c3
|
||||||
|
|
||||||
|
return v
|
||||||
|
}
|
||||||
+343
@@ -0,0 +1,343 @@
|
|||||||
|
// Copyright (c) 2016 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package edwards25519
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
// A Scalar is an integer modulo
|
||||||
|
//
|
||||||
|
// l = 2^252 + 27742317777372353535851937790883648493
|
||||||
|
//
|
||||||
|
// which is the prime order of the edwards25519 group.
|
||||||
|
//
|
||||||
|
// This type works similarly to math/big.Int, and all arguments and
|
||||||
|
// receivers are allowed to alias.
|
||||||
|
//
|
||||||
|
// The zero value is a valid zero element.
|
||||||
|
type Scalar struct {
|
||||||
|
// s is the scalar in the Montgomery domain, in the format of the
|
||||||
|
// fiat-crypto implementation.
|
||||||
|
s fiatScalarMontgomeryDomainFieldElement
|
||||||
|
}
|
||||||
|
|
||||||
|
// The field implementation in scalar_fiat.go is generated by the fiat-crypto
|
||||||
|
// project (https://github.com/mit-plv/fiat-crypto) at version v0.0.9 (23d2dbc)
|
||||||
|
// from a formally verified model.
|
||||||
|
//
|
||||||
|
// fiat-crypto code comes under the following license.
|
||||||
|
//
|
||||||
|
// Copyright (c) 2015-2020 The fiat-crypto Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// Redistribution and use in source and binary forms, with or without
|
||||||
|
// modification, are permitted provided that the following conditions are
|
||||||
|
// met:
|
||||||
|
//
|
||||||
|
// 1. Redistributions of source code must retain the above copyright
|
||||||
|
// notice, this list of conditions and the following disclaimer.
|
||||||
|
//
|
||||||
|
// THIS SOFTWARE IS PROVIDED BY the fiat-crypto authors "AS IS"
|
||||||
|
// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO,
|
||||||
|
// THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
|
||||||
|
// PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL Berkeley Software Design,
|
||||||
|
// Inc. BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL,
|
||||||
|
// EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO,
|
||||||
|
// PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR
|
||||||
|
// PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF
|
||||||
|
// LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING
|
||||||
|
// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS
|
||||||
|
// SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||||
|
//
|
||||||
|
|
||||||
|
// NewScalar returns a new zero Scalar.
|
||||||
|
func NewScalar() *Scalar {
|
||||||
|
return &Scalar{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// MultiplyAdd sets s = x * y + z mod l, and returns s. It is equivalent to
|
||||||
|
// using Multiply and then Add.
|
||||||
|
func (s *Scalar) MultiplyAdd(x, y, z *Scalar) *Scalar {
|
||||||
|
// Make a copy of z in case it aliases s.
|
||||||
|
zCopy := new(Scalar).Set(z)
|
||||||
|
return s.Multiply(x, y).Add(s, zCopy)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add sets s = x + y mod l, and returns s.
|
||||||
|
func (s *Scalar) Add(x, y *Scalar) *Scalar {
|
||||||
|
// s = 1 * x + y mod l
|
||||||
|
fiatScalarAdd(&s.s, &x.s, &y.s)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Subtract sets s = x - y mod l, and returns s.
|
||||||
|
func (s *Scalar) Subtract(x, y *Scalar) *Scalar {
|
||||||
|
// s = -1 * y + x mod l
|
||||||
|
fiatScalarSub(&s.s, &x.s, &y.s)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Negate sets s = -x mod l, and returns s.
|
||||||
|
func (s *Scalar) Negate(x *Scalar) *Scalar {
|
||||||
|
// s = -1 * x + 0 mod l
|
||||||
|
fiatScalarOpp(&s.s, &x.s)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiply sets s = x * y mod l, and returns s.
|
||||||
|
func (s *Scalar) Multiply(x, y *Scalar) *Scalar {
|
||||||
|
// s = x * y + 0 mod l
|
||||||
|
fiatScalarMul(&s.s, &x.s, &y.s)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set sets s = x, and returns s.
|
||||||
|
func (s *Scalar) Set(x *Scalar) *Scalar {
|
||||||
|
*s = *x
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetUniformBytes sets s = x mod l, where x is a 64-byte little-endian integer.
|
||||||
|
// If x is not of the right length, SetUniformBytes returns nil and an error,
|
||||||
|
// and the receiver is unchanged.
|
||||||
|
//
|
||||||
|
// SetUniformBytes can be used to set s to a uniformly distributed value given
|
||||||
|
// 64 uniformly distributed random bytes.
|
||||||
|
func (s *Scalar) SetUniformBytes(x []byte) (*Scalar, error) {
|
||||||
|
if len(x) != 64 {
|
||||||
|
return nil, errors.New("edwards25519: invalid SetUniformBytes input length")
|
||||||
|
}
|
||||||
|
|
||||||
|
// We have a value x of 512 bits, but our fiatScalarFromBytes function
|
||||||
|
// expects an input lower than l, which is a little over 252 bits.
|
||||||
|
//
|
||||||
|
// Instead of writing a reduction function that operates on wider inputs, we
|
||||||
|
// can interpret x as the sum of three shorter values a, b, and c.
|
||||||
|
//
|
||||||
|
// x = a + b * 2^168 + c * 2^336 mod l
|
||||||
|
//
|
||||||
|
// We then precompute 2^168 and 2^336 modulo l, and perform the reduction
|
||||||
|
// with two multiplications and two additions.
|
||||||
|
|
||||||
|
s.setShortBytes(x[:21])
|
||||||
|
t := new(Scalar).setShortBytes(x[21:42])
|
||||||
|
s.Add(s, t.Multiply(t, scalarTwo168))
|
||||||
|
t.setShortBytes(x[42:])
|
||||||
|
s.Add(s, t.Multiply(t, scalarTwo336))
|
||||||
|
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// scalarTwo168 and scalarTwo336 are 2^168 and 2^336 modulo l, encoded as a
|
||||||
|
// fiatScalarMontgomeryDomainFieldElement, which is a little-endian 4-limb value
|
||||||
|
// in the 2^256 Montgomery domain.
|
||||||
|
var scalarTwo168 = &Scalar{s: [4]uint64{0x5b8ab432eac74798, 0x38afddd6de59d5d7,
|
||||||
|
0xa2c131b399411b7c, 0x6329a7ed9ce5a30}}
|
||||||
|
var scalarTwo336 = &Scalar{s: [4]uint64{0xbd3d108e2b35ecc5, 0x5c3a3718bdf9c90b,
|
||||||
|
0x63aa97a331b4f2ee, 0x3d217f5be65cb5c}}
|
||||||
|
|
||||||
|
// setShortBytes sets s = x mod l, where x is a little-endian integer shorter
|
||||||
|
// than 32 bytes.
|
||||||
|
func (s *Scalar) setShortBytes(x []byte) *Scalar {
|
||||||
|
if len(x) >= 32 {
|
||||||
|
panic("edwards25519: internal error: setShortBytes called with a long string")
|
||||||
|
}
|
||||||
|
var buf [32]byte
|
||||||
|
copy(buf[:], x)
|
||||||
|
fiatScalarFromBytes((*[4]uint64)(&s.s), &buf)
|
||||||
|
fiatScalarToMontgomery(&s.s, (*fiatScalarNonMontgomeryDomainFieldElement)(&s.s))
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCanonicalBytes sets s = x, where x is a 32-byte little-endian encoding of
|
||||||
|
// s, and returns s. If x is not a canonical encoding of s, SetCanonicalBytes
|
||||||
|
// returns nil and an error, and the receiver is unchanged.
|
||||||
|
func (s *Scalar) SetCanonicalBytes(x []byte) (*Scalar, error) {
|
||||||
|
if len(x) != 32 {
|
||||||
|
return nil, errors.New("invalid scalar length")
|
||||||
|
}
|
||||||
|
if !isReduced(x) {
|
||||||
|
return nil, errors.New("invalid scalar encoding")
|
||||||
|
}
|
||||||
|
|
||||||
|
fiatScalarFromBytes((*[4]uint64)(&s.s), (*[32]byte)(x))
|
||||||
|
fiatScalarToMontgomery(&s.s, (*fiatScalarNonMontgomeryDomainFieldElement)(&s.s))
|
||||||
|
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// scalarMinusOneBytes is l - 1 in little endian.
|
||||||
|
var scalarMinusOneBytes = [32]byte{236, 211, 245, 92, 26, 99, 18, 88, 214, 156, 247, 162, 222, 249, 222, 20, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 16}
|
||||||
|
|
||||||
|
// isReduced returns whether the given scalar in 32-byte little endian encoded
|
||||||
|
// form is reduced modulo l.
|
||||||
|
func isReduced(s []byte) bool {
|
||||||
|
if len(s) != 32 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := len(s) - 1; i >= 0; i-- {
|
||||||
|
switch {
|
||||||
|
case s[i] > scalarMinusOneBytes[i]:
|
||||||
|
return false
|
||||||
|
case s[i] < scalarMinusOneBytes[i]:
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetBytesWithClamping applies the buffer pruning described in RFC 8032,
|
||||||
|
// Section 5.1.5 (also known as clamping) and sets s to the result. The input
|
||||||
|
// must be 32 bytes, and it is not modified. If x is not of the right length,
|
||||||
|
// SetBytesWithClamping returns nil and an error, and the receiver is unchanged.
|
||||||
|
//
|
||||||
|
// Note that since Scalar values are always reduced modulo the prime order of
|
||||||
|
// the curve, the resulting value will not preserve any of the cofactor-clearing
|
||||||
|
// properties that clamping is meant to provide. It will however work as
|
||||||
|
// expected as long as it is applied to points on the prime order subgroup, like
|
||||||
|
// in Ed25519. In fact, it is lost to history why RFC 8032 adopted the
|
||||||
|
// irrelevant RFC 7748 clamping, but it is now required for compatibility.
|
||||||
|
func (s *Scalar) SetBytesWithClamping(x []byte) (*Scalar, error) {
|
||||||
|
// The description above omits the purpose of the high bits of the clamping
|
||||||
|
// for brevity, but those are also lost to reductions, and are also
|
||||||
|
// irrelevant to edwards25519 as they protect against a specific
|
||||||
|
// implementation bug that was once observed in a generic Montgomery ladder.
|
||||||
|
if len(x) != 32 {
|
||||||
|
return nil, errors.New("edwards25519: invalid SetBytesWithClamping input length")
|
||||||
|
}
|
||||||
|
|
||||||
|
// We need to use the wide reduction from SetUniformBytes, since clamping
|
||||||
|
// sets the 2^254 bit, making the value higher than the order.
|
||||||
|
var wideBytes [64]byte
|
||||||
|
copy(wideBytes[:], x[:])
|
||||||
|
wideBytes[0] &= 248
|
||||||
|
wideBytes[31] &= 63
|
||||||
|
wideBytes[31] |= 64
|
||||||
|
return s.SetUniformBytes(wideBytes[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Bytes returns the canonical 32-byte little-endian encoding of s.
|
||||||
|
func (s *Scalar) Bytes() []byte {
|
||||||
|
// This function is outlined to make the allocations inline in the caller
|
||||||
|
// rather than happen on the heap.
|
||||||
|
var encoded [32]byte
|
||||||
|
return s.bytes(&encoded)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Scalar) bytes(out *[32]byte) []byte {
|
||||||
|
var ss fiatScalarNonMontgomeryDomainFieldElement
|
||||||
|
fiatScalarFromMontgomery(&ss, &s.s)
|
||||||
|
fiatScalarToBytes(out, (*[4]uint64)(&ss))
|
||||||
|
return out[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Equal returns 1 if s and t are equal, and 0 otherwise.
|
||||||
|
func (s *Scalar) Equal(t *Scalar) int {
|
||||||
|
var diff fiatScalarMontgomeryDomainFieldElement
|
||||||
|
fiatScalarSub(&diff, &s.s, &t.s)
|
||||||
|
var nonzero uint64
|
||||||
|
fiatScalarNonzero(&nonzero, (*[4]uint64)(&diff))
|
||||||
|
nonzero |= nonzero >> 32
|
||||||
|
nonzero |= nonzero >> 16
|
||||||
|
nonzero |= nonzero >> 8
|
||||||
|
nonzero |= nonzero >> 4
|
||||||
|
nonzero |= nonzero >> 2
|
||||||
|
nonzero |= nonzero >> 1
|
||||||
|
return int(^nonzero) & 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// nonAdjacentForm computes a width-w non-adjacent form for this scalar.
|
||||||
|
//
|
||||||
|
// w must be between 2 and 8, or nonAdjacentForm will panic.
|
||||||
|
func (s *Scalar) nonAdjacentForm(w uint) [256]int8 {
|
||||||
|
// This implementation is adapted from the one
|
||||||
|
// in curve25519-dalek and is documented there:
|
||||||
|
// https://github.com/dalek-cryptography/curve25519-dalek/blob/f630041af28e9a405255f98a8a93adca18e4315b/src/scalar.rs#L800-L871
|
||||||
|
b := s.Bytes()
|
||||||
|
if b[31] > 127 {
|
||||||
|
panic("scalar has high bit set illegally")
|
||||||
|
}
|
||||||
|
if w < 2 {
|
||||||
|
panic("w must be at least 2 by the definition of NAF")
|
||||||
|
} else if w > 8 {
|
||||||
|
panic("NAF digits must fit in int8")
|
||||||
|
}
|
||||||
|
|
||||||
|
var naf [256]int8
|
||||||
|
var digits [5]uint64
|
||||||
|
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
digits[i] = binary.LittleEndian.Uint64(b[i*8:])
|
||||||
|
}
|
||||||
|
|
||||||
|
width := uint64(1 << w)
|
||||||
|
windowMask := uint64(width - 1)
|
||||||
|
|
||||||
|
pos := uint(0)
|
||||||
|
carry := uint64(0)
|
||||||
|
for pos < 256 {
|
||||||
|
indexU64 := pos / 64
|
||||||
|
indexBit := pos % 64
|
||||||
|
var bitBuf uint64
|
||||||
|
if indexBit < 64-w {
|
||||||
|
// This window's bits are contained in a single u64
|
||||||
|
bitBuf = digits[indexU64] >> indexBit
|
||||||
|
} else {
|
||||||
|
// Combine the current 64 bits with bits from the next 64
|
||||||
|
bitBuf = (digits[indexU64] >> indexBit) | (digits[1+indexU64] << (64 - indexBit))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add carry into the current window
|
||||||
|
window := carry + (bitBuf & windowMask)
|
||||||
|
|
||||||
|
if window&1 == 0 {
|
||||||
|
// If the window value is even, preserve the carry and continue.
|
||||||
|
// Why is the carry preserved?
|
||||||
|
// If carry == 0 and window & 1 == 0,
|
||||||
|
// then the next carry should be 0
|
||||||
|
// If carry == 1 and window & 1 == 0,
|
||||||
|
// then bit_buf & 1 == 1 so the next carry should be 1
|
||||||
|
pos += 1
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if window < width/2 {
|
||||||
|
carry = 0
|
||||||
|
naf[pos] = int8(window)
|
||||||
|
} else {
|
||||||
|
carry = 1
|
||||||
|
naf[pos] = int8(window) - int8(width)
|
||||||
|
}
|
||||||
|
|
||||||
|
pos += w
|
||||||
|
}
|
||||||
|
return naf
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Scalar) signedRadix16() [64]int8 {
|
||||||
|
b := s.Bytes()
|
||||||
|
if b[31] > 127 {
|
||||||
|
panic("scalar has high bit set illegally")
|
||||||
|
}
|
||||||
|
|
||||||
|
var digits [64]int8
|
||||||
|
|
||||||
|
// Compute unsigned radix-16 digits:
|
||||||
|
for i := 0; i < 32; i++ {
|
||||||
|
digits[2*i] = int8(b[i] & 15)
|
||||||
|
digits[2*i+1] = int8((b[i] >> 4) & 15)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Recenter coefficients:
|
||||||
|
for i := 0; i < 63; i++ {
|
||||||
|
carry := (digits[i] + 8) >> 4
|
||||||
|
digits[i] -= carry << 4
|
||||||
|
digits[i+1] += carry
|
||||||
|
}
|
||||||
|
|
||||||
|
return digits
|
||||||
|
}
|
||||||
+1147
File diff suppressed because it is too large
Load Diff
+214
@@ -0,0 +1,214 @@
|
|||||||
|
// Copyright (c) 2019 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package edwards25519
|
||||||
|
|
||||||
|
import "sync"
|
||||||
|
|
||||||
|
// basepointTable is a set of 32 affineLookupTables, where table i is generated
|
||||||
|
// from 256i * basepoint. It is precomputed the first time it's used.
|
||||||
|
func basepointTable() *[32]affineLookupTable {
|
||||||
|
basepointTablePrecomp.initOnce.Do(func() {
|
||||||
|
p := NewGeneratorPoint()
|
||||||
|
for i := 0; i < 32; i++ {
|
||||||
|
basepointTablePrecomp.table[i].FromP3(p)
|
||||||
|
for j := 0; j < 8; j++ {
|
||||||
|
p.Add(p, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return &basepointTablePrecomp.table
|
||||||
|
}
|
||||||
|
|
||||||
|
var basepointTablePrecomp struct {
|
||||||
|
table [32]affineLookupTable
|
||||||
|
initOnce sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
// ScalarBaseMult sets v = x * B, where B is the canonical generator, and
|
||||||
|
// returns v.
|
||||||
|
//
|
||||||
|
// The scalar multiplication is done in constant time.
|
||||||
|
func (v *Point) ScalarBaseMult(x *Scalar) *Point {
|
||||||
|
basepointTable := basepointTable()
|
||||||
|
|
||||||
|
// Write x = sum(x_i * 16^i) so x*B = sum( B*x_i*16^i )
|
||||||
|
// as described in the Ed25519 paper
|
||||||
|
//
|
||||||
|
// Group even and odd coefficients
|
||||||
|
// x*B = x_0*16^0*B + x_2*16^2*B + ... + x_62*16^62*B
|
||||||
|
// + x_1*16^1*B + x_3*16^3*B + ... + x_63*16^63*B
|
||||||
|
// x*B = x_0*16^0*B + x_2*16^2*B + ... + x_62*16^62*B
|
||||||
|
// + 16*( x_1*16^0*B + x_3*16^2*B + ... + x_63*16^62*B)
|
||||||
|
//
|
||||||
|
// We use a lookup table for each i to get x_i*16^(2*i)*B
|
||||||
|
// and do four doublings to multiply by 16.
|
||||||
|
digits := x.signedRadix16()
|
||||||
|
|
||||||
|
multiple := &affineCached{}
|
||||||
|
tmp1 := &projP1xP1{}
|
||||||
|
tmp2 := &projP2{}
|
||||||
|
|
||||||
|
// Accumulate the odd components first
|
||||||
|
v.Set(NewIdentityPoint())
|
||||||
|
for i := 1; i < 64; i += 2 {
|
||||||
|
basepointTable[i/2].SelectInto(multiple, digits[i])
|
||||||
|
tmp1.AddAffine(v, multiple)
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Multiply by 16
|
||||||
|
tmp2.FromP3(v) // tmp2 = v in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 2*v in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 2*v in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 4*v in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 4*v in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 8*v in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 8*v in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 16*v in P1xP1 coords
|
||||||
|
v.fromP1xP1(tmp1) // now v = 16*(odd components)
|
||||||
|
|
||||||
|
// Accumulate the even components
|
||||||
|
for i := 0; i < 64; i += 2 {
|
||||||
|
basepointTable[i/2].SelectInto(multiple, digits[i])
|
||||||
|
tmp1.AddAffine(v, multiple)
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
}
|
||||||
|
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// ScalarMult sets v = x * q, and returns v.
|
||||||
|
//
|
||||||
|
// The scalar multiplication is done in constant time.
|
||||||
|
func (v *Point) ScalarMult(x *Scalar, q *Point) *Point {
|
||||||
|
checkInitialized(q)
|
||||||
|
|
||||||
|
var table projLookupTable
|
||||||
|
table.FromP3(q)
|
||||||
|
|
||||||
|
// Write x = sum(x_i * 16^i)
|
||||||
|
// so x*Q = sum( Q*x_i*16^i )
|
||||||
|
// = Q*x_0 + 16*(Q*x_1 + 16*( ... + Q*x_63) ... )
|
||||||
|
// <------compute inside out---------
|
||||||
|
//
|
||||||
|
// We use the lookup table to get the x_i*Q values
|
||||||
|
// and do four doublings to compute 16*Q
|
||||||
|
digits := x.signedRadix16()
|
||||||
|
|
||||||
|
// Unwrap first loop iteration to save computing 16*identity
|
||||||
|
multiple := &projCached{}
|
||||||
|
tmp1 := &projP1xP1{}
|
||||||
|
tmp2 := &projP2{}
|
||||||
|
table.SelectInto(multiple, digits[63])
|
||||||
|
|
||||||
|
v.Set(NewIdentityPoint())
|
||||||
|
tmp1.Add(v, multiple) // tmp1 = x_63*Q in P1xP1 coords
|
||||||
|
for i := 62; i >= 0; i-- {
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = (prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 2*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 2*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 4*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 4*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 8*(prev) in P1xP1 coords
|
||||||
|
tmp2.FromP1xP1(tmp1) // tmp2 = 8*(prev) in P2 coords
|
||||||
|
tmp1.Double(tmp2) // tmp1 = 16*(prev) in P1xP1 coords
|
||||||
|
v.fromP1xP1(tmp1) // v = 16*(prev) in P3 coords
|
||||||
|
table.SelectInto(multiple, digits[i])
|
||||||
|
tmp1.Add(v, multiple) // tmp1 = x_i*Q + 16*(prev) in P1xP1 coords
|
||||||
|
}
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
return v
|
||||||
|
}
|
||||||
|
|
||||||
|
// basepointNafTable is the nafLookupTable8 for the basepoint.
|
||||||
|
// It is precomputed the first time it's used.
|
||||||
|
func basepointNafTable() *nafLookupTable8 {
|
||||||
|
basepointNafTablePrecomp.initOnce.Do(func() {
|
||||||
|
basepointNafTablePrecomp.table.FromP3(NewGeneratorPoint())
|
||||||
|
})
|
||||||
|
return &basepointNafTablePrecomp.table
|
||||||
|
}
|
||||||
|
|
||||||
|
var basepointNafTablePrecomp struct {
|
||||||
|
table nafLookupTable8
|
||||||
|
initOnce sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
// VarTimeDoubleScalarBaseMult sets v = a * A + b * B, where B is the canonical
|
||||||
|
// generator, and returns v.
|
||||||
|
//
|
||||||
|
// Execution time depends on the inputs.
|
||||||
|
func (v *Point) VarTimeDoubleScalarBaseMult(a *Scalar, A *Point, b *Scalar) *Point {
|
||||||
|
checkInitialized(A)
|
||||||
|
|
||||||
|
// Similarly to the single variable-base approach, we compute
|
||||||
|
// digits and use them with a lookup table. However, because
|
||||||
|
// we are allowed to do variable-time operations, we don't
|
||||||
|
// need constant-time lookups or constant-time digit
|
||||||
|
// computations.
|
||||||
|
//
|
||||||
|
// So we use a non-adjacent form of some width w instead of
|
||||||
|
// radix 16. This is like a binary representation (one digit
|
||||||
|
// for each binary place) but we allow the digits to grow in
|
||||||
|
// magnitude up to 2^{w-1} so that the nonzero digits are as
|
||||||
|
// sparse as possible. Intuitively, this "condenses" the
|
||||||
|
// "mass" of the scalar onto sparse coefficients (meaning
|
||||||
|
// fewer additions).
|
||||||
|
|
||||||
|
basepointNafTable := basepointNafTable()
|
||||||
|
var aTable nafLookupTable5
|
||||||
|
aTable.FromP3(A)
|
||||||
|
// Because the basepoint is fixed, we can use a wider NAF
|
||||||
|
// corresponding to a bigger table.
|
||||||
|
aNaf := a.nonAdjacentForm(5)
|
||||||
|
bNaf := b.nonAdjacentForm(8)
|
||||||
|
|
||||||
|
// Find the first nonzero coefficient.
|
||||||
|
i := 255
|
||||||
|
for j := i; j >= 0; j-- {
|
||||||
|
if aNaf[j] != 0 || bNaf[j] != 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
multA := &projCached{}
|
||||||
|
multB := &affineCached{}
|
||||||
|
tmp1 := &projP1xP1{}
|
||||||
|
tmp2 := &projP2{}
|
||||||
|
tmp2.Zero()
|
||||||
|
|
||||||
|
// Move from high to low bits, doubling the accumulator
|
||||||
|
// at each iteration and checking whether there is a nonzero
|
||||||
|
// coefficient to look up a multiple of.
|
||||||
|
for ; i >= 0; i-- {
|
||||||
|
tmp1.Double(tmp2)
|
||||||
|
|
||||||
|
// Only update v if we have a nonzero coeff to add in.
|
||||||
|
if aNaf[i] > 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
aTable.SelectInto(multA, aNaf[i])
|
||||||
|
tmp1.Add(v, multA)
|
||||||
|
} else if aNaf[i] < 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
aTable.SelectInto(multA, -aNaf[i])
|
||||||
|
tmp1.Sub(v, multA)
|
||||||
|
}
|
||||||
|
|
||||||
|
if bNaf[i] > 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
basepointNafTable.SelectInto(multB, bNaf[i])
|
||||||
|
tmp1.AddAffine(v, multB)
|
||||||
|
} else if bNaf[i] < 0 {
|
||||||
|
v.fromP1xP1(tmp1)
|
||||||
|
basepointNafTable.SelectInto(multB, -bNaf[i])
|
||||||
|
tmp1.SubAffine(v, multB)
|
||||||
|
}
|
||||||
|
|
||||||
|
tmp2.FromP1xP1(tmp1)
|
||||||
|
}
|
||||||
|
|
||||||
|
v.fromP2(tmp2)
|
||||||
|
return v
|
||||||
|
}
|
||||||
+129
@@ -0,0 +1,129 @@
|
|||||||
|
// Copyright (c) 2019 The Go Authors. All rights reserved.
|
||||||
|
// Use of this source code is governed by a BSD-style
|
||||||
|
// license that can be found in the LICENSE file.
|
||||||
|
|
||||||
|
package edwards25519
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/subtle"
|
||||||
|
)
|
||||||
|
|
||||||
|
// A dynamic lookup table for variable-base, constant-time scalar muls.
|
||||||
|
type projLookupTable struct {
|
||||||
|
points [8]projCached
|
||||||
|
}
|
||||||
|
|
||||||
|
// A precomputed lookup table for fixed-base, constant-time scalar muls.
|
||||||
|
type affineLookupTable struct {
|
||||||
|
points [8]affineCached
|
||||||
|
}
|
||||||
|
|
||||||
|
// A dynamic lookup table for variable-base, variable-time scalar muls.
|
||||||
|
type nafLookupTable5 struct {
|
||||||
|
points [8]projCached
|
||||||
|
}
|
||||||
|
|
||||||
|
// A precomputed lookup table for fixed-base, variable-time scalar muls.
|
||||||
|
type nafLookupTable8 struct {
|
||||||
|
points [64]affineCached
|
||||||
|
}
|
||||||
|
|
||||||
|
// Constructors.
|
||||||
|
|
||||||
|
// Builds a lookup table at runtime. Fast.
|
||||||
|
func (v *projLookupTable) FromP3(q *Point) {
|
||||||
|
// Goal: v.points[i] = (i+1)*Q, i.e., Q, 2Q, ..., 8Q
|
||||||
|
// This allows lookup of -8Q, ..., -Q, 0, Q, ..., 8Q
|
||||||
|
v.points[0].FromP3(q)
|
||||||
|
tmpP3 := Point{}
|
||||||
|
tmpP1xP1 := projP1xP1{}
|
||||||
|
for i := 0; i < 7; i++ {
|
||||||
|
// Compute (i+1)*Q as Q + i*Q and convert to a projCached
|
||||||
|
// This is needlessly complicated because the API has explicit
|
||||||
|
// receivers instead of creating stack objects and relying on RVO
|
||||||
|
v.points[i+1].FromP3(tmpP3.fromP1xP1(tmpP1xP1.Add(q, &v.points[i])))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// This is not optimised for speed; fixed-base tables should be precomputed.
|
||||||
|
func (v *affineLookupTable) FromP3(q *Point) {
|
||||||
|
// Goal: v.points[i] = (i+1)*Q, i.e., Q, 2Q, ..., 8Q
|
||||||
|
// This allows lookup of -8Q, ..., -Q, 0, Q, ..., 8Q
|
||||||
|
v.points[0].FromP3(q)
|
||||||
|
tmpP3 := Point{}
|
||||||
|
tmpP1xP1 := projP1xP1{}
|
||||||
|
for i := 0; i < 7; i++ {
|
||||||
|
// Compute (i+1)*Q as Q + i*Q and convert to affineCached
|
||||||
|
v.points[i+1].FromP3(tmpP3.fromP1xP1(tmpP1xP1.AddAffine(q, &v.points[i])))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Builds a lookup table at runtime. Fast.
|
||||||
|
func (v *nafLookupTable5) FromP3(q *Point) {
|
||||||
|
// Goal: v.points[i] = (2*i+1)*Q, i.e., Q, 3Q, 5Q, ..., 15Q
|
||||||
|
// This allows lookup of -15Q, ..., -3Q, -Q, 0, Q, 3Q, ..., 15Q
|
||||||
|
v.points[0].FromP3(q)
|
||||||
|
q2 := Point{}
|
||||||
|
q2.Add(q, q)
|
||||||
|
tmpP3 := Point{}
|
||||||
|
tmpP1xP1 := projP1xP1{}
|
||||||
|
for i := 0; i < 7; i++ {
|
||||||
|
v.points[i+1].FromP3(tmpP3.fromP1xP1(tmpP1xP1.Add(&q2, &v.points[i])))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// This is not optimised for speed; fixed-base tables should be precomputed.
|
||||||
|
func (v *nafLookupTable8) FromP3(q *Point) {
|
||||||
|
v.points[0].FromP3(q)
|
||||||
|
q2 := Point{}
|
||||||
|
q2.Add(q, q)
|
||||||
|
tmpP3 := Point{}
|
||||||
|
tmpP1xP1 := projP1xP1{}
|
||||||
|
for i := 0; i < 63; i++ {
|
||||||
|
v.points[i+1].FromP3(tmpP3.fromP1xP1(tmpP1xP1.AddAffine(&q2, &v.points[i])))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Selectors.
|
||||||
|
|
||||||
|
// Set dest to x*Q, where -8 <= x <= 8, in constant time.
|
||||||
|
func (v *projLookupTable) SelectInto(dest *projCached, x int8) {
|
||||||
|
// Compute xabs = |x|
|
||||||
|
xmask := x >> 7
|
||||||
|
xabs := uint8((x + xmask) ^ xmask)
|
||||||
|
|
||||||
|
dest.Zero()
|
||||||
|
for j := 1; j <= 8; j++ {
|
||||||
|
// Set dest = j*Q if |x| = j
|
||||||
|
cond := subtle.ConstantTimeByteEq(xabs, uint8(j))
|
||||||
|
dest.Select(&v.points[j-1], dest, cond)
|
||||||
|
}
|
||||||
|
// Now dest = |x|*Q, conditionally negate to get x*Q
|
||||||
|
dest.CondNeg(int(xmask & 1))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set dest to x*Q, where -8 <= x <= 8, in constant time.
|
||||||
|
func (v *affineLookupTable) SelectInto(dest *affineCached, x int8) {
|
||||||
|
// Compute xabs = |x|
|
||||||
|
xmask := x >> 7
|
||||||
|
xabs := uint8((x + xmask) ^ xmask)
|
||||||
|
|
||||||
|
dest.Zero()
|
||||||
|
for j := 1; j <= 8; j++ {
|
||||||
|
// Set dest = j*Q if |x| = j
|
||||||
|
cond := subtle.ConstantTimeByteEq(xabs, uint8(j))
|
||||||
|
dest.Select(&v.points[j-1], dest, cond)
|
||||||
|
}
|
||||||
|
// Now dest = |x|*Q, conditionally negate to get x*Q
|
||||||
|
dest.CondNeg(int(xmask & 1))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Given odd x with 0 < x < 2^4, return x*Q (in variable time).
|
||||||
|
func (v *nafLookupTable5) SelectInto(dest *projCached, x int8) {
|
||||||
|
*dest = v.points[x/2]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Given odd x with 0 < x < 2^7, return x*Q (in variable time).
|
||||||
|
func (v *nafLookupTable8) SelectInto(dest *affineCached, x int8) {
|
||||||
|
*dest = v.points[x/2]
|
||||||
|
}
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
.DS_Store
|
||||||
|
.DS_Store?
|
||||||
|
._*
|
||||||
|
.Spotlight-V100
|
||||||
|
.Trashes
|
||||||
|
Icon?
|
||||||
|
ehthumbs.db
|
||||||
|
Thumbs.db
|
||||||
|
.idea
|
||||||
+155
@@ -0,0 +1,155 @@
|
|||||||
|
# This is the official list of Go-MySQL-Driver authors for copyright purposes.
|
||||||
|
|
||||||
|
# If you are submitting a patch, please add your name or the name of the
|
||||||
|
# organization which holds the copyright to this list in alphabetical order.
|
||||||
|
|
||||||
|
# Names should be added to this file as
|
||||||
|
# Name <email address>
|
||||||
|
# The email address is not required for organizations.
|
||||||
|
# Please keep the list sorted.
|
||||||
|
|
||||||
|
|
||||||
|
# Individual Persons
|
||||||
|
|
||||||
|
Aaron Hopkins <go-sql-driver at die.net>
|
||||||
|
Achille Roussel <achille.roussel at gmail.com>
|
||||||
|
Aidan <aidan.liu at pingcap.com>
|
||||||
|
Alex Snast <alexsn at fb.com>
|
||||||
|
Alexey Palazhchenko <alexey.palazhchenko at gmail.com>
|
||||||
|
Andrew Reid <andrew.reid at tixtrack.com>
|
||||||
|
Animesh Ray <mail.rayanimesh at gmail.com>
|
||||||
|
Arne Hormann <arnehormann at gmail.com>
|
||||||
|
Ariel Mashraki <ariel at mashraki.co.il>
|
||||||
|
Artur Melanchyk <artur.melanchyk@gmail.com>
|
||||||
|
Asta Xie <xiemengjun at gmail.com>
|
||||||
|
B Lamarche <blam413 at gmail.com>
|
||||||
|
Bes Dollma <bdollma@thousandeyes.com>
|
||||||
|
Bogdan Constantinescu <bog.con.bc at gmail.com>
|
||||||
|
Brad Higgins <brad at defined.net>
|
||||||
|
Brian Hendriks <brian at dolthub.com>
|
||||||
|
Bulat Gaifullin <gaifullinbf at gmail.com>
|
||||||
|
Caine Jette <jette at alum.mit.edu>
|
||||||
|
Carlos Nieto <jose.carlos at menteslibres.net>
|
||||||
|
Chris Kirkland <chriskirkland at github.com>
|
||||||
|
Chris Moos <chris at tech9computers.com>
|
||||||
|
Craig Wilson <craiggwilson at gmail.com>
|
||||||
|
Daemonxiao <735462752 at qq.com>
|
||||||
|
Daniel Montoya <dsmontoyam at gmail.com>
|
||||||
|
Daniel Nichter <nil at codenode.com>
|
||||||
|
Daniël van Eeden <git at myname.nl>
|
||||||
|
Dave Protasowski <dprotaso at gmail.com>
|
||||||
|
Diego Dupin <diego.dupin at gmail.com>
|
||||||
|
Dirkjan Bussink <d.bussink at gmail.com>
|
||||||
|
DisposaBoy <disposaboy at dby.me>
|
||||||
|
Egor Smolyakov <egorsmkv at gmail.com>
|
||||||
|
Erwan Martin <hello at erwan.io>
|
||||||
|
Evan Elias <evan at skeema.net>
|
||||||
|
Evan Shaw <evan at vendhq.com>
|
||||||
|
Frederick Mayle <frederickmayle at gmail.com>
|
||||||
|
Gustavo Kristic <gkristic at gmail.com>
|
||||||
|
Gusted <postmaster at gusted.xyz>
|
||||||
|
Hajime Nakagami <nakagami at gmail.com>
|
||||||
|
Hanno Braun <mail at hannobraun.com>
|
||||||
|
Henri Yandell <flamefew at gmail.com>
|
||||||
|
Hirotaka Yamamoto <ymmt2005 at gmail.com>
|
||||||
|
Huyiguang <hyg at webterren.com>
|
||||||
|
ICHINOSE Shogo <shogo82148 at gmail.com>
|
||||||
|
Ilia Cimpoes <ichimpoesh at gmail.com>
|
||||||
|
INADA Naoki <songofacandy at gmail.com>
|
||||||
|
Jacek Szwec <szwec.jacek at gmail.com>
|
||||||
|
Jakub Adamus <kratky at zobak.cz>
|
||||||
|
James Harr <james.harr at gmail.com>
|
||||||
|
Janek Vedock <janekvedock at comcast.net>
|
||||||
|
Jason Ng <oblitorum at gmail.com>
|
||||||
|
Jean-Yves Pellé <jy at pelle.link>
|
||||||
|
Jeff Hodges <jeff at somethingsimilar.com>
|
||||||
|
Jeffrey Charles <jeffreycharles at gmail.com>
|
||||||
|
Jennifer Purevsuren <jennifer at dolthub.com>
|
||||||
|
Jerome Meyer <jxmeyer at gmail.com>
|
||||||
|
Jiajia Zhong <zhong2plus at gmail.com>
|
||||||
|
Jian Zhen <zhenjl at gmail.com>
|
||||||
|
Joe Mann <contact at joemann.co.uk>
|
||||||
|
Joshua Prunier <joshua.prunier at gmail.com>
|
||||||
|
Julien Lefevre <julien.lefevr at gmail.com>
|
||||||
|
Julien Schmidt <go-sql-driver at julienschmidt.com>
|
||||||
|
Justin Li <jli at j-li.net>
|
||||||
|
Justin Nuß <nuss.justin at gmail.com>
|
||||||
|
Kamil Dziedzic <kamil at klecza.pl>
|
||||||
|
Kei Kamikawa <x00.x7f.x86 at gmail.com>
|
||||||
|
Kevin Malachowski <kevin at chowski.com>
|
||||||
|
Kieron Woodhouse <kieron.woodhouse at infosum.com>
|
||||||
|
Lance Tian <lance6716 at gmail.com>
|
||||||
|
Lennart Rudolph <lrudolph at hmc.edu>
|
||||||
|
Leonardo YongUk Kim <dalinaum at gmail.com>
|
||||||
|
Linh Tran Tuan <linhduonggnu at gmail.com>
|
||||||
|
Lion Yang <lion at aosc.xyz>
|
||||||
|
Luca Looz <luca.looz92 at gmail.com>
|
||||||
|
Lucas Liu <extrafliu at gmail.com>
|
||||||
|
Lunny Xiao <xiaolunwen at gmail.com>
|
||||||
|
Luke Scott <luke at webconnex.com>
|
||||||
|
Maciej Zimnoch <maciej.zimnoch at codilime.com>
|
||||||
|
Michael Woolnough <michael.woolnough at gmail.com>
|
||||||
|
Nao Yokotsuka <yokotukanao at gmail.com>
|
||||||
|
Nathanial Murphy <nathanial.murphy at gmail.com>
|
||||||
|
Nicola Peduzzi <thenikso at gmail.com>
|
||||||
|
Oliver Bone <owbone at github.com>
|
||||||
|
Olivier Mengué <dolmen at cpan.org>
|
||||||
|
oscarzhao <oscarzhaosl at gmail.com>
|
||||||
|
Paul Bonser <misterpib at gmail.com>
|
||||||
|
Paulius Lozys <pauliuslozys at gmail.com>
|
||||||
|
Peter Schultz <peter.schultz at classmarkets.com>
|
||||||
|
Phil Porada <philporada at gmail.com>
|
||||||
|
Minh Quang <minhquang4334 at gmail.com>
|
||||||
|
Rebecca Chin <rchin at pivotal.io>
|
||||||
|
Reed Allman <rdallman10 at gmail.com>
|
||||||
|
Richard Wilkes <wilkes at me.com>
|
||||||
|
Robert Russell <robert at rrbrussell.com>
|
||||||
|
Runrioter Wung <runrioter at gmail.com>
|
||||||
|
Samantha Frank <hello at entropy.cat>
|
||||||
|
Santhosh Kumar Tekuri <santhosh.tekuri at gmail.com>
|
||||||
|
Sho Iizuka <sho.i518 at gmail.com>
|
||||||
|
Sho Ikeda <suicaicoca at gmail.com>
|
||||||
|
Shuode Li <elemount at qq.com>
|
||||||
|
Simon J Mudd <sjmudd at pobox.com>
|
||||||
|
Soroush Pour <me at soroushjp.com>
|
||||||
|
Stan Putrya <root.vagner at gmail.com>
|
||||||
|
Stanley Gunawan <gunawan.stanley at gmail.com>
|
||||||
|
Steven Hartland <steven.hartland at multiplay.co.uk>
|
||||||
|
Tan Jinhua <312841925 at qq.com>
|
||||||
|
Tetsuro Aoki <t.aoki1130 at gmail.com>
|
||||||
|
Thomas Wodarek <wodarekwebpage at gmail.com>
|
||||||
|
Tim Ruffles <timruffles at gmail.com>
|
||||||
|
Tom Jenkinson <tom at tjenkinson.me>
|
||||||
|
Vladimir Kovpak <cn007b at gmail.com>
|
||||||
|
Vladyslav Zhelezniak <zhvladi at gmail.com>
|
||||||
|
Xiangyu Hu <xiangyu.hu at outlook.com>
|
||||||
|
Xiaobing Jiang <s7v7nislands at gmail.com>
|
||||||
|
Xiuming Chen <cc at cxm.cc>
|
||||||
|
Xuehong Chan <chanxuehong at gmail.com>
|
||||||
|
Zhang Xiang <angwerzx at 126.com>
|
||||||
|
Zhenye Xie <xiezhenye at gmail.com>
|
||||||
|
Zhixin Wen <john.wenzhixin at gmail.com>
|
||||||
|
Ziheng Lyu <zihenglv at gmail.com>
|
||||||
|
|
||||||
|
# Organizations
|
||||||
|
|
||||||
|
Barracuda Networks, Inc.
|
||||||
|
Counting Ltd.
|
||||||
|
Defined Networking Inc.
|
||||||
|
DigitalOcean Inc.
|
||||||
|
Dolthub Inc.
|
||||||
|
dyves labs AG
|
||||||
|
Facebook Inc.
|
||||||
|
GitHub Inc.
|
||||||
|
Google Inc.
|
||||||
|
InfoSum Ltd.
|
||||||
|
Keybase Inc.
|
||||||
|
Microsoft Corp.
|
||||||
|
Multiplay Ltd.
|
||||||
|
Percona LLC
|
||||||
|
PingCAP Inc.
|
||||||
|
Pivotal Inc.
|
||||||
|
Shattered Silicon Ltd.
|
||||||
|
Stripe Inc.
|
||||||
|
ThousandEyes
|
||||||
|
Zendesk Inc.
|
||||||
+360
@@ -0,0 +1,360 @@
|
|||||||
|
# Changelog
|
||||||
|
|
||||||
|
## v1.9.3 (2025-06-13)
|
||||||
|
|
||||||
|
* `tx.Commit()` and `tx.Rollback()` returned `ErrInvalidConn` always.
|
||||||
|
Now they return cached real error if present. (#1690)
|
||||||
|
|
||||||
|
* Optimize reading small resultsets to fix performance regression
|
||||||
|
introduced by compression protocol support. (#1707)
|
||||||
|
|
||||||
|
* Fix `db.Ping()` on compressed connection. (#1723)
|
||||||
|
|
||||||
|
|
||||||
|
## v1.9.2 (2025-04-07)
|
||||||
|
|
||||||
|
v1.9.2 is a re-release of v1.9.1 due to a release process issue; no changes were made to the content.
|
||||||
|
|
||||||
|
|
||||||
|
## v1.9.1 (2025-03-21)
|
||||||
|
|
||||||
|
### Major Changes
|
||||||
|
|
||||||
|
* Add Charset() option. (#1679)
|
||||||
|
|
||||||
|
### Bugfixes
|
||||||
|
|
||||||
|
* go.mod: fix go version format (#1682)
|
||||||
|
* Fix FormatDSN missing ConnectionAttributes (#1619)
|
||||||
|
|
||||||
|
## v1.9.0 (2025-02-18)
|
||||||
|
|
||||||
|
### Major Changes
|
||||||
|
|
||||||
|
- Implement zlib compression. (#1487)
|
||||||
|
- Supported Go version is updated to Go 1.21+. (#1639)
|
||||||
|
- Add support for VECTOR type introduced in MySQL 9.0. (#1609)
|
||||||
|
- Config object can have custom dial function. (#1527)
|
||||||
|
|
||||||
|
### Bugfixes
|
||||||
|
|
||||||
|
- Fix auth errors when username/password are too long. (#1625)
|
||||||
|
- Check if MySQL supports CLIENT_CONNECT_ATTRS before sending client attributes. (#1640)
|
||||||
|
- Fix auth switch request handling. (#1666)
|
||||||
|
|
||||||
|
### Other changes
|
||||||
|
|
||||||
|
- Add "filename:line" prefix to log in go-mysql. Custom loggers now show it. (#1589)
|
||||||
|
- Improve error handling. It reduces the "busy buffer" errors. (#1595, #1601, #1641)
|
||||||
|
- Use `strconv.Atoi` to parse max_allowed_packet. (#1661)
|
||||||
|
- `rejectReadOnly` option now handles ER_READ_ONLY_MODE (1290) error too. (#1660)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.8.1 (2024-03-26)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- fix race condition when context is canceled in [#1562](https://github.com/go-sql-driver/mysql/pull/1562) and [#1570](https://github.com/go-sql-driver/mysql/pull/1570)
|
||||||
|
|
||||||
|
## Version 1.8.0 (2024-03-09)
|
||||||
|
|
||||||
|
Major Changes:
|
||||||
|
|
||||||
|
- Use `SET NAMES charset COLLATE collation`. by @methane in [#1437](https://github.com/go-sql-driver/mysql/pull/1437)
|
||||||
|
- Older go-mysql-driver used `collation_id` in the handshake packet. But it caused collation mismatch in some situation.
|
||||||
|
- If you don't specify charset nor collation, go-mysql-driver sends `SET NAMES utf8mb4` for new connection. This uses server's default collation for utf8mb4.
|
||||||
|
- If you specify charset, go-mysql-driver sends `SET NAMES <charset>`. This uses the server's default collation for `<charset>`.
|
||||||
|
- If you specify collation and/or charset, go-mysql-driver sends `SET NAMES charset COLLATE collation`.
|
||||||
|
- PathEscape dbname in DSN. by @methane in [#1432](https://github.com/go-sql-driver/mysql/pull/1432)
|
||||||
|
- This is backward incompatible in rare case. Check your DSN.
|
||||||
|
- Drop Go 1.13-17 support by @methane in [#1420](https://github.com/go-sql-driver/mysql/pull/1420)
|
||||||
|
- Use Go 1.18+
|
||||||
|
- Parse numbers on text protocol too by @methane in [#1452](https://github.com/go-sql-driver/mysql/pull/1452)
|
||||||
|
- When text protocol is used, go-mysql-driver passed bare `[]byte` to database/sql for avoid unnecessary allocation and conversion.
|
||||||
|
- If user specified `*any` to `Scan()`, database/sql passed the `[]byte` into the target variable.
|
||||||
|
- This confused users because most user doesn't know when text/binary protocol used.
|
||||||
|
- go-mysql-driver 1.8 converts integer/float values into int64/double even in text protocol. This doesn't increase allocation compared to `[]byte` and conversion cost is negatable.
|
||||||
|
- New options start using the Functional Option Pattern to avoid increasing technical debt in the Config object. Future version may introduce Functional Option for existing options, but not for now.
|
||||||
|
- Make TimeTruncate functional option by @methane in [1552](https://github.com/go-sql-driver/mysql/pull/1552)
|
||||||
|
- Add BeforeConnect callback to configuration object by @ItalyPaleAle in [#1469](https://github.com/go-sql-driver/mysql/pull/1469)
|
||||||
|
|
||||||
|
|
||||||
|
Other changes:
|
||||||
|
|
||||||
|
- Adding DeregisterDialContext to prevent memory leaks with dialers we don't need anymore by @jypelle in https://github.com/go-sql-driver/mysql/pull/1422
|
||||||
|
- Make logger configurable per connection by @frozenbonito in https://github.com/go-sql-driver/mysql/pull/1408
|
||||||
|
- Fix ColumnType.DatabaseTypeName for mediumint unsigned by @evanelias in https://github.com/go-sql-driver/mysql/pull/1428
|
||||||
|
- Add connection attributes by @Daemonxiao in https://github.com/go-sql-driver/mysql/pull/1389
|
||||||
|
- Stop `ColumnTypeScanType()` from returning `sql.RawBytes` by @methane in https://github.com/go-sql-driver/mysql/pull/1424
|
||||||
|
- Exec() now provides access to status of multiple statements. by @mherr-google in https://github.com/go-sql-driver/mysql/pull/1309
|
||||||
|
- Allow to change (or disable) the default driver name for registration by @dolmen in https://github.com/go-sql-driver/mysql/pull/1499
|
||||||
|
- Add default connection attribute '_server_host' by @oblitorum in https://github.com/go-sql-driver/mysql/pull/1506
|
||||||
|
- QueryUnescape DSN ConnectionAttribute value by @zhangyangyu in https://github.com/go-sql-driver/mysql/pull/1470
|
||||||
|
- Add client_ed25519 authentication by @Gusted in https://github.com/go-sql-driver/mysql/pull/1518
|
||||||
|
|
||||||
|
## Version 1.7.1 (2023-04-25)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- bump actions/checkout@v3 and actions/setup-go@v3 (#1375)
|
||||||
|
- Add go1.20 and mariadb10.11 to the testing matrix (#1403)
|
||||||
|
- Increase default maxAllowedPacket size. (#1411)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Use SET syntax as specified in the MySQL documentation (#1402)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.7 (2022-11-29)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Drop support of Go 1.12 (#1211)
|
||||||
|
- Refactoring `(*textRows).readRow` in a more clear way (#1230)
|
||||||
|
- util: Reduce boundary check in escape functions. (#1316)
|
||||||
|
- enhancement for mysqlConn handleAuthResult (#1250)
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- support Is comparison on MySQLError (#1210)
|
||||||
|
- return unsigned in database type name when necessary (#1238)
|
||||||
|
- Add API to express like a --ssl-mode=PREFERRED MySQL client (#1370)
|
||||||
|
- Add SQLState to MySQLError (#1321)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Fix parsing 0 year. (#1257)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.6 (2021-04-01)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Migrate the CI service from travis-ci to GitHub Actions (#1176, #1183, #1190)
|
||||||
|
- `NullTime` is deprecated (#960, #1144)
|
||||||
|
- Reduce allocations when building SET command (#1111)
|
||||||
|
- Performance improvement for time formatting (#1118)
|
||||||
|
- Performance improvement for time parsing (#1098, #1113)
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- Implement `driver.Validator` interface (#1106, #1174)
|
||||||
|
- Support returning `uint64` from `Valuer` in `ConvertValue` (#1143)
|
||||||
|
- Add `json.RawMessage` for converter and prepared statement (#1059)
|
||||||
|
- Interpolate `json.RawMessage` as `string` (#1058)
|
||||||
|
- Implements `CheckNamedValue` (#1090)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Stop rounding times (#1121, #1172)
|
||||||
|
- Put zero filler into the SSL handshake packet (#1066)
|
||||||
|
- Fix checking cancelled connections back into the connection pool (#1095)
|
||||||
|
- Fix remove last 0 byte for mysql_old_password when password is empty (#1133)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.5 (2020-01-07)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Dropped support Go 1.9 and lower (#823, #829, #886, #1016, #1017)
|
||||||
|
- Improve buffer handling (#890)
|
||||||
|
- Document potentially insecure TLS configs (#901)
|
||||||
|
- Use a double-buffering scheme to prevent data races (#943)
|
||||||
|
- Pass uint64 values without converting them to string (#838, #955)
|
||||||
|
- Update collations and make utf8mb4 default (#877, #1054)
|
||||||
|
- Make NullTime compatible with sql.NullTime in Go 1.13+ (#995)
|
||||||
|
- Removed CloudSQL support (#993, #1007)
|
||||||
|
- Add Go Module support (#1003)
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- Implement support of optional TLS (#900)
|
||||||
|
- Check connection liveness (#934, #964, #997, #1048, #1051, #1052)
|
||||||
|
- Implement Connector Interface (#941, #958, #1020, #1035)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Mark connections as bad on error during ping (#875)
|
||||||
|
- Mark connections as bad on error during dial (#867)
|
||||||
|
- Fix connection leak caused by rapid context cancellation (#1024)
|
||||||
|
- Mark connections as bad on error during Conn.Prepare (#1030)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.4.1 (2018-11-14)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Fix TIME format for binary columns (#818)
|
||||||
|
- Fix handling of empty auth plugin names (#835)
|
||||||
|
- Fix caching_sha2_password with empty password (#826)
|
||||||
|
- Fix canceled context broke mysqlConn (#862)
|
||||||
|
- Fix OldAuthSwitchRequest support (#870)
|
||||||
|
- Fix Auth Response packet for cleartext password (#887)
|
||||||
|
|
||||||
|
## Version 1.4 (2018-06-03)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Documentation fixes (#530, #535, #567)
|
||||||
|
- Refactoring (#575, #579, #580, #581, #603, #615, #704)
|
||||||
|
- Cache column names (#444)
|
||||||
|
- Sort the DSN parameters in DSNs generated from a config (#637)
|
||||||
|
- Allow native password authentication by default (#644)
|
||||||
|
- Use the default port if it is missing in the DSN (#668)
|
||||||
|
- Removed the `strict` mode (#676)
|
||||||
|
- Do not query `max_allowed_packet` by default (#680)
|
||||||
|
- Dropped support Go 1.6 and lower (#696)
|
||||||
|
- Updated `ConvertValue()` to match the database/sql/driver implementation (#760)
|
||||||
|
- Document the usage of `0000-00-00T00:00:00` as the time.Time zero value (#783)
|
||||||
|
- Improved the compatibility of the authentication system (#807)
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- Multi-Results support (#537)
|
||||||
|
- `rejectReadOnly` DSN option (#604)
|
||||||
|
- `context.Context` support (#608, #612, #627, #761)
|
||||||
|
- Transaction isolation level support (#619, #744)
|
||||||
|
- Read-Only transactions support (#618, #634)
|
||||||
|
- `NewConfig` function which initializes a config with default values (#679)
|
||||||
|
- Implemented the `ColumnType` interfaces (#667, #724)
|
||||||
|
- Support for custom string types in `ConvertValue` (#623)
|
||||||
|
- Implemented `NamedValueChecker`, improving support for uint64 with high bit set (#690, #709, #710)
|
||||||
|
- `caching_sha2_password` authentication plugin support (#794, #800, #801, #802)
|
||||||
|
- Implemented `driver.SessionResetter` (#779)
|
||||||
|
- `sha256_password` authentication plugin support (#808)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Use the DSN hostname as TLS default ServerName if `tls=true` (#564, #718)
|
||||||
|
- Fixed LOAD LOCAL DATA INFILE for empty files (#590)
|
||||||
|
- Removed columns definition cache since it sometimes cached invalid data (#592)
|
||||||
|
- Don't mutate registered TLS configs (#600)
|
||||||
|
- Make RegisterTLSConfig concurrency-safe (#613)
|
||||||
|
- Handle missing auth data in the handshake packet correctly (#646)
|
||||||
|
- Do not retry queries when data was written to avoid data corruption (#302, #736)
|
||||||
|
- Cache the connection pointer for error handling before invalidating it (#678)
|
||||||
|
- Fixed imports for appengine/cloudsql (#700)
|
||||||
|
- Fix sending STMT_LONG_DATA for 0 byte data (#734)
|
||||||
|
- Set correct capacity for []bytes read from length-encoded strings (#766)
|
||||||
|
- Make RegisterDial concurrency-safe (#773)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.3 (2016-12-01)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Go 1.1 is no longer supported
|
||||||
|
- Use decimals fields in MySQL to format time types (#249)
|
||||||
|
- Buffer optimizations (#269)
|
||||||
|
- TLS ServerName defaults to the host (#283)
|
||||||
|
- Refactoring (#400, #410, #437)
|
||||||
|
- Adjusted documentation for second generation CloudSQL (#485)
|
||||||
|
- Documented DSN system var quoting rules (#502)
|
||||||
|
- Made statement.Close() calls idempotent to avoid errors in Go 1.6+ (#512)
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- Enable microsecond resolution on TIME, DATETIME and TIMESTAMP (#249)
|
||||||
|
- Support for returning table alias on Columns() (#289, #359, #382)
|
||||||
|
- Placeholder interpolation, can be activated with the DSN parameter `interpolateParams=true` (#309, #318, #490)
|
||||||
|
- Support for uint64 parameters with high bit set (#332, #345)
|
||||||
|
- Cleartext authentication plugin support (#327)
|
||||||
|
- Exported ParseDSN function and the Config struct (#403, #419, #429)
|
||||||
|
- Read / Write timeouts (#401)
|
||||||
|
- Support for JSON field type (#414)
|
||||||
|
- Support for multi-statements and multi-results (#411, #431)
|
||||||
|
- DSN parameter to set the driver-side max_allowed_packet value manually (#489)
|
||||||
|
- Native password authentication plugin support (#494, #524)
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Fixed handling of queries without columns and rows (#255)
|
||||||
|
- Fixed a panic when SetKeepAlive() failed (#298)
|
||||||
|
- Handle ERR packets while reading rows (#321)
|
||||||
|
- Fixed reading NULL length-encoded integers in MySQL 5.6+ (#349)
|
||||||
|
- Fixed absolute paths support in LOAD LOCAL DATA INFILE (#356)
|
||||||
|
- Actually zero out bytes in handshake response (#378)
|
||||||
|
- Fixed race condition in registering LOAD DATA INFILE handler (#383)
|
||||||
|
- Fixed tests with MySQL 5.7.9+ (#380)
|
||||||
|
- QueryUnescape TLS config names (#397)
|
||||||
|
- Fixed "broken pipe" error by writing to closed socket (#390)
|
||||||
|
- Fixed LOAD LOCAL DATA INFILE buffering (#424)
|
||||||
|
- Fixed parsing of floats into float64 when placeholders are used (#434)
|
||||||
|
- Fixed DSN tests with Go 1.7+ (#459)
|
||||||
|
- Handle ERR packets while waiting for EOF (#473)
|
||||||
|
- Invalidate connection on error while discarding additional results (#513)
|
||||||
|
- Allow terminating packets of length 0 (#516)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.2 (2014-06-03)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- We switched back to a "rolling release". `go get` installs the current master branch again
|
||||||
|
- Version v1 of the driver will not be maintained anymore. Go 1.0 is no longer supported by this driver
|
||||||
|
- Exported errors to allow easy checking from application code
|
||||||
|
- Enabled TCP Keepalives on TCP connections
|
||||||
|
- Optimized INFILE handling (better buffer size calculation, lazy init, ...)
|
||||||
|
- The DSN parser also checks for a missing separating slash
|
||||||
|
- Faster binary date / datetime to string formatting
|
||||||
|
- Also exported the MySQLWarning type
|
||||||
|
- mysqlConn.Close returns the first error encountered instead of ignoring all errors
|
||||||
|
- writePacket() automatically writes the packet size to the header
|
||||||
|
- readPacket() uses an iterative approach instead of the recursive approach to merge split packets
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- `RegisterDial` allows the usage of a custom dial function to establish the network connection
|
||||||
|
- Setting the connection collation is possible with the `collation` DSN parameter. This parameter should be preferred over the `charset` parameter
|
||||||
|
- Logging of critical errors is configurable with `SetLogger`
|
||||||
|
- Google CloudSQL support
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Allow more than 32 parameters in prepared statements
|
||||||
|
- Various old_password fixes
|
||||||
|
- Fixed TestConcurrent test to pass Go's race detection
|
||||||
|
- Fixed appendLengthEncodedInteger for large numbers
|
||||||
|
- Renamed readLengthEnodedString to readLengthEncodedString and skipLengthEnodedString to skipLengthEncodedString (fixed typo)
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.1 (2013-11-02)
|
||||||
|
|
||||||
|
Changes:
|
||||||
|
|
||||||
|
- Go-MySQL-Driver now requires Go 1.1
|
||||||
|
- Connections now use the collation `utf8_general_ci` by default. Adding `&charset=UTF8` to the DSN should not be necessary anymore
|
||||||
|
- Made closing rows and connections error tolerant. This allows for example deferring rows.Close() without checking for errors
|
||||||
|
- `[]byte(nil)` is now treated as a NULL value. Before, it was treated like an empty string / `[]byte("")`
|
||||||
|
- DSN parameter values must now be url.QueryEscape'ed. This allows text values to contain special characters, such as '&'.
|
||||||
|
- Use the IO buffer also for writing. This results in zero allocations (by the driver) for most queries
|
||||||
|
- Optimized the buffer for reading
|
||||||
|
- stmt.Query now caches column metadata
|
||||||
|
- New Logo
|
||||||
|
- Changed the copyright header to include all contributors
|
||||||
|
- Improved the LOAD INFILE documentation
|
||||||
|
- The driver struct is now exported to make the driver directly accessible
|
||||||
|
- Refactored the driver tests
|
||||||
|
- Added more benchmarks and moved all to a separate file
|
||||||
|
- Other small refactoring
|
||||||
|
|
||||||
|
New Features:
|
||||||
|
|
||||||
|
- Added *old_passwords* support: Required in some cases, but must be enabled by adding `allowOldPasswords=true` to the DSN since it is insecure
|
||||||
|
- Added a `clientFoundRows` parameter: Return the number of matching rows instead of the number of rows changed on UPDATEs
|
||||||
|
- Added TLS/SSL support: Use a TLS/SSL encrypted connection to the server. Custom TLS configs can be registered and used
|
||||||
|
|
||||||
|
Bugfixes:
|
||||||
|
|
||||||
|
- Fixed MySQL 4.1 support: MySQL 4.1 sends packets with lengths which differ from the specification
|
||||||
|
- Convert to DB timezone when inserting `time.Time`
|
||||||
|
- Split packets (more than 16MB) are now merged correctly
|
||||||
|
- Fixed false positive `io.EOF` errors when the data was fully read
|
||||||
|
- Avoid panics on reuse of closed connections
|
||||||
|
- Fixed empty string producing false nil values
|
||||||
|
- Fixed sign byte for positive TIME fields
|
||||||
|
|
||||||
|
|
||||||
|
## Version 1.0 (2013-05-14)
|
||||||
|
|
||||||
|
Initial Release
|
||||||
+373
@@ -0,0 +1,373 @@
|
|||||||
|
Mozilla Public License Version 2.0
|
||||||
|
==================================
|
||||||
|
|
||||||
|
1. Definitions
|
||||||
|
--------------
|
||||||
|
|
||||||
|
1.1. "Contributor"
|
||||||
|
means each individual or legal entity that creates, contributes to
|
||||||
|
the creation of, or owns Covered Software.
|
||||||
|
|
||||||
|
1.2. "Contributor Version"
|
||||||
|
means the combination of the Contributions of others (if any) used
|
||||||
|
by a Contributor and that particular Contributor's Contribution.
|
||||||
|
|
||||||
|
1.3. "Contribution"
|
||||||
|
means Covered Software of a particular Contributor.
|
||||||
|
|
||||||
|
1.4. "Covered Software"
|
||||||
|
means Source Code Form to which the initial Contributor has attached
|
||||||
|
the notice in Exhibit A, the Executable Form of such Source Code
|
||||||
|
Form, and Modifications of such Source Code Form, in each case
|
||||||
|
including portions thereof.
|
||||||
|
|
||||||
|
1.5. "Incompatible With Secondary Licenses"
|
||||||
|
means
|
||||||
|
|
||||||
|
(a) that the initial Contributor has attached the notice described
|
||||||
|
in Exhibit B to the Covered Software; or
|
||||||
|
|
||||||
|
(b) that the Covered Software was made available under the terms of
|
||||||
|
version 1.1 or earlier of the License, but not also under the
|
||||||
|
terms of a Secondary License.
|
||||||
|
|
||||||
|
1.6. "Executable Form"
|
||||||
|
means any form of the work other than Source Code Form.
|
||||||
|
|
||||||
|
1.7. "Larger Work"
|
||||||
|
means a work that combines Covered Software with other material, in
|
||||||
|
a separate file or files, that is not Covered Software.
|
||||||
|
|
||||||
|
1.8. "License"
|
||||||
|
means this document.
|
||||||
|
|
||||||
|
1.9. "Licensable"
|
||||||
|
means having the right to grant, to the maximum extent possible,
|
||||||
|
whether at the time of the initial grant or subsequently, any and
|
||||||
|
all of the rights conveyed by this License.
|
||||||
|
|
||||||
|
1.10. "Modifications"
|
||||||
|
means any of the following:
|
||||||
|
|
||||||
|
(a) any file in Source Code Form that results from an addition to,
|
||||||
|
deletion from, or modification of the contents of Covered
|
||||||
|
Software; or
|
||||||
|
|
||||||
|
(b) any new file in Source Code Form that contains any Covered
|
||||||
|
Software.
|
||||||
|
|
||||||
|
1.11. "Patent Claims" of a Contributor
|
||||||
|
means any patent claim(s), including without limitation, method,
|
||||||
|
process, and apparatus claims, in any patent Licensable by such
|
||||||
|
Contributor that would be infringed, but for the grant of the
|
||||||
|
License, by the making, using, selling, offering for sale, having
|
||||||
|
made, import, or transfer of either its Contributions or its
|
||||||
|
Contributor Version.
|
||||||
|
|
||||||
|
1.12. "Secondary License"
|
||||||
|
means either the GNU General Public License, Version 2.0, the GNU
|
||||||
|
Lesser General Public License, Version 2.1, the GNU Affero General
|
||||||
|
Public License, Version 3.0, or any later versions of those
|
||||||
|
licenses.
|
||||||
|
|
||||||
|
1.13. "Source Code Form"
|
||||||
|
means the form of the work preferred for making modifications.
|
||||||
|
|
||||||
|
1.14. "You" (or "Your")
|
||||||
|
means an individual or a legal entity exercising rights under this
|
||||||
|
License. For legal entities, "You" includes any entity that
|
||||||
|
controls, is controlled by, or is under common control with You. For
|
||||||
|
purposes of this definition, "control" means (a) the power, direct
|
||||||
|
or indirect, to cause the direction or management of such entity,
|
||||||
|
whether by contract or otherwise, or (b) ownership of more than
|
||||||
|
fifty percent (50%) of the outstanding shares or beneficial
|
||||||
|
ownership of such entity.
|
||||||
|
|
||||||
|
2. License Grants and Conditions
|
||||||
|
--------------------------------
|
||||||
|
|
||||||
|
2.1. Grants
|
||||||
|
|
||||||
|
Each Contributor hereby grants You a world-wide, royalty-free,
|
||||||
|
non-exclusive license:
|
||||||
|
|
||||||
|
(a) under intellectual property rights (other than patent or trademark)
|
||||||
|
Licensable by such Contributor to use, reproduce, make available,
|
||||||
|
modify, display, perform, distribute, and otherwise exploit its
|
||||||
|
Contributions, either on an unmodified basis, with Modifications, or
|
||||||
|
as part of a Larger Work; and
|
||||||
|
|
||||||
|
(b) under Patent Claims of such Contributor to make, use, sell, offer
|
||||||
|
for sale, have made, import, and otherwise transfer either its
|
||||||
|
Contributions or its Contributor Version.
|
||||||
|
|
||||||
|
2.2. Effective Date
|
||||||
|
|
||||||
|
The licenses granted in Section 2.1 with respect to any Contribution
|
||||||
|
become effective for each Contribution on the date the Contributor first
|
||||||
|
distributes such Contribution.
|
||||||
|
|
||||||
|
2.3. Limitations on Grant Scope
|
||||||
|
|
||||||
|
The licenses granted in this Section 2 are the only rights granted under
|
||||||
|
this License. No additional rights or licenses will be implied from the
|
||||||
|
distribution or licensing of Covered Software under this License.
|
||||||
|
Notwithstanding Section 2.1(b) above, no patent license is granted by a
|
||||||
|
Contributor:
|
||||||
|
|
||||||
|
(a) for any code that a Contributor has removed from Covered Software;
|
||||||
|
or
|
||||||
|
|
||||||
|
(b) for infringements caused by: (i) Your and any other third party's
|
||||||
|
modifications of Covered Software, or (ii) the combination of its
|
||||||
|
Contributions with other software (except as part of its Contributor
|
||||||
|
Version); or
|
||||||
|
|
||||||
|
(c) under Patent Claims infringed by Covered Software in the absence of
|
||||||
|
its Contributions.
|
||||||
|
|
||||||
|
This License does not grant any rights in the trademarks, service marks,
|
||||||
|
or logos of any Contributor (except as may be necessary to comply with
|
||||||
|
the notice requirements in Section 3.4).
|
||||||
|
|
||||||
|
2.4. Subsequent Licenses
|
||||||
|
|
||||||
|
No Contributor makes additional grants as a result of Your choice to
|
||||||
|
distribute the Covered Software under a subsequent version of this
|
||||||
|
License (see Section 10.2) or under the terms of a Secondary License (if
|
||||||
|
permitted under the terms of Section 3.3).
|
||||||
|
|
||||||
|
2.5. Representation
|
||||||
|
|
||||||
|
Each Contributor represents that the Contributor believes its
|
||||||
|
Contributions are its original creation(s) or it has sufficient rights
|
||||||
|
to grant the rights to its Contributions conveyed by this License.
|
||||||
|
|
||||||
|
2.6. Fair Use
|
||||||
|
|
||||||
|
This License is not intended to limit any rights You have under
|
||||||
|
applicable copyright doctrines of fair use, fair dealing, or other
|
||||||
|
equivalents.
|
||||||
|
|
||||||
|
2.7. Conditions
|
||||||
|
|
||||||
|
Sections 3.1, 3.2, 3.3, and 3.4 are conditions of the licenses granted
|
||||||
|
in Section 2.1.
|
||||||
|
|
||||||
|
3. Responsibilities
|
||||||
|
-------------------
|
||||||
|
|
||||||
|
3.1. Distribution of Source Form
|
||||||
|
|
||||||
|
All distribution of Covered Software in Source Code Form, including any
|
||||||
|
Modifications that You create or to which You contribute, must be under
|
||||||
|
the terms of this License. You must inform recipients that the Source
|
||||||
|
Code Form of the Covered Software is governed by the terms of this
|
||||||
|
License, and how they can obtain a copy of this License. You may not
|
||||||
|
attempt to alter or restrict the recipients' rights in the Source Code
|
||||||
|
Form.
|
||||||
|
|
||||||
|
3.2. Distribution of Executable Form
|
||||||
|
|
||||||
|
If You distribute Covered Software in Executable Form then:
|
||||||
|
|
||||||
|
(a) such Covered Software must also be made available in Source Code
|
||||||
|
Form, as described in Section 3.1, and You must inform recipients of
|
||||||
|
the Executable Form how they can obtain a copy of such Source Code
|
||||||
|
Form by reasonable means in a timely manner, at a charge no more
|
||||||
|
than the cost of distribution to the recipient; and
|
||||||
|
|
||||||
|
(b) You may distribute such Executable Form under the terms of this
|
||||||
|
License, or sublicense it under different terms, provided that the
|
||||||
|
license for the Executable Form does not attempt to limit or alter
|
||||||
|
the recipients' rights in the Source Code Form under this License.
|
||||||
|
|
||||||
|
3.3. Distribution of a Larger Work
|
||||||
|
|
||||||
|
You may create and distribute a Larger Work under terms of Your choice,
|
||||||
|
provided that You also comply with the requirements of this License for
|
||||||
|
the Covered Software. If the Larger Work is a combination of Covered
|
||||||
|
Software with a work governed by one or more Secondary Licenses, and the
|
||||||
|
Covered Software is not Incompatible With Secondary Licenses, this
|
||||||
|
License permits You to additionally distribute such Covered Software
|
||||||
|
under the terms of such Secondary License(s), so that the recipient of
|
||||||
|
the Larger Work may, at their option, further distribute the Covered
|
||||||
|
Software under the terms of either this License or such Secondary
|
||||||
|
License(s).
|
||||||
|
|
||||||
|
3.4. Notices
|
||||||
|
|
||||||
|
You may not remove or alter the substance of any license notices
|
||||||
|
(including copyright notices, patent notices, disclaimers of warranty,
|
||||||
|
or limitations of liability) contained within the Source Code Form of
|
||||||
|
the Covered Software, except that You may alter any license notices to
|
||||||
|
the extent required to remedy known factual inaccuracies.
|
||||||
|
|
||||||
|
3.5. Application of Additional Terms
|
||||||
|
|
||||||
|
You may choose to offer, and to charge a fee for, warranty, support,
|
||||||
|
indemnity or liability obligations to one or more recipients of Covered
|
||||||
|
Software. However, You may do so only on Your own behalf, and not on
|
||||||
|
behalf of any Contributor. You must make it absolutely clear that any
|
||||||
|
such warranty, support, indemnity, or liability obligation is offered by
|
||||||
|
You alone, and You hereby agree to indemnify every Contributor for any
|
||||||
|
liability incurred by such Contributor as a result of warranty, support,
|
||||||
|
indemnity or liability terms You offer. You may include additional
|
||||||
|
disclaimers of warranty and limitations of liability specific to any
|
||||||
|
jurisdiction.
|
||||||
|
|
||||||
|
4. Inability to Comply Due to Statute or Regulation
|
||||||
|
---------------------------------------------------
|
||||||
|
|
||||||
|
If it is impossible for You to comply with any of the terms of this
|
||||||
|
License with respect to some or all of the Covered Software due to
|
||||||
|
statute, judicial order, or regulation then You must: (a) comply with
|
||||||
|
the terms of this License to the maximum extent possible; and (b)
|
||||||
|
describe the limitations and the code they affect. Such description must
|
||||||
|
be placed in a text file included with all distributions of the Covered
|
||||||
|
Software under this License. Except to the extent prohibited by statute
|
||||||
|
or regulation, such description must be sufficiently detailed for a
|
||||||
|
recipient of ordinary skill to be able to understand it.
|
||||||
|
|
||||||
|
5. Termination
|
||||||
|
--------------
|
||||||
|
|
||||||
|
5.1. The rights granted under this License will terminate automatically
|
||||||
|
if You fail to comply with any of its terms. However, if You become
|
||||||
|
compliant, then the rights granted under this License from a particular
|
||||||
|
Contributor are reinstated (a) provisionally, unless and until such
|
||||||
|
Contributor explicitly and finally terminates Your grants, and (b) on an
|
||||||
|
ongoing basis, if such Contributor fails to notify You of the
|
||||||
|
non-compliance by some reasonable means prior to 60 days after You have
|
||||||
|
come back into compliance. Moreover, Your grants from a particular
|
||||||
|
Contributor are reinstated on an ongoing basis if such Contributor
|
||||||
|
notifies You of the non-compliance by some reasonable means, this is the
|
||||||
|
first time You have received notice of non-compliance with this License
|
||||||
|
from such Contributor, and You become compliant prior to 30 days after
|
||||||
|
Your receipt of the notice.
|
||||||
|
|
||||||
|
5.2. If You initiate litigation against any entity by asserting a patent
|
||||||
|
infringement claim (excluding declaratory judgment actions,
|
||||||
|
counter-claims, and cross-claims) alleging that a Contributor Version
|
||||||
|
directly or indirectly infringes any patent, then the rights granted to
|
||||||
|
You by any and all Contributors for the Covered Software under Section
|
||||||
|
2.1 of this License shall terminate.
|
||||||
|
|
||||||
|
5.3. In the event of termination under Sections 5.1 or 5.2 above, all
|
||||||
|
end user license agreements (excluding distributors and resellers) which
|
||||||
|
have been validly granted by You or Your distributors under this License
|
||||||
|
prior to termination shall survive termination.
|
||||||
|
|
||||||
|
************************************************************************
|
||||||
|
* *
|
||||||
|
* 6. Disclaimer of Warranty *
|
||||||
|
* ------------------------- *
|
||||||
|
* *
|
||||||
|
* Covered Software is provided under this License on an "as is" *
|
||||||
|
* basis, without warranty of any kind, either expressed, implied, or *
|
||||||
|
* statutory, including, without limitation, warranties that the *
|
||||||
|
* Covered Software is free of defects, merchantable, fit for a *
|
||||||
|
* particular purpose or non-infringing. The entire risk as to the *
|
||||||
|
* quality and performance of the Covered Software is with You. *
|
||||||
|
* Should any Covered Software prove defective in any respect, You *
|
||||||
|
* (not any Contributor) assume the cost of any necessary servicing, *
|
||||||
|
* repair, or correction. This disclaimer of warranty constitutes an *
|
||||||
|
* essential part of this License. No use of any Covered Software is *
|
||||||
|
* authorized under this License except under this disclaimer. *
|
||||||
|
* *
|
||||||
|
************************************************************************
|
||||||
|
|
||||||
|
************************************************************************
|
||||||
|
* *
|
||||||
|
* 7. Limitation of Liability *
|
||||||
|
* -------------------------- *
|
||||||
|
* *
|
||||||
|
* Under no circumstances and under no legal theory, whether tort *
|
||||||
|
* (including negligence), contract, or otherwise, shall any *
|
||||||
|
* Contributor, or anyone who distributes Covered Software as *
|
||||||
|
* permitted above, be liable to You for any direct, indirect, *
|
||||||
|
* special, incidental, or consequential damages of any character *
|
||||||
|
* including, without limitation, damages for lost profits, loss of *
|
||||||
|
* goodwill, work stoppage, computer failure or malfunction, or any *
|
||||||
|
* and all other commercial damages or losses, even if such party *
|
||||||
|
* shall have been informed of the possibility of such damages. This *
|
||||||
|
* limitation of liability shall not apply to liability for death or *
|
||||||
|
* personal injury resulting from such party's negligence to the *
|
||||||
|
* extent applicable law prohibits such limitation. Some *
|
||||||
|
* jurisdictions do not allow the exclusion or limitation of *
|
||||||
|
* incidental or consequential damages, so this exclusion and *
|
||||||
|
* limitation may not apply to You. *
|
||||||
|
* *
|
||||||
|
************************************************************************
|
||||||
|
|
||||||
|
8. Litigation
|
||||||
|
-------------
|
||||||
|
|
||||||
|
Any litigation relating to this License may be brought only in the
|
||||||
|
courts of a jurisdiction where the defendant maintains its principal
|
||||||
|
place of business and such litigation shall be governed by laws of that
|
||||||
|
jurisdiction, without reference to its conflict-of-law provisions.
|
||||||
|
Nothing in this Section shall prevent a party's ability to bring
|
||||||
|
cross-claims or counter-claims.
|
||||||
|
|
||||||
|
9. Miscellaneous
|
||||||
|
----------------
|
||||||
|
|
||||||
|
This License represents the complete agreement concerning the subject
|
||||||
|
matter hereof. If any provision of this License is held to be
|
||||||
|
unenforceable, such provision shall be reformed only to the extent
|
||||||
|
necessary to make it enforceable. Any law or regulation which provides
|
||||||
|
that the language of a contract shall be construed against the drafter
|
||||||
|
shall not be used to construe this License against a Contributor.
|
||||||
|
|
||||||
|
10. Versions of the License
|
||||||
|
---------------------------
|
||||||
|
|
||||||
|
10.1. New Versions
|
||||||
|
|
||||||
|
Mozilla Foundation is the license steward. Except as provided in Section
|
||||||
|
10.3, no one other than the license steward has the right to modify or
|
||||||
|
publish new versions of this License. Each version will be given a
|
||||||
|
distinguishing version number.
|
||||||
|
|
||||||
|
10.2. Effect of New Versions
|
||||||
|
|
||||||
|
You may distribute the Covered Software under the terms of the version
|
||||||
|
of the License under which You originally received the Covered Software,
|
||||||
|
or under the terms of any subsequent version published by the license
|
||||||
|
steward.
|
||||||
|
|
||||||
|
10.3. Modified Versions
|
||||||
|
|
||||||
|
If you create software not governed by this License, and you want to
|
||||||
|
create a new license for such software, you may create and use a
|
||||||
|
modified version of this License if you rename the license and remove
|
||||||
|
any references to the name of the license steward (except to note that
|
||||||
|
such modified license differs from this License).
|
||||||
|
|
||||||
|
10.4. Distributing Source Code Form that is Incompatible With Secondary
|
||||||
|
Licenses
|
||||||
|
|
||||||
|
If You choose to distribute Source Code Form that is Incompatible With
|
||||||
|
Secondary Licenses under the terms of this version of the License, the
|
||||||
|
notice described in Exhibit B of this License must be attached.
|
||||||
|
|
||||||
|
Exhibit A - Source Code Form License Notice
|
||||||
|
-------------------------------------------
|
||||||
|
|
||||||
|
This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
License, v. 2.0. If a copy of the MPL was not distributed with this
|
||||||
|
file, You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
If it is not possible or desirable to put the notice in a particular
|
||||||
|
file, then You may include the notice in a location (such as a LICENSE
|
||||||
|
file in a relevant directory) where a recipient would be likely to look
|
||||||
|
for such a notice.
|
||||||
|
|
||||||
|
You may add additional accurate notices of copyright ownership.
|
||||||
|
|
||||||
|
Exhibit B - "Incompatible With Secondary Licenses" Notice
|
||||||
|
---------------------------------------------------------
|
||||||
|
|
||||||
|
This Source Code Form is "Incompatible With Secondary Licenses", as
|
||||||
|
defined by the Mozilla Public License, v. 2.0.
|
||||||
+595
@@ -0,0 +1,595 @@
|
|||||||
|
# Go-MySQL-Driver
|
||||||
|
|
||||||
|
A MySQL-Driver for Go's [database/sql](https://golang.org/pkg/database/sql/) package
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
---------------------------------------
|
||||||
|
* [Features](#features)
|
||||||
|
* [Requirements](#requirements)
|
||||||
|
* [Installation](#installation)
|
||||||
|
* [Usage](#usage)
|
||||||
|
* [DSN (Data Source Name)](#dsn-data-source-name)
|
||||||
|
* [Password](#password)
|
||||||
|
* [Protocol](#protocol)
|
||||||
|
* [Address](#address)
|
||||||
|
* [Parameters](#parameters)
|
||||||
|
* [Examples](#examples)
|
||||||
|
* [Connection pool and timeouts](#connection-pool-and-timeouts)
|
||||||
|
* [context.Context Support](#contextcontext-support)
|
||||||
|
* [ColumnType Support](#columntype-support)
|
||||||
|
* [LOAD DATA LOCAL INFILE support](#load-data-local-infile-support)
|
||||||
|
* [time.Time support](#timetime-support)
|
||||||
|
* [Unicode support](#unicode-support)
|
||||||
|
* [Testing / Development](#testing--development)
|
||||||
|
* [License](#license)
|
||||||
|
|
||||||
|
---------------------------------------
|
||||||
|
|
||||||
|
## Features
|
||||||
|
* Lightweight and [fast](https://github.com/go-sql-driver/sql-benchmark "golang MySQL-Driver performance")
|
||||||
|
* Native Go implementation. No C-bindings, just pure Go
|
||||||
|
* Connections over TCP/IPv4, TCP/IPv6, Unix domain sockets or [custom protocols](https://godoc.org/github.com/go-sql-driver/mysql#DialFunc)
|
||||||
|
* Automatic handling of broken connections
|
||||||
|
* Automatic Connection Pooling *(by database/sql package)*
|
||||||
|
* Supports queries larger than 16MB
|
||||||
|
* Full [`sql.RawBytes`](https://golang.org/pkg/database/sql/#RawBytes) support.
|
||||||
|
* Intelligent `LONG DATA` handling in prepared statements
|
||||||
|
* Secure `LOAD DATA LOCAL INFILE` support with file allowlisting and `io.Reader` support
|
||||||
|
* Optional `time.Time` parsing
|
||||||
|
* Optional placeholder interpolation
|
||||||
|
* Supports zlib compression.
|
||||||
|
|
||||||
|
## Requirements
|
||||||
|
|
||||||
|
* Go 1.21 or higher. We aim to support the 3 latest versions of Go.
|
||||||
|
* MySQL (5.7+) and MariaDB (10.5+) are supported.
|
||||||
|
* [TiDB](https://github.com/pingcap/tidb) is supported by PingCAP.
|
||||||
|
* Do not ask questions about TiDB in our issue tracker or forum.
|
||||||
|
* [Document](https://docs.pingcap.com/tidb/v6.1/dev-guide-sample-application-golang)
|
||||||
|
* [Forum](https://ask.pingcap.com/)
|
||||||
|
* go-mysql would work with Percona Server, Google CloudSQL or Sphinx (2.2.3+).
|
||||||
|
* Maintainers won't support them. Do not expect issues are investigated and resolved by maintainers.
|
||||||
|
* Investigate issues yourself and please send a pull request to fix it.
|
||||||
|
|
||||||
|
---------------------------------------
|
||||||
|
|
||||||
|
## Installation
|
||||||
|
Simple install the package to your [$GOPATH](https://github.com/golang/go/wiki/GOPATH "GOPATH") with the [go tool](https://golang.org/cmd/go/ "go command") from shell:
|
||||||
|
```bash
|
||||||
|
go get -u github.com/go-sql-driver/mysql
|
||||||
|
```
|
||||||
|
Make sure [Git is installed](https://git-scm.com/downloads) on your machine and in your system's `PATH`.
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
_Go MySQL Driver_ is an implementation of Go's `database/sql/driver` interface. You only need to import the driver and can use the full [`database/sql`](https://golang.org/pkg/database/sql/) API then.
|
||||||
|
|
||||||
|
Use `mysql` as `driverName` and a valid [DSN](#dsn-data-source-name) as `dataSourceName`:
|
||||||
|
|
||||||
|
```go
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
_ "github.com/go-sql-driver/mysql"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ...
|
||||||
|
|
||||||
|
db, err := sql.Open("mysql", "user:password@/dbname")
|
||||||
|
if err != nil {
|
||||||
|
panic(err)
|
||||||
|
}
|
||||||
|
// See "Important settings" section.
|
||||||
|
db.SetConnMaxLifetime(time.Minute * 3)
|
||||||
|
db.SetMaxOpenConns(10)
|
||||||
|
db.SetMaxIdleConns(10)
|
||||||
|
```
|
||||||
|
|
||||||
|
[Examples are available in our Wiki](https://github.com/go-sql-driver/mysql/wiki/Examples "Go-MySQL-Driver Examples").
|
||||||
|
|
||||||
|
### Important settings
|
||||||
|
|
||||||
|
`db.SetConnMaxLifetime()` is required to ensure connections are closed by the driver safely before connection is closed by MySQL server, OS, or other middlewares. Since some middlewares close idle connections by 5 minutes, we recommend timeout shorter than 5 minutes. This setting helps load balancing and changing system variables too.
|
||||||
|
|
||||||
|
`db.SetMaxOpenConns()` is highly recommended to limit the number of connection used by the application. There is no recommended limit number because it depends on application and MySQL server.
|
||||||
|
|
||||||
|
`db.SetMaxIdleConns()` is recommended to be set same to `db.SetMaxOpenConns()`. When it is smaller than `SetMaxOpenConns()`, connections can be opened and closed much more frequently than you expect. Idle connections can be closed by the `db.SetConnMaxLifetime()`. If you want to close idle connections more rapidly, you can use `db.SetConnMaxIdleTime()` since Go 1.15.
|
||||||
|
|
||||||
|
|
||||||
|
### DSN (Data Source Name)
|
||||||
|
|
||||||
|
The Data Source Name has a common format, like e.g. [PEAR DB](http://pear.php.net/manual/en/package.database.db.intro-dsn.php) uses it, but without type-prefix (optional parts marked by squared brackets):
|
||||||
|
```
|
||||||
|
[username[:password]@][protocol[(address)]]/dbname[?param1=value1&...¶mN=valueN]
|
||||||
|
```
|
||||||
|
|
||||||
|
A DSN in its fullest form:
|
||||||
|
```
|
||||||
|
username:password@protocol(address)/dbname?param=value
|
||||||
|
```
|
||||||
|
|
||||||
|
Except for the databasename, all values are optional. So the minimal DSN is:
|
||||||
|
```
|
||||||
|
/dbname
|
||||||
|
```
|
||||||
|
|
||||||
|
If you do not want to preselect a database, leave `dbname` empty:
|
||||||
|
```
|
||||||
|
/
|
||||||
|
```
|
||||||
|
This has the same effect as an empty DSN string:
|
||||||
|
```
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
`dbname` is escaped by [PathEscape()](https://pkg.go.dev/net/url#PathEscape) since v1.8.0. If your database name is `dbname/withslash`, it becomes:
|
||||||
|
|
||||||
|
```
|
||||||
|
/dbname%2Fwithslash
|
||||||
|
```
|
||||||
|
|
||||||
|
Alternatively, [Config.FormatDSN](https://godoc.org/github.com/go-sql-driver/mysql#Config.FormatDSN) can be used to create a DSN string by filling a struct.
|
||||||
|
|
||||||
|
#### Password
|
||||||
|
Passwords can consist of any character. Escaping is **not** necessary.
|
||||||
|
|
||||||
|
#### Protocol
|
||||||
|
See [net.Dial](https://golang.org/pkg/net/#Dial) for more information which networks are available.
|
||||||
|
In general you should use a Unix domain socket if available and TCP otherwise for best performance.
|
||||||
|
|
||||||
|
#### Address
|
||||||
|
For TCP and UDP networks, addresses have the form `host[:port]`.
|
||||||
|
If `port` is omitted, the default port will be used.
|
||||||
|
If `host` is a literal IPv6 address, it must be enclosed in square brackets.
|
||||||
|
The functions [net.JoinHostPort](https://golang.org/pkg/net/#JoinHostPort) and [net.SplitHostPort](https://golang.org/pkg/net/#SplitHostPort) manipulate addresses in this form.
|
||||||
|
|
||||||
|
For Unix domain sockets the address is the absolute path to the MySQL-Server-socket, e.g. `/var/run/mysqld/mysqld.sock` or `/tmp/mysql.sock`.
|
||||||
|
|
||||||
|
#### Parameters
|
||||||
|
*Parameters are case-sensitive!*
|
||||||
|
|
||||||
|
Notice that any of `true`, `TRUE`, `True` or `1` is accepted to stand for a true boolean value. Not surprisingly, false can be specified as any of: `false`, `FALSE`, `False` or `0`.
|
||||||
|
|
||||||
|
##### `allowAllFiles`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`allowAllFiles=true` disables the file allowlist for `LOAD DATA LOCAL INFILE` and allows *all* files.
|
||||||
|
[*Might be insecure!*](https://dev.mysql.com/doc/refman/8.0/en/load-data.html#load-data-local)
|
||||||
|
|
||||||
|
##### `allowCleartextPasswords`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`allowCleartextPasswords=true` allows using the [cleartext client side plugin](https://dev.mysql.com/doc/en/cleartext-pluggable-authentication.html) if required by an account, such as one defined with the [PAM authentication plugin](http://dev.mysql.com/doc/en/pam-authentication-plugin.html). Sending passwords in clear text may be a security problem in some configurations. To avoid problems if there is any possibility that the password would be intercepted, clients should connect to MySQL Server using a method that protects the password. Possibilities include [TLS / SSL](#tls), IPsec, or a private network.
|
||||||
|
|
||||||
|
|
||||||
|
##### `allowFallbackToPlaintext`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`allowFallbackToPlaintext=true` acts like a `--ssl-mode=PREFERRED` MySQL client as described in [Command Options for Connecting to the Server](https://dev.mysql.com/doc/refman/5.7/en/connection-options.html#option_general_ssl-mode)
|
||||||
|
|
||||||
|
##### `allowNativePasswords`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: true
|
||||||
|
```
|
||||||
|
`allowNativePasswords=false` disallows the usage of MySQL native password method.
|
||||||
|
|
||||||
|
##### `allowOldPasswords`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
`allowOldPasswords=true` allows the usage of the insecure old password method. This should be avoided, but is necessary in some cases. See also [the old_passwords wiki page](https://github.com/go-sql-driver/mysql/wiki/old_passwords).
|
||||||
|
|
||||||
|
##### `charset`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: string
|
||||||
|
Valid Values: <name>
|
||||||
|
Default: none
|
||||||
|
```
|
||||||
|
|
||||||
|
Sets the charset used for client-server interaction (`"SET NAMES <value>"`). If multiple charsets are set (separated by a comma), the following charset is used if setting the charset fails. This enables for example support for `utf8mb4` ([introduced in MySQL 5.5.3](http://dev.mysql.com/doc/refman/5.5/en/charset-unicode-utf8mb4.html)) with fallback to `utf8` for older servers (`charset=utf8mb4,utf8`).
|
||||||
|
|
||||||
|
See also [Unicode Support](#unicode-support).
|
||||||
|
|
||||||
|
##### `checkConnLiveness`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: true
|
||||||
|
```
|
||||||
|
|
||||||
|
On supported platforms connections retrieved from the connection pool are checked for liveness before using them. If the check fails, the respective connection is marked as bad and the query retried with another connection.
|
||||||
|
`checkConnLiveness=false` disables this liveness check of connections.
|
||||||
|
|
||||||
|
##### `collation`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: string
|
||||||
|
Valid Values: <name>
|
||||||
|
Default: utf8mb4_general_ci
|
||||||
|
```
|
||||||
|
|
||||||
|
Sets the collation used for client-server interaction on connection. In contrast to `charset`, `collation` does not issue additional queries. If the specified collation is unavailable on the target server, the connection will fail.
|
||||||
|
|
||||||
|
A list of valid charsets for a server is retrievable with `SHOW COLLATION`.
|
||||||
|
|
||||||
|
The default collation (`utf8mb4_general_ci`) is supported from MySQL 5.5. You should use an older collation (e.g. `utf8_general_ci`) for older MySQL.
|
||||||
|
|
||||||
|
Collations for charset "ucs2", "utf16", "utf16le", and "utf32" can not be used ([ref](https://dev.mysql.com/doc/refman/5.7/en/charset-connection.html#charset-connection-impermissible-client-charset)).
|
||||||
|
|
||||||
|
See also [Unicode Support](#unicode-support).
|
||||||
|
|
||||||
|
##### `clientFoundRows`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`clientFoundRows=true` causes an UPDATE to return the number of matching rows instead of the number of rows changed.
|
||||||
|
|
||||||
|
##### `columnsWithAlias`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
When `columnsWithAlias` is true, calls to `sql.Rows.Columns()` will return the table alias and the column name separated by a dot. For example:
|
||||||
|
|
||||||
|
```
|
||||||
|
SELECT u.id FROM users as u
|
||||||
|
```
|
||||||
|
|
||||||
|
will return `u.id` instead of just `id` if `columnsWithAlias=true`.
|
||||||
|
|
||||||
|
##### `compress`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
Toggles zlib compression. false by default.
|
||||||
|
|
||||||
|
##### `interpolateParams`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
If `interpolateParams` is true, placeholders (`?`) in calls to `db.Query()` and `db.Exec()` are interpolated into a single query string with given parameters. This reduces the number of roundtrips, since the driver has to prepare a statement, execute it with given parameters and close the statement again with `interpolateParams=false`.
|
||||||
|
|
||||||
|
*This can not be used together with the multibyte encodings BIG5, CP932, GB2312, GBK or SJIS. These are rejected as they may [introduce a SQL injection vulnerability](http://stackoverflow.com/a/12118602/3430118)!*
|
||||||
|
|
||||||
|
##### `loc`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: string
|
||||||
|
Valid Values: <escaped name>
|
||||||
|
Default: UTC
|
||||||
|
```
|
||||||
|
|
||||||
|
Sets the location for time.Time values (when using `parseTime=true`). *"Local"* sets the system's location. See [time.LoadLocation](https://golang.org/pkg/time/#LoadLocation) for details.
|
||||||
|
|
||||||
|
Note that this sets the location for time.Time values but does not change MySQL's [time_zone setting](https://dev.mysql.com/doc/refman/5.5/en/time-zone-support.html). For that see the [time_zone system variable](#system-variables), which can also be set as a DSN parameter.
|
||||||
|
|
||||||
|
Please keep in mind, that param values must be [url.QueryEscape](https://golang.org/pkg/net/url/#QueryEscape)'ed. Alternatively you can manually replace the `/` with `%2F`. For example `US/Pacific` would be `loc=US%2FPacific`.
|
||||||
|
|
||||||
|
##### `timeTruncate`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: duration
|
||||||
|
Default: 0
|
||||||
|
```
|
||||||
|
|
||||||
|
[Truncate time values](https://pkg.go.dev/time#Duration.Truncate) to the specified duration. The value must be a decimal number with a unit suffix (*"ms"*, *"s"*, *"m"*, *"h"*), such as *"30s"*, *"0.5m"* or *"1m30s"*.
|
||||||
|
|
||||||
|
##### `maxAllowedPacket`
|
||||||
|
```
|
||||||
|
Type: decimal number
|
||||||
|
Default: 64*1024*1024
|
||||||
|
```
|
||||||
|
|
||||||
|
Max packet size allowed in bytes. The default value is 64 MiB and should be adjusted to match the server settings. `maxAllowedPacket=0` can be used to automatically fetch the `max_allowed_packet` variable from server *on every connection*.
|
||||||
|
|
||||||
|
##### `multiStatements`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
Allow multiple statements in one query. This can be used to bach multiple queries. Use [Rows.NextResultSet()](https://pkg.go.dev/database/sql#Rows.NextResultSet) to get result of the second and subsequent queries.
|
||||||
|
|
||||||
|
When `multiStatements` is used, `?` parameters must only be used in the first statement. [interpolateParams](#interpolateparams) can be used to avoid this limitation unless prepared statement is used explicitly.
|
||||||
|
|
||||||
|
It's possible to access the last inserted ID and number of affected rows for multiple statements by using `sql.Conn.Raw()` and the `mysql.Result`. For example:
|
||||||
|
|
||||||
|
```go
|
||||||
|
conn, _ := db.Conn(ctx)
|
||||||
|
conn.Raw(func(conn any) error {
|
||||||
|
ex := conn.(driver.Execer)
|
||||||
|
res, err := ex.Exec(`
|
||||||
|
UPDATE point SET x = 1 WHERE y = 2;
|
||||||
|
UPDATE point SET x = 2 WHERE y = 3;
|
||||||
|
`, nil)
|
||||||
|
// Both slices have 2 elements.
|
||||||
|
log.Print(res.(mysql.Result).AllRowsAffected())
|
||||||
|
log.Print(res.(mysql.Result).AllLastInsertIds())
|
||||||
|
})
|
||||||
|
```
|
||||||
|
|
||||||
|
##### `parseTime`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`parseTime=true` changes the output type of `DATE` and `DATETIME` values to `time.Time` instead of `[]byte` / `string`
|
||||||
|
The date or datetime like `0000-00-00 00:00:00` is converted into zero value of `time.Time`.
|
||||||
|
|
||||||
|
|
||||||
|
##### `readTimeout`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: duration
|
||||||
|
Default: 0
|
||||||
|
```
|
||||||
|
|
||||||
|
I/O read timeout. The value must be a decimal number with a unit suffix (*"ms"*, *"s"*, *"m"*, *"h"*), such as *"30s"*, *"0.5m"* or *"1m30s"*.
|
||||||
|
|
||||||
|
##### `rejectReadOnly`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool
|
||||||
|
Valid Values: true, false
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
|
`rejectReadOnly=true` causes the driver to reject read-only connections. This
|
||||||
|
is for a possible race condition during an automatic failover, where the mysql
|
||||||
|
client gets connected to a read-only replica after the failover.
|
||||||
|
|
||||||
|
Note that this should be a fairly rare case, as an automatic failover normally
|
||||||
|
happens when the primary is down, and the race condition shouldn't happen
|
||||||
|
unless it comes back up online as soon as the failover is kicked off. On the
|
||||||
|
other hand, when this happens, a MySQL application can get stuck on a
|
||||||
|
read-only connection until restarted. It is however fairly easy to reproduce,
|
||||||
|
for example, using a manual failover on AWS Aurora's MySQL-compatible cluster.
|
||||||
|
|
||||||
|
If you are not relying on read-only transactions to reject writes that aren't
|
||||||
|
supposed to happen, setting this on some MySQL providers (such as AWS Aurora)
|
||||||
|
is safer for failovers.
|
||||||
|
|
||||||
|
Note that ERROR 1290 can be returned for a `read-only` server and this option will
|
||||||
|
cause a retry for that error. However the same error number is used for some
|
||||||
|
other cases. You should ensure your application will never cause an ERROR 1290
|
||||||
|
except for `read-only` mode when enabling this option.
|
||||||
|
|
||||||
|
|
||||||
|
##### `serverPubKey`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: string
|
||||||
|
Valid Values: <name>
|
||||||
|
Default: none
|
||||||
|
```
|
||||||
|
|
||||||
|
Server public keys can be registered with [`mysql.RegisterServerPubKey`](https://godoc.org/github.com/go-sql-driver/mysql#RegisterServerPubKey), which can then be used by the assigned name in the DSN.
|
||||||
|
Public keys are used to transmit encrypted data, e.g. for authentication.
|
||||||
|
If the server's public key is known, it should be set manually to avoid expensive and potentially insecure transmissions of the public key from the server to the client each time it is required.
|
||||||
|
|
||||||
|
|
||||||
|
##### `timeout`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: duration
|
||||||
|
Default: OS default
|
||||||
|
```
|
||||||
|
|
||||||
|
Timeout for establishing connections, aka dial timeout. The value must be a decimal number with a unit suffix (*"ms"*, *"s"*, *"m"*, *"h"*), such as *"30s"*, *"0.5m"* or *"1m30s"*.
|
||||||
|
|
||||||
|
|
||||||
|
##### `tls`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: bool / string
|
||||||
|
Valid Values: true, false, skip-verify, preferred, <name>
|
||||||
|
Default: false
|
||||||
|
```
|
||||||
|
|
||||||
|
`tls=true` enables TLS / SSL encrypted connection to the server. Use `skip-verify` if you want to use a self-signed or invalid certificate (server side) or use `preferred` to use TLS only when advertised by the server. This is similar to `skip-verify`, but additionally allows a fallback to a connection which is not encrypted. Neither `skip-verify` nor `preferred` add any reliable security. You can use a custom TLS config after registering it with [`mysql.RegisterTLSConfig`](https://godoc.org/github.com/go-sql-driver/mysql#RegisterTLSConfig).
|
||||||
|
|
||||||
|
|
||||||
|
##### `writeTimeout`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: duration
|
||||||
|
Default: 0
|
||||||
|
```
|
||||||
|
|
||||||
|
I/O write timeout. The value must be a decimal number with a unit suffix (*"ms"*, *"s"*, *"m"*, *"h"*), such as *"30s"*, *"0.5m"* or *"1m30s"*.
|
||||||
|
|
||||||
|
##### `connectionAttributes`
|
||||||
|
|
||||||
|
```
|
||||||
|
Type: comma-delimited string of user-defined "key:value" pairs
|
||||||
|
Valid Values: (<name1>:<value1>,<name2>:<value2>,...)
|
||||||
|
Default: none
|
||||||
|
```
|
||||||
|
|
||||||
|
[Connection attributes](https://dev.mysql.com/doc/refman/8.0/en/performance-schema-connection-attribute-tables.html) are key-value pairs that application programs can pass to the server at connect time.
|
||||||
|
|
||||||
|
##### System Variables
|
||||||
|
|
||||||
|
Any other parameters are interpreted as system variables:
|
||||||
|
* `<boolean_var>=<value>`: `SET <boolean_var>=<value>`
|
||||||
|
* `<enum_var>=<value>`: `SET <enum_var>=<value>`
|
||||||
|
* `<string_var>=%27<value>%27`: `SET <string_var>='<value>'`
|
||||||
|
|
||||||
|
Rules:
|
||||||
|
* The values for string variables must be quoted with `'`.
|
||||||
|
* The values must also be [url.QueryEscape](http://golang.org/pkg/net/url/#QueryEscape)'ed!
|
||||||
|
(which implies values of string variables must be wrapped with `%27`).
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
* `autocommit=1`: `SET autocommit=1`
|
||||||
|
* [`time_zone=%27Europe%2FParis%27`](https://dev.mysql.com/doc/refman/5.5/en/time-zone-support.html): `SET time_zone='Europe/Paris'`
|
||||||
|
* [`transaction_isolation=%27REPEATABLE-READ%27`](https://dev.mysql.com/doc/refman/5.7/en/server-system-variables.html#sysvar_transaction_isolation): `SET transaction_isolation='REPEATABLE-READ'`
|
||||||
|
|
||||||
|
|
||||||
|
#### Examples
|
||||||
|
```
|
||||||
|
user@unix(/path/to/socket)/dbname
|
||||||
|
```
|
||||||
|
|
||||||
|
```
|
||||||
|
root:pw@unix(/tmp/mysql.sock)/myDatabase?loc=Local
|
||||||
|
```
|
||||||
|
|
||||||
|
```
|
||||||
|
user:password@tcp(localhost:5555)/dbname?tls=skip-verify&autocommit=true
|
||||||
|
```
|
||||||
|
|
||||||
|
Treat warnings as errors by setting the system variable [`sql_mode`](https://dev.mysql.com/doc/refman/5.7/en/sql-mode.html):
|
||||||
|
```
|
||||||
|
user:password@/dbname?sql_mode=TRADITIONAL
|
||||||
|
```
|
||||||
|
|
||||||
|
TCP via IPv6:
|
||||||
|
```
|
||||||
|
user:password@tcp([de:ad:be:ef::ca:fe]:80)/dbname?timeout=90s&collation=utf8mb4_unicode_ci
|
||||||
|
```
|
||||||
|
|
||||||
|
TCP on a remote host, e.g. Amazon RDS:
|
||||||
|
```
|
||||||
|
id:password@tcp(your-amazonaws-uri.com:3306)/dbname
|
||||||
|
```
|
||||||
|
|
||||||
|
Google Cloud SQL on App Engine:
|
||||||
|
```
|
||||||
|
user:password@unix(/cloudsql/project-id:region-name:instance-name)/dbname
|
||||||
|
```
|
||||||
|
|
||||||
|
TCP using default port (3306) on localhost:
|
||||||
|
```
|
||||||
|
user:password@tcp/dbname?charset=utf8mb4,utf8&sys_var=esc%40ped
|
||||||
|
```
|
||||||
|
|
||||||
|
Use the default protocol (tcp) and host (localhost:3306):
|
||||||
|
```
|
||||||
|
user:password@/dbname
|
||||||
|
```
|
||||||
|
|
||||||
|
No Database preselected:
|
||||||
|
```
|
||||||
|
user:password@/
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
|
### Connection pool and timeouts
|
||||||
|
The connection pool is managed by Go's database/sql package. For details on how to configure the size of the pool and how long connections stay in the pool see `*DB.SetMaxOpenConns`, `*DB.SetMaxIdleConns`, and `*DB.SetConnMaxLifetime` in the [database/sql documentation](https://golang.org/pkg/database/sql/). The read, write, and dial timeouts for each individual connection are configured with the DSN parameters [`readTimeout`](#readtimeout), [`writeTimeout`](#writetimeout), and [`timeout`](#timeout), respectively.
|
||||||
|
|
||||||
|
## `ColumnType` Support
|
||||||
|
This driver supports the [`ColumnType` interface](https://golang.org/pkg/database/sql/#ColumnType) introduced in Go 1.8, with the exception of [`ColumnType.Length()`](https://golang.org/pkg/database/sql/#ColumnType.Length), which is currently not supported. All Unsigned database type names will be returned `UNSIGNED ` with `INT`, `TINYINT`, `SMALLINT`, `MEDIUMINT`, `BIGINT`.
|
||||||
|
|
||||||
|
## `context.Context` Support
|
||||||
|
Go 1.8 added `database/sql` support for `context.Context`. This driver supports query timeouts and cancellation via contexts.
|
||||||
|
See [context support in the database/sql package](https://golang.org/doc/go1.8#database_sql) for more details.
|
||||||
|
|
||||||
|
> [!IMPORTANT]
|
||||||
|
> The `QueryContext`, `ExecContext`, etc. variants provided by `database/sql` will cause the connection to be closed if the provided context is cancelled or timed out before the result is received by the driver.
|
||||||
|
|
||||||
|
|
||||||
|
### `LOAD DATA LOCAL INFILE` support
|
||||||
|
For this feature you need direct access to the package. Therefore you must change the import path (no `_`):
|
||||||
|
```go
|
||||||
|
import "github.com/go-sql-driver/mysql"
|
||||||
|
```
|
||||||
|
|
||||||
|
Files must be explicitly allowed by registering them with `mysql.RegisterLocalFile(filepath)` (recommended) or the allowlist check must be deactivated by using the DSN parameter `allowAllFiles=true` ([*Might be insecure!*](https://dev.mysql.com/doc/refman/8.0/en/load-data.html#load-data-local)).
|
||||||
|
|
||||||
|
To use a `io.Reader` a handler function must be registered with `mysql.RegisterReaderHandler(name, handler)` which returns a `io.Reader` or `io.ReadCloser`. The Reader is available with the filepath `Reader::<name>` then. Choose different names for different handlers and `DeregisterReaderHandler` when you don't need it anymore.
|
||||||
|
|
||||||
|
See the [godoc of Go-MySQL-Driver](https://godoc.org/github.com/go-sql-driver/mysql "golang mysql driver documentation") for details.
|
||||||
|
|
||||||
|
|
||||||
|
### `time.Time` support
|
||||||
|
The default internal output type of MySQL `DATE` and `DATETIME` values is `[]byte` which allows you to scan the value into a `[]byte`, `string` or `sql.RawBytes` variable in your program.
|
||||||
|
|
||||||
|
However, many want to scan MySQL `DATE` and `DATETIME` values into `time.Time` variables, which is the logical equivalent in Go to `DATE` and `DATETIME` in MySQL. You can do that by changing the internal output type from `[]byte` to `time.Time` with the DSN parameter `parseTime=true`. You can set the default [`time.Time` location](https://golang.org/pkg/time/#Location) with the `loc` DSN parameter.
|
||||||
|
|
||||||
|
**Caution:** As of Go 1.1, this makes `time.Time` the only variable type you can scan `DATE` and `DATETIME` values into. This breaks for example [`sql.RawBytes` support](https://github.com/go-sql-driver/mysql/wiki/Examples#rawbytes).
|
||||||
|
|
||||||
|
|
||||||
|
### Unicode support
|
||||||
|
Since version 1.5 Go-MySQL-Driver automatically uses the collation ` utf8mb4_general_ci` by default.
|
||||||
|
|
||||||
|
Other charsets / collations can be set using the [`charset`](#charset) or [`collation`](#collation) DSN parameter.
|
||||||
|
|
||||||
|
- When only the `charset` is specified, the `SET NAMES <charset>` query is sent and the server's default collation is used.
|
||||||
|
- When both the `charset` and `collation` are specified, the `SET NAMES <charset> COLLATE <collation>` query is sent.
|
||||||
|
- When only the `collation` is specified, the collation is specified in the protocol handshake and the `SET NAMES` query is not sent. This can save one roundtrip, but note that the server may ignore the specified collation silently and use the server's default charset/collation instead.
|
||||||
|
|
||||||
|
See http://dev.mysql.com/doc/refman/8.0/en/charset-unicode.html for more details on MySQL's Unicode support.
|
||||||
|
|
||||||
|
## Testing / Development
|
||||||
|
To run the driver tests you may need to adjust the configuration. See the [Testing Wiki-Page](https://github.com/go-sql-driver/mysql/wiki/Testing "Testing") for details.
|
||||||
|
|
||||||
|
Go-MySQL-Driver is not feature-complete yet. Your help is very appreciated.
|
||||||
|
If you want to contribute, you can work on an [open issue](https://github.com/go-sql-driver/mysql/issues?state=open) or review a [pull request](https://github.com/go-sql-driver/mysql/pulls).
|
||||||
|
|
||||||
|
See the [Contribution Guidelines](https://github.com/go-sql-driver/mysql/blob/master/.github/CONTRIBUTING.md) for details.
|
||||||
|
|
||||||
|
---------------------------------------
|
||||||
|
|
||||||
|
## License
|
||||||
|
Go-MySQL-Driver is licensed under the [Mozilla Public License Version 2.0](https://raw.github.com/go-sql-driver/mysql/master/LICENSE)
|
||||||
|
|
||||||
|
Mozilla summarizes the license scope as follows:
|
||||||
|
> MPL: The copyleft applies to any files containing MPLed code.
|
||||||
|
|
||||||
|
|
||||||
|
That means:
|
||||||
|
* You can **use** the **unchanged** source code both in private and commercially.
|
||||||
|
* When distributing, you **must publish** the source code of any **changed files** licensed under the MPL 2.0 under a) the MPL 2.0 itself or b) a compatible license (e.g. GPL 3.0 or Apache License 2.0).
|
||||||
|
* You **needn't publish** the source code of your library as long as the files licensed under the MPL 2.0 are **unchanged**.
|
||||||
|
|
||||||
|
Please read the [MPL 2.0 FAQ](https://www.mozilla.org/en-US/MPL/2.0/FAQ/) if you have further questions regarding the license.
|
||||||
|
|
||||||
|
You can read the full terms here: [LICENSE](https://raw.github.com/go-sql-driver/mysql/master/LICENSE).
|
||||||
|
|
||||||
|

|
||||||
+484
@@ -0,0 +1,484 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2018 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/rsa"
|
||||||
|
"crypto/sha1"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/sha512"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"filippo.io/edwards25519"
|
||||||
|
)
|
||||||
|
|
||||||
|
// server pub keys registry
|
||||||
|
var (
|
||||||
|
serverPubKeyLock sync.RWMutex
|
||||||
|
serverPubKeyRegistry map[string]*rsa.PublicKey
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterServerPubKey registers a server RSA public key which can be used to
|
||||||
|
// send data in a secure manner to the server without receiving the public key
|
||||||
|
// in a potentially insecure way from the server first.
|
||||||
|
// Registered keys can afterwards be used adding serverPubKey=<name> to the DSN.
|
||||||
|
//
|
||||||
|
// Note: The provided rsa.PublicKey instance is exclusively owned by the driver
|
||||||
|
// after registering it and may not be modified.
|
||||||
|
//
|
||||||
|
// data, err := os.ReadFile("mykey.pem")
|
||||||
|
// if err != nil {
|
||||||
|
// log.Fatal(err)
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// block, _ := pem.Decode(data)
|
||||||
|
// if block == nil || block.Type != "PUBLIC KEY" {
|
||||||
|
// log.Fatal("failed to decode PEM block containing public key")
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// pub, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||||
|
// if err != nil {
|
||||||
|
// log.Fatal(err)
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// if rsaPubKey, ok := pub.(*rsa.PublicKey); ok {
|
||||||
|
// mysql.RegisterServerPubKey("mykey", rsaPubKey)
|
||||||
|
// } else {
|
||||||
|
// log.Fatal("not a RSA public key")
|
||||||
|
// }
|
||||||
|
func RegisterServerPubKey(name string, pubKey *rsa.PublicKey) {
|
||||||
|
serverPubKeyLock.Lock()
|
||||||
|
if serverPubKeyRegistry == nil {
|
||||||
|
serverPubKeyRegistry = make(map[string]*rsa.PublicKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
serverPubKeyRegistry[name] = pubKey
|
||||||
|
serverPubKeyLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeregisterServerPubKey removes the public key registered with the given name.
|
||||||
|
func DeregisterServerPubKey(name string) {
|
||||||
|
serverPubKeyLock.Lock()
|
||||||
|
if serverPubKeyRegistry != nil {
|
||||||
|
delete(serverPubKeyRegistry, name)
|
||||||
|
}
|
||||||
|
serverPubKeyLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func getServerPubKey(name string) (pubKey *rsa.PublicKey) {
|
||||||
|
serverPubKeyLock.RLock()
|
||||||
|
if v, ok := serverPubKeyRegistry[name]; ok {
|
||||||
|
pubKey = v
|
||||||
|
}
|
||||||
|
serverPubKeyLock.RUnlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash password using pre 4.1 (old password) method
|
||||||
|
// https://github.com/atcurtis/mariadb/blob/master/mysys/my_rnd.c
|
||||||
|
type myRnd struct {
|
||||||
|
seed1, seed2 uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
const myRndMaxVal = 0x3FFFFFFF
|
||||||
|
|
||||||
|
// Pseudo random number generator
|
||||||
|
func newMyRnd(seed1, seed2 uint32) *myRnd {
|
||||||
|
return &myRnd{
|
||||||
|
seed1: seed1 % myRndMaxVal,
|
||||||
|
seed2: seed2 % myRndMaxVal,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tested to be equivalent to MariaDB's floating point variant
|
||||||
|
// http://play.golang.org/p/QHvhd4qved
|
||||||
|
// http://play.golang.org/p/RG0q4ElWDx
|
||||||
|
func (r *myRnd) NextByte() byte {
|
||||||
|
r.seed1 = (r.seed1*3 + r.seed2) % myRndMaxVal
|
||||||
|
r.seed2 = (r.seed1 + r.seed2 + 33) % myRndMaxVal
|
||||||
|
|
||||||
|
return byte(uint64(r.seed1) * 31 / myRndMaxVal)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generate binary hash from byte string using insecure pre 4.1 method
|
||||||
|
func pwHash(password []byte) (result [2]uint32) {
|
||||||
|
var add uint32 = 7
|
||||||
|
var tmp uint32
|
||||||
|
|
||||||
|
result[0] = 1345345333
|
||||||
|
result[1] = 0x12345671
|
||||||
|
|
||||||
|
for _, c := range password {
|
||||||
|
// skip spaces and tabs in password
|
||||||
|
if c == ' ' || c == '\t' {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
tmp = uint32(c)
|
||||||
|
result[0] ^= (((result[0] & 63) + add) * tmp) + (result[0] << 8)
|
||||||
|
result[1] += (result[1] << 8) ^ result[0]
|
||||||
|
add += tmp
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove sign bit (1<<31)-1)
|
||||||
|
result[0] &= 0x7FFFFFFF
|
||||||
|
result[1] &= 0x7FFFFFFF
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash password using insecure pre 4.1 method
|
||||||
|
func scrambleOldPassword(scramble []byte, password string) []byte {
|
||||||
|
scramble = scramble[:8]
|
||||||
|
|
||||||
|
hashPw := pwHash([]byte(password))
|
||||||
|
hashSc := pwHash(scramble)
|
||||||
|
|
||||||
|
r := newMyRnd(hashPw[0]^hashSc[0], hashPw[1]^hashSc[1])
|
||||||
|
|
||||||
|
var out [8]byte
|
||||||
|
for i := range out {
|
||||||
|
out[i] = r.NextByte() + 64
|
||||||
|
}
|
||||||
|
|
||||||
|
mask := r.NextByte()
|
||||||
|
for i := range out {
|
||||||
|
out[i] ^= mask
|
||||||
|
}
|
||||||
|
|
||||||
|
return out[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash password using 4.1+ method (SHA1)
|
||||||
|
func scramblePassword(scramble []byte, password string) []byte {
|
||||||
|
if len(password) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stage1Hash = SHA1(password)
|
||||||
|
crypt := sha1.New()
|
||||||
|
crypt.Write([]byte(password))
|
||||||
|
stage1 := crypt.Sum(nil)
|
||||||
|
|
||||||
|
// scrambleHash = SHA1(scramble + SHA1(stage1Hash))
|
||||||
|
// inner Hash
|
||||||
|
crypt.Reset()
|
||||||
|
crypt.Write(stage1)
|
||||||
|
hash := crypt.Sum(nil)
|
||||||
|
|
||||||
|
// outer Hash
|
||||||
|
crypt.Reset()
|
||||||
|
crypt.Write(scramble)
|
||||||
|
crypt.Write(hash)
|
||||||
|
scramble = crypt.Sum(nil)
|
||||||
|
|
||||||
|
// token = scrambleHash XOR stage1Hash
|
||||||
|
for i := range scramble {
|
||||||
|
scramble[i] ^= stage1[i]
|
||||||
|
}
|
||||||
|
return scramble
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash password using MySQL 8+ method (SHA256)
|
||||||
|
func scrambleSHA256Password(scramble []byte, password string) []byte {
|
||||||
|
if len(password) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// XOR(SHA256(password), SHA256(SHA256(SHA256(password)), scramble))
|
||||||
|
|
||||||
|
crypt := sha256.New()
|
||||||
|
crypt.Write([]byte(password))
|
||||||
|
message1 := crypt.Sum(nil)
|
||||||
|
|
||||||
|
crypt.Reset()
|
||||||
|
crypt.Write(message1)
|
||||||
|
message1Hash := crypt.Sum(nil)
|
||||||
|
|
||||||
|
crypt.Reset()
|
||||||
|
crypt.Write(message1Hash)
|
||||||
|
crypt.Write(scramble)
|
||||||
|
message2 := crypt.Sum(nil)
|
||||||
|
|
||||||
|
for i := range message1 {
|
||||||
|
message1[i] ^= message2[i]
|
||||||
|
}
|
||||||
|
|
||||||
|
return message1
|
||||||
|
}
|
||||||
|
|
||||||
|
func encryptPassword(password string, seed []byte, pub *rsa.PublicKey) ([]byte, error) {
|
||||||
|
plain := make([]byte, len(password)+1)
|
||||||
|
copy(plain, password)
|
||||||
|
for i := range plain {
|
||||||
|
j := i % len(seed)
|
||||||
|
plain[i] ^= seed[j]
|
||||||
|
}
|
||||||
|
sha1 := sha1.New()
|
||||||
|
return rsa.EncryptOAEP(sha1, rand.Reader, pub, plain, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// authEd25519 does ed25519 authentication used by MariaDB.
|
||||||
|
func authEd25519(scramble []byte, password string) ([]byte, error) {
|
||||||
|
// Derived from https://github.com/MariaDB/server/blob/d8e6bb00888b1f82c031938f4c8ac5d97f6874c3/plugin/auth_ed25519/ref10/sign.c
|
||||||
|
// Code style is from https://cs.opensource.google/go/go/+/refs/tags/go1.21.5:src/crypto/ed25519/ed25519.go;l=207
|
||||||
|
h := sha512.Sum512([]byte(password))
|
||||||
|
|
||||||
|
s, err := edwards25519.NewScalar().SetBytesWithClamping(h[:32])
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
A := (&edwards25519.Point{}).ScalarBaseMult(s)
|
||||||
|
|
||||||
|
mh := sha512.New()
|
||||||
|
mh.Write(h[32:])
|
||||||
|
mh.Write(scramble)
|
||||||
|
messageDigest := mh.Sum(nil)
|
||||||
|
r, err := edwards25519.NewScalar().SetUniformBytes(messageDigest)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
R := (&edwards25519.Point{}).ScalarBaseMult(r)
|
||||||
|
|
||||||
|
kh := sha512.New()
|
||||||
|
kh.Write(R.Bytes())
|
||||||
|
kh.Write(A.Bytes())
|
||||||
|
kh.Write(scramble)
|
||||||
|
hramDigest := kh.Sum(nil)
|
||||||
|
k, err := edwards25519.NewScalar().SetUniformBytes(hramDigest)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
S := k.MultiplyAdd(k, s, r)
|
||||||
|
|
||||||
|
return append(R.Bytes(), S.Bytes()...), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) sendEncryptedPassword(seed []byte, pub *rsa.PublicKey) error {
|
||||||
|
enc, err := encryptPassword(mc.cfg.Passwd, seed, pub)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return mc.writeAuthSwitchPacket(enc)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) auth(authData []byte, plugin string) ([]byte, error) {
|
||||||
|
switch plugin {
|
||||||
|
case "caching_sha2_password":
|
||||||
|
authResp := scrambleSHA256Password(authData, mc.cfg.Passwd)
|
||||||
|
return authResp, nil
|
||||||
|
|
||||||
|
case "mysql_old_password":
|
||||||
|
if !mc.cfg.AllowOldPasswords {
|
||||||
|
return nil, ErrOldPassword
|
||||||
|
}
|
||||||
|
if len(mc.cfg.Passwd) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
// Note: there are edge cases where this should work but doesn't;
|
||||||
|
// this is currently "wontfix":
|
||||||
|
// https://github.com/go-sql-driver/mysql/issues/184
|
||||||
|
authResp := append(scrambleOldPassword(authData[:8], mc.cfg.Passwd), 0)
|
||||||
|
return authResp, nil
|
||||||
|
|
||||||
|
case "mysql_clear_password":
|
||||||
|
if !mc.cfg.AllowCleartextPasswords {
|
||||||
|
return nil, ErrCleartextPassword
|
||||||
|
}
|
||||||
|
// http://dev.mysql.com/doc/refman/5.7/en/cleartext-authentication-plugin.html
|
||||||
|
// http://dev.mysql.com/doc/refman/5.7/en/pam-authentication-plugin.html
|
||||||
|
return append([]byte(mc.cfg.Passwd), 0), nil
|
||||||
|
|
||||||
|
case "mysql_native_password":
|
||||||
|
if !mc.cfg.AllowNativePasswords {
|
||||||
|
return nil, ErrNativePassword
|
||||||
|
}
|
||||||
|
// https://dev.mysql.com/doc/internals/en/secure-password-authentication.html
|
||||||
|
// Native password authentication only need and will need 20-byte challenge.
|
||||||
|
authResp := scramblePassword(authData[:20], mc.cfg.Passwd)
|
||||||
|
return authResp, nil
|
||||||
|
|
||||||
|
case "sha256_password":
|
||||||
|
if len(mc.cfg.Passwd) == 0 {
|
||||||
|
return []byte{0}, nil
|
||||||
|
}
|
||||||
|
// unlike caching_sha2_password, sha256_password does not accept
|
||||||
|
// cleartext password on unix transport.
|
||||||
|
if mc.cfg.TLS != nil {
|
||||||
|
// write cleartext auth packet
|
||||||
|
return append([]byte(mc.cfg.Passwd), 0), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
pubKey := mc.cfg.pubKey
|
||||||
|
if pubKey == nil {
|
||||||
|
// request public key from server
|
||||||
|
return []byte{1}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// encrypted password
|
||||||
|
enc, err := encryptPassword(mc.cfg.Passwd, authData, pubKey)
|
||||||
|
return enc, err
|
||||||
|
|
||||||
|
case "client_ed25519":
|
||||||
|
if len(authData) != 32 {
|
||||||
|
return nil, ErrMalformPkt
|
||||||
|
}
|
||||||
|
return authEd25519(authData, mc.cfg.Passwd)
|
||||||
|
|
||||||
|
default:
|
||||||
|
mc.log("unknown auth plugin:", plugin)
|
||||||
|
return nil, ErrUnknownPlugin
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) handleAuthResult(oldAuthData []byte, plugin string) error {
|
||||||
|
// Read Result Packet
|
||||||
|
authData, newPlugin, err := mc.readAuthResult()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// handle auth plugin switch, if requested
|
||||||
|
if newPlugin != "" {
|
||||||
|
// If CLIENT_PLUGIN_AUTH capability is not supported, no new cipher is
|
||||||
|
// sent and we have to keep using the cipher sent in the init packet.
|
||||||
|
if authData == nil {
|
||||||
|
authData = oldAuthData
|
||||||
|
} else {
|
||||||
|
// copy data from read buffer to owned slice
|
||||||
|
copy(oldAuthData, authData)
|
||||||
|
}
|
||||||
|
|
||||||
|
plugin = newPlugin
|
||||||
|
|
||||||
|
authResp, err := mc.auth(authData, plugin)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err = mc.writeAuthSwitchPacket(authResp); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read Result Packet
|
||||||
|
authData, newPlugin, err = mc.readAuthResult()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Do not allow to change the auth plugin more than once
|
||||||
|
if newPlugin != "" {
|
||||||
|
return ErrMalformPkt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch plugin {
|
||||||
|
|
||||||
|
// https://dev.mysql.com/blog-archive/preparing-your-community-connector-for-mysql-8-part-2-sha256/
|
||||||
|
case "caching_sha2_password":
|
||||||
|
switch len(authData) {
|
||||||
|
case 0:
|
||||||
|
return nil // auth successful
|
||||||
|
case 1:
|
||||||
|
switch authData[0] {
|
||||||
|
case cachingSha2PasswordFastAuthSuccess:
|
||||||
|
if err = mc.resultUnchanged().readResultOK(); err == nil {
|
||||||
|
return nil // auth successful
|
||||||
|
}
|
||||||
|
|
||||||
|
case cachingSha2PasswordPerformFullAuthentication:
|
||||||
|
if mc.cfg.TLS != nil || mc.cfg.Net == "unix" {
|
||||||
|
// write cleartext auth packet
|
||||||
|
err = mc.writeAuthSwitchPacket(append([]byte(mc.cfg.Passwd), 0))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
pubKey := mc.cfg.pubKey
|
||||||
|
if pubKey == nil {
|
||||||
|
// request public key from server
|
||||||
|
data, err := mc.buf.takeSmallBuffer(4 + 1)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
data[4] = cachingSha2PasswordRequestPublicKey
|
||||||
|
err = mc.writePacket(data)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if data, err = mc.readPacket(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if data[0] != iAuthMoreData {
|
||||||
|
return fmt.Errorf("unexpected resp from server for caching_sha2_password, perform full authentication")
|
||||||
|
}
|
||||||
|
|
||||||
|
// parse public key
|
||||||
|
block, rest := pem.Decode(data[1:])
|
||||||
|
if block == nil {
|
||||||
|
return fmt.Errorf("no pem data found, data: %s", rest)
|
||||||
|
}
|
||||||
|
pkix, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
pubKey = pkix.(*rsa.PublicKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
// send encrypted password
|
||||||
|
err = mc.sendEncryptedPassword(oldAuthData, pubKey)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return mc.resultUnchanged().readResultOK()
|
||||||
|
|
||||||
|
default:
|
||||||
|
return ErrMalformPkt
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return ErrMalformPkt
|
||||||
|
}
|
||||||
|
|
||||||
|
case "sha256_password":
|
||||||
|
switch len(authData) {
|
||||||
|
case 0:
|
||||||
|
return nil // auth successful
|
||||||
|
default:
|
||||||
|
block, _ := pem.Decode(authData)
|
||||||
|
if block == nil {
|
||||||
|
return fmt.Errorf("no Pem data found, data: %s", authData)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// send encrypted password
|
||||||
|
err = mc.sendEncryptedPassword(oldAuthData, pub.(*rsa.PublicKey))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return mc.resultUnchanged().readResultOK()
|
||||||
|
}
|
||||||
|
|
||||||
|
default:
|
||||||
|
return nil // auth successful
|
||||||
|
}
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
+149
@@ -0,0 +1,149 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2013 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
const defaultBufSize = 4096
|
||||||
|
const maxCachedBufSize = 256 * 1024
|
||||||
|
|
||||||
|
// readerFunc is a function that compatible with io.Reader.
|
||||||
|
// We use this function type instead of io.Reader because we want to
|
||||||
|
// just pass mc.readWithTimeout.
|
||||||
|
type readerFunc func([]byte) (int, error)
|
||||||
|
|
||||||
|
// A buffer which is used for both reading and writing.
|
||||||
|
// This is possible since communication on each connection is synchronous.
|
||||||
|
// In other words, we can't write and read simultaneously on the same connection.
|
||||||
|
// The buffer is similar to bufio.Reader / Writer but zero-copy-ish
|
||||||
|
// Also highly optimized for this particular use case.
|
||||||
|
type buffer struct {
|
||||||
|
buf []byte // read buffer.
|
||||||
|
cachedBuf []byte // buffer that will be reused. len(cachedBuf) <= maxCachedBufSize.
|
||||||
|
}
|
||||||
|
|
||||||
|
// newBuffer allocates and returns a new buffer.
|
||||||
|
func newBuffer() buffer {
|
||||||
|
return buffer{
|
||||||
|
cachedBuf: make([]byte, defaultBufSize),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// busy returns true if the read buffer is not empty.
|
||||||
|
func (b *buffer) busy() bool {
|
||||||
|
return len(b.buf) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// len returns how many bytes in the read buffer.
|
||||||
|
func (b *buffer) len() int {
|
||||||
|
return len(b.buf)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fill reads into the read buffer until at least _need_ bytes are in it.
|
||||||
|
func (b *buffer) fill(need int, r readerFunc) error {
|
||||||
|
// we'll move the contents of the current buffer to dest before filling it.
|
||||||
|
dest := b.cachedBuf
|
||||||
|
|
||||||
|
// grow buffer if necessary to fit the whole packet.
|
||||||
|
if need > len(dest) {
|
||||||
|
// Round up to the next multiple of the default size
|
||||||
|
dest = make([]byte, ((need/defaultBufSize)+1)*defaultBufSize)
|
||||||
|
|
||||||
|
// if the allocated buffer is not too large, move it to backing storage
|
||||||
|
// to prevent extra allocations on applications that perform large reads
|
||||||
|
if len(dest) <= maxCachedBufSize {
|
||||||
|
b.cachedBuf = dest
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// move the existing data to the start of the buffer.
|
||||||
|
n := len(b.buf)
|
||||||
|
copy(dest[:n], b.buf)
|
||||||
|
|
||||||
|
for {
|
||||||
|
nn, err := r(dest[n:])
|
||||||
|
n += nn
|
||||||
|
|
||||||
|
if err == nil && n < need {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b.buf = dest[:n]
|
||||||
|
|
||||||
|
if err == io.EOF {
|
||||||
|
if n < need {
|
||||||
|
err = io.ErrUnexpectedEOF
|
||||||
|
} else {
|
||||||
|
err = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// returns next N bytes from buffer.
|
||||||
|
// The returned slice is only guaranteed to be valid until the next read
|
||||||
|
func (b *buffer) readNext(need int) []byte {
|
||||||
|
data := b.buf[:need:need]
|
||||||
|
b.buf = b.buf[need:]
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeBuffer returns a buffer with the requested size.
|
||||||
|
// If possible, a slice from the existing buffer is returned.
|
||||||
|
// Otherwise a bigger buffer is made.
|
||||||
|
// Only one buffer (total) can be used at a time.
|
||||||
|
func (b *buffer) takeBuffer(length int) ([]byte, error) {
|
||||||
|
if b.busy() {
|
||||||
|
return nil, ErrBusyBuffer
|
||||||
|
}
|
||||||
|
|
||||||
|
// test (cheap) general case first
|
||||||
|
if length <= len(b.cachedBuf) {
|
||||||
|
return b.cachedBuf[:length], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if length < maxCachedBufSize {
|
||||||
|
b.cachedBuf = make([]byte, length)
|
||||||
|
return b.cachedBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// buffer is larger than we want to store.
|
||||||
|
return make([]byte, length), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeSmallBuffer is shortcut which can be used if length is
|
||||||
|
// known to be smaller than defaultBufSize.
|
||||||
|
// Only one buffer (total) can be used at a time.
|
||||||
|
func (b *buffer) takeSmallBuffer(length int) ([]byte, error) {
|
||||||
|
if b.busy() {
|
||||||
|
return nil, ErrBusyBuffer
|
||||||
|
}
|
||||||
|
return b.cachedBuf[:length], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// takeCompleteBuffer returns the complete existing buffer.
|
||||||
|
// This can be used if the necessary buffer size is unknown.
|
||||||
|
// cap and len of the returned buffer will be equal.
|
||||||
|
// Only one buffer (total) can be used at a time.
|
||||||
|
func (b *buffer) takeCompleteBuffer() ([]byte, error) {
|
||||||
|
if b.busy() {
|
||||||
|
return nil, ErrBusyBuffer
|
||||||
|
}
|
||||||
|
return b.cachedBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// store stores buf, an updated buffer, if its suitable to do so.
|
||||||
|
func (b *buffer) store(buf []byte) {
|
||||||
|
if cap(buf) <= maxCachedBufSize && cap(buf) > cap(b.cachedBuf) {
|
||||||
|
b.cachedBuf = buf[:cap(buf)]
|
||||||
|
}
|
||||||
|
}
|
||||||
+266
@@ -0,0 +1,266 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2014 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
const defaultCollationID = 45 // utf8mb4_general_ci
|
||||||
|
const binaryCollationID = 63
|
||||||
|
|
||||||
|
// A list of available collations mapped to the internal ID.
|
||||||
|
// To update this map use the following MySQL query:
|
||||||
|
//
|
||||||
|
// SELECT COLLATION_NAME, ID FROM information_schema.COLLATIONS WHERE ID<256 ORDER BY ID
|
||||||
|
//
|
||||||
|
// Handshake packet have only 1 byte for collation_id. So we can't use collations with ID > 255.
|
||||||
|
//
|
||||||
|
// ucs2, utf16, and utf32 can't be used for connection charset.
|
||||||
|
// https://dev.mysql.com/doc/refman/5.7/en/charset-connection.html#charset-connection-impermissible-client-charset
|
||||||
|
// They are commented out to reduce this map.
|
||||||
|
var collations = map[string]byte{
|
||||||
|
"big5_chinese_ci": 1,
|
||||||
|
"latin2_czech_cs": 2,
|
||||||
|
"dec8_swedish_ci": 3,
|
||||||
|
"cp850_general_ci": 4,
|
||||||
|
"latin1_german1_ci": 5,
|
||||||
|
"hp8_english_ci": 6,
|
||||||
|
"koi8r_general_ci": 7,
|
||||||
|
"latin1_swedish_ci": 8,
|
||||||
|
"latin2_general_ci": 9,
|
||||||
|
"swe7_swedish_ci": 10,
|
||||||
|
"ascii_general_ci": 11,
|
||||||
|
"ujis_japanese_ci": 12,
|
||||||
|
"sjis_japanese_ci": 13,
|
||||||
|
"cp1251_bulgarian_ci": 14,
|
||||||
|
"latin1_danish_ci": 15,
|
||||||
|
"hebrew_general_ci": 16,
|
||||||
|
"tis620_thai_ci": 18,
|
||||||
|
"euckr_korean_ci": 19,
|
||||||
|
"latin7_estonian_cs": 20,
|
||||||
|
"latin2_hungarian_ci": 21,
|
||||||
|
"koi8u_general_ci": 22,
|
||||||
|
"cp1251_ukrainian_ci": 23,
|
||||||
|
"gb2312_chinese_ci": 24,
|
||||||
|
"greek_general_ci": 25,
|
||||||
|
"cp1250_general_ci": 26,
|
||||||
|
"latin2_croatian_ci": 27,
|
||||||
|
"gbk_chinese_ci": 28,
|
||||||
|
"cp1257_lithuanian_ci": 29,
|
||||||
|
"latin5_turkish_ci": 30,
|
||||||
|
"latin1_german2_ci": 31,
|
||||||
|
"armscii8_general_ci": 32,
|
||||||
|
"utf8_general_ci": 33,
|
||||||
|
"cp1250_czech_cs": 34,
|
||||||
|
//"ucs2_general_ci": 35,
|
||||||
|
"cp866_general_ci": 36,
|
||||||
|
"keybcs2_general_ci": 37,
|
||||||
|
"macce_general_ci": 38,
|
||||||
|
"macroman_general_ci": 39,
|
||||||
|
"cp852_general_ci": 40,
|
||||||
|
"latin7_general_ci": 41,
|
||||||
|
"latin7_general_cs": 42,
|
||||||
|
"macce_bin": 43,
|
||||||
|
"cp1250_croatian_ci": 44,
|
||||||
|
"utf8mb4_general_ci": 45,
|
||||||
|
"utf8mb4_bin": 46,
|
||||||
|
"latin1_bin": 47,
|
||||||
|
"latin1_general_ci": 48,
|
||||||
|
"latin1_general_cs": 49,
|
||||||
|
"cp1251_bin": 50,
|
||||||
|
"cp1251_general_ci": 51,
|
||||||
|
"cp1251_general_cs": 52,
|
||||||
|
"macroman_bin": 53,
|
||||||
|
//"utf16_general_ci": 54,
|
||||||
|
//"utf16_bin": 55,
|
||||||
|
//"utf16le_general_ci": 56,
|
||||||
|
"cp1256_general_ci": 57,
|
||||||
|
"cp1257_bin": 58,
|
||||||
|
"cp1257_general_ci": 59,
|
||||||
|
//"utf32_general_ci": 60,
|
||||||
|
//"utf32_bin": 61,
|
||||||
|
//"utf16le_bin": 62,
|
||||||
|
"binary": 63,
|
||||||
|
"armscii8_bin": 64,
|
||||||
|
"ascii_bin": 65,
|
||||||
|
"cp1250_bin": 66,
|
||||||
|
"cp1256_bin": 67,
|
||||||
|
"cp866_bin": 68,
|
||||||
|
"dec8_bin": 69,
|
||||||
|
"greek_bin": 70,
|
||||||
|
"hebrew_bin": 71,
|
||||||
|
"hp8_bin": 72,
|
||||||
|
"keybcs2_bin": 73,
|
||||||
|
"koi8r_bin": 74,
|
||||||
|
"koi8u_bin": 75,
|
||||||
|
"utf8_tolower_ci": 76,
|
||||||
|
"latin2_bin": 77,
|
||||||
|
"latin5_bin": 78,
|
||||||
|
"latin7_bin": 79,
|
||||||
|
"cp850_bin": 80,
|
||||||
|
"cp852_bin": 81,
|
||||||
|
"swe7_bin": 82,
|
||||||
|
"utf8_bin": 83,
|
||||||
|
"big5_bin": 84,
|
||||||
|
"euckr_bin": 85,
|
||||||
|
"gb2312_bin": 86,
|
||||||
|
"gbk_bin": 87,
|
||||||
|
"sjis_bin": 88,
|
||||||
|
"tis620_bin": 89,
|
||||||
|
//"ucs2_bin": 90,
|
||||||
|
"ujis_bin": 91,
|
||||||
|
"geostd8_general_ci": 92,
|
||||||
|
"geostd8_bin": 93,
|
||||||
|
"latin1_spanish_ci": 94,
|
||||||
|
"cp932_japanese_ci": 95,
|
||||||
|
"cp932_bin": 96,
|
||||||
|
"eucjpms_japanese_ci": 97,
|
||||||
|
"eucjpms_bin": 98,
|
||||||
|
"cp1250_polish_ci": 99,
|
||||||
|
//"utf16_unicode_ci": 101,
|
||||||
|
//"utf16_icelandic_ci": 102,
|
||||||
|
//"utf16_latvian_ci": 103,
|
||||||
|
//"utf16_romanian_ci": 104,
|
||||||
|
//"utf16_slovenian_ci": 105,
|
||||||
|
//"utf16_polish_ci": 106,
|
||||||
|
//"utf16_estonian_ci": 107,
|
||||||
|
//"utf16_spanish_ci": 108,
|
||||||
|
//"utf16_swedish_ci": 109,
|
||||||
|
//"utf16_turkish_ci": 110,
|
||||||
|
//"utf16_czech_ci": 111,
|
||||||
|
//"utf16_danish_ci": 112,
|
||||||
|
//"utf16_lithuanian_ci": 113,
|
||||||
|
//"utf16_slovak_ci": 114,
|
||||||
|
//"utf16_spanish2_ci": 115,
|
||||||
|
//"utf16_roman_ci": 116,
|
||||||
|
//"utf16_persian_ci": 117,
|
||||||
|
//"utf16_esperanto_ci": 118,
|
||||||
|
//"utf16_hungarian_ci": 119,
|
||||||
|
//"utf16_sinhala_ci": 120,
|
||||||
|
//"utf16_german2_ci": 121,
|
||||||
|
//"utf16_croatian_ci": 122,
|
||||||
|
//"utf16_unicode_520_ci": 123,
|
||||||
|
//"utf16_vietnamese_ci": 124,
|
||||||
|
//"ucs2_unicode_ci": 128,
|
||||||
|
//"ucs2_icelandic_ci": 129,
|
||||||
|
//"ucs2_latvian_ci": 130,
|
||||||
|
//"ucs2_romanian_ci": 131,
|
||||||
|
//"ucs2_slovenian_ci": 132,
|
||||||
|
//"ucs2_polish_ci": 133,
|
||||||
|
//"ucs2_estonian_ci": 134,
|
||||||
|
//"ucs2_spanish_ci": 135,
|
||||||
|
//"ucs2_swedish_ci": 136,
|
||||||
|
//"ucs2_turkish_ci": 137,
|
||||||
|
//"ucs2_czech_ci": 138,
|
||||||
|
//"ucs2_danish_ci": 139,
|
||||||
|
//"ucs2_lithuanian_ci": 140,
|
||||||
|
//"ucs2_slovak_ci": 141,
|
||||||
|
//"ucs2_spanish2_ci": 142,
|
||||||
|
//"ucs2_roman_ci": 143,
|
||||||
|
//"ucs2_persian_ci": 144,
|
||||||
|
//"ucs2_esperanto_ci": 145,
|
||||||
|
//"ucs2_hungarian_ci": 146,
|
||||||
|
//"ucs2_sinhala_ci": 147,
|
||||||
|
//"ucs2_german2_ci": 148,
|
||||||
|
//"ucs2_croatian_ci": 149,
|
||||||
|
//"ucs2_unicode_520_ci": 150,
|
||||||
|
//"ucs2_vietnamese_ci": 151,
|
||||||
|
//"ucs2_general_mysql500_ci": 159,
|
||||||
|
//"utf32_unicode_ci": 160,
|
||||||
|
//"utf32_icelandic_ci": 161,
|
||||||
|
//"utf32_latvian_ci": 162,
|
||||||
|
//"utf32_romanian_ci": 163,
|
||||||
|
//"utf32_slovenian_ci": 164,
|
||||||
|
//"utf32_polish_ci": 165,
|
||||||
|
//"utf32_estonian_ci": 166,
|
||||||
|
//"utf32_spanish_ci": 167,
|
||||||
|
//"utf32_swedish_ci": 168,
|
||||||
|
//"utf32_turkish_ci": 169,
|
||||||
|
//"utf32_czech_ci": 170,
|
||||||
|
//"utf32_danish_ci": 171,
|
||||||
|
//"utf32_lithuanian_ci": 172,
|
||||||
|
//"utf32_slovak_ci": 173,
|
||||||
|
//"utf32_spanish2_ci": 174,
|
||||||
|
//"utf32_roman_ci": 175,
|
||||||
|
//"utf32_persian_ci": 176,
|
||||||
|
//"utf32_esperanto_ci": 177,
|
||||||
|
//"utf32_hungarian_ci": 178,
|
||||||
|
//"utf32_sinhala_ci": 179,
|
||||||
|
//"utf32_german2_ci": 180,
|
||||||
|
//"utf32_croatian_ci": 181,
|
||||||
|
//"utf32_unicode_520_ci": 182,
|
||||||
|
//"utf32_vietnamese_ci": 183,
|
||||||
|
"utf8_unicode_ci": 192,
|
||||||
|
"utf8_icelandic_ci": 193,
|
||||||
|
"utf8_latvian_ci": 194,
|
||||||
|
"utf8_romanian_ci": 195,
|
||||||
|
"utf8_slovenian_ci": 196,
|
||||||
|
"utf8_polish_ci": 197,
|
||||||
|
"utf8_estonian_ci": 198,
|
||||||
|
"utf8_spanish_ci": 199,
|
||||||
|
"utf8_swedish_ci": 200,
|
||||||
|
"utf8_turkish_ci": 201,
|
||||||
|
"utf8_czech_ci": 202,
|
||||||
|
"utf8_danish_ci": 203,
|
||||||
|
"utf8_lithuanian_ci": 204,
|
||||||
|
"utf8_slovak_ci": 205,
|
||||||
|
"utf8_spanish2_ci": 206,
|
||||||
|
"utf8_roman_ci": 207,
|
||||||
|
"utf8_persian_ci": 208,
|
||||||
|
"utf8_esperanto_ci": 209,
|
||||||
|
"utf8_hungarian_ci": 210,
|
||||||
|
"utf8_sinhala_ci": 211,
|
||||||
|
"utf8_german2_ci": 212,
|
||||||
|
"utf8_croatian_ci": 213,
|
||||||
|
"utf8_unicode_520_ci": 214,
|
||||||
|
"utf8_vietnamese_ci": 215,
|
||||||
|
"utf8_general_mysql500_ci": 223,
|
||||||
|
"utf8mb4_unicode_ci": 224,
|
||||||
|
"utf8mb4_icelandic_ci": 225,
|
||||||
|
"utf8mb4_latvian_ci": 226,
|
||||||
|
"utf8mb4_romanian_ci": 227,
|
||||||
|
"utf8mb4_slovenian_ci": 228,
|
||||||
|
"utf8mb4_polish_ci": 229,
|
||||||
|
"utf8mb4_estonian_ci": 230,
|
||||||
|
"utf8mb4_spanish_ci": 231,
|
||||||
|
"utf8mb4_swedish_ci": 232,
|
||||||
|
"utf8mb4_turkish_ci": 233,
|
||||||
|
"utf8mb4_czech_ci": 234,
|
||||||
|
"utf8mb4_danish_ci": 235,
|
||||||
|
"utf8mb4_lithuanian_ci": 236,
|
||||||
|
"utf8mb4_slovak_ci": 237,
|
||||||
|
"utf8mb4_spanish2_ci": 238,
|
||||||
|
"utf8mb4_roman_ci": 239,
|
||||||
|
"utf8mb4_persian_ci": 240,
|
||||||
|
"utf8mb4_esperanto_ci": 241,
|
||||||
|
"utf8mb4_hungarian_ci": 242,
|
||||||
|
"utf8mb4_sinhala_ci": 243,
|
||||||
|
"utf8mb4_german2_ci": 244,
|
||||||
|
"utf8mb4_croatian_ci": 245,
|
||||||
|
"utf8mb4_unicode_520_ci": 246,
|
||||||
|
"utf8mb4_vietnamese_ci": 247,
|
||||||
|
"gb18030_chinese_ci": 248,
|
||||||
|
"gb18030_bin": 249,
|
||||||
|
"gb18030_unicode_520_ci": 250,
|
||||||
|
"utf8mb4_0900_ai_ci": 255,
|
||||||
|
}
|
||||||
|
|
||||||
|
// A denylist of collations which is unsafe to interpolate parameters.
|
||||||
|
// These multibyte encodings may contains 0x5c (`\`) in their trailing bytes.
|
||||||
|
var unsafeCollations = map[string]bool{
|
||||||
|
"big5_chinese_ci": true,
|
||||||
|
"sjis_japanese_ci": true,
|
||||||
|
"gbk_chinese_ci": true,
|
||||||
|
"big5_bin": true,
|
||||||
|
"gb2312_bin": true,
|
||||||
|
"gbk_bin": true,
|
||||||
|
"sjis_bin": true,
|
||||||
|
"cp932_japanese_ci": true,
|
||||||
|
"cp932_bin": true,
|
||||||
|
"gb18030_chinese_ci": true,
|
||||||
|
"gb18030_bin": true,
|
||||||
|
"gb18030_unicode_520_ci": true,
|
||||||
|
}
|
||||||
+213
@@ -0,0 +1,213 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2024 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"compress/zlib"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
zrPool *sync.Pool // Do not use directly. Use zDecompress() instead.
|
||||||
|
zwPool *sync.Pool // Do not use directly. Use zCompress() instead.
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
zrPool = &sync.Pool{
|
||||||
|
New: func() any { return nil },
|
||||||
|
}
|
||||||
|
zwPool = &sync.Pool{
|
||||||
|
New: func() any {
|
||||||
|
zw, err := zlib.NewWriterLevel(new(bytes.Buffer), 2)
|
||||||
|
if err != nil {
|
||||||
|
panic(err) // compress/zlib return non-nil error only if level is invalid
|
||||||
|
}
|
||||||
|
return zw
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func zDecompress(src []byte, dst *bytes.Buffer) (int, error) {
|
||||||
|
br := bytes.NewReader(src)
|
||||||
|
var zr io.ReadCloser
|
||||||
|
var err error
|
||||||
|
|
||||||
|
if a := zrPool.Get(); a == nil {
|
||||||
|
if zr, err = zlib.NewReader(br); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
zr = a.(io.ReadCloser)
|
||||||
|
if err := zr.(zlib.Resetter).Reset(br, nil); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
n, _ := dst.ReadFrom(zr) // ignore err because zr.Close() will return it again.
|
||||||
|
err = zr.Close() // zr.Close() may return chuecksum error.
|
||||||
|
zrPool.Put(zr)
|
||||||
|
return int(n), err
|
||||||
|
}
|
||||||
|
|
||||||
|
func zCompress(src []byte, dst io.Writer) error {
|
||||||
|
zw := zwPool.Get().(*zlib.Writer)
|
||||||
|
zw.Reset(dst)
|
||||||
|
if _, err := zw.Write(src); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
err := zw.Close()
|
||||||
|
zwPool.Put(zw)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
type compIO struct {
|
||||||
|
mc *mysqlConn
|
||||||
|
buff bytes.Buffer
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCompIO(mc *mysqlConn) *compIO {
|
||||||
|
return &compIO{
|
||||||
|
mc: mc,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *compIO) reset() {
|
||||||
|
c.buff.Reset()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *compIO) readNext(need int) ([]byte, error) {
|
||||||
|
for c.buff.Len() < need {
|
||||||
|
if err := c.readCompressedPacket(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
data := c.buff.Next(need)
|
||||||
|
return data[:need:need], nil // prevent caller writes into c.buff
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *compIO) readCompressedPacket() error {
|
||||||
|
header, err := c.mc.readNext(7)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_ = header[6] // bounds check hint to compiler; guaranteed by readNext
|
||||||
|
|
||||||
|
// compressed header structure
|
||||||
|
comprLength := getUint24(header[0:3])
|
||||||
|
compressionSequence := header[3]
|
||||||
|
uncompressedLength := getUint24(header[4:7])
|
||||||
|
if debug {
|
||||||
|
fmt.Printf("uncompress cmplen=%v uncomplen=%v pkt_cmp_seq=%v expected_cmp_seq=%v\n",
|
||||||
|
comprLength, uncompressedLength, compressionSequence, c.mc.sequence)
|
||||||
|
}
|
||||||
|
// Do not return ErrPktSync here.
|
||||||
|
// Server may return error packet (e.g. 1153 Got a packet bigger than 'max_allowed_packet' bytes)
|
||||||
|
// before receiving all packets from client. In this case, seqnr is younger than expected.
|
||||||
|
// NOTE: Both of mariadbclient and mysqlclient do not check seqnr. Only server checks it.
|
||||||
|
if debug && compressionSequence != c.mc.compressSequence {
|
||||||
|
fmt.Printf("WARN: unexpected cmpress seq nr: expected %v, got %v",
|
||||||
|
c.mc.compressSequence, compressionSequence)
|
||||||
|
}
|
||||||
|
c.mc.compressSequence = compressionSequence + 1
|
||||||
|
|
||||||
|
comprData, err := c.mc.readNext(comprLength)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// if payload is uncompressed, its length will be specified as zero, and its
|
||||||
|
// true length is contained in comprLength
|
||||||
|
if uncompressedLength == 0 {
|
||||||
|
c.buff.Write(comprData)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// use existing capacity in bytesBuf if possible
|
||||||
|
c.buff.Grow(uncompressedLength)
|
||||||
|
nread, err := zDecompress(comprData, &c.buff)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if nread != uncompressedLength {
|
||||||
|
return fmt.Errorf("invalid compressed packet: uncompressed length in header is %d, actual %d",
|
||||||
|
uncompressedLength, nread)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
const minCompressLength = 150
|
||||||
|
const maxPayloadLen = maxPacketSize - 4
|
||||||
|
|
||||||
|
// writePackets sends one or some packets with compression.
|
||||||
|
// Use this instead of mc.netConn.Write() when mc.compress is true.
|
||||||
|
func (c *compIO) writePackets(packets []byte) (int, error) {
|
||||||
|
totalBytes := len(packets)
|
||||||
|
blankHeader := make([]byte, 7)
|
||||||
|
buf := &c.buff
|
||||||
|
|
||||||
|
for len(packets) > 0 {
|
||||||
|
payloadLen := min(maxPayloadLen, len(packets))
|
||||||
|
payload := packets[:payloadLen]
|
||||||
|
uncompressedLen := payloadLen
|
||||||
|
|
||||||
|
buf.Reset()
|
||||||
|
buf.Write(blankHeader) // Buffer.Write() never returns error
|
||||||
|
|
||||||
|
// If payload is less than minCompressLength, don't compress.
|
||||||
|
if uncompressedLen < minCompressLength {
|
||||||
|
buf.Write(payload)
|
||||||
|
uncompressedLen = 0
|
||||||
|
} else {
|
||||||
|
err := zCompress(payload, buf)
|
||||||
|
if debug && err != nil {
|
||||||
|
fmt.Printf("zCompress error: %v", err)
|
||||||
|
}
|
||||||
|
// do not compress if compressed data is larger than uncompressed data
|
||||||
|
// I intentionally miss 7 byte header in the buf; zCompress must compress more than 7 bytes.
|
||||||
|
if err != nil || buf.Len() >= uncompressedLen {
|
||||||
|
buf.Reset()
|
||||||
|
buf.Write(blankHeader)
|
||||||
|
buf.Write(payload)
|
||||||
|
uncompressedLen = 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if n, err := c.writeCompressedPacket(buf.Bytes(), uncompressedLen); err != nil {
|
||||||
|
// To allow returning ErrBadConn when sending really 0 bytes, we sum
|
||||||
|
// up compressed bytes that is returned by underlying Write().
|
||||||
|
return totalBytes - len(packets) + n, err
|
||||||
|
}
|
||||||
|
packets = packets[payloadLen:]
|
||||||
|
}
|
||||||
|
|
||||||
|
return totalBytes, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeCompressedPacket writes a compressed packet with header.
|
||||||
|
// data should start with 7 size space for header followed by payload.
|
||||||
|
func (c *compIO) writeCompressedPacket(data []byte, uncompressedLen int) (int, error) {
|
||||||
|
mc := c.mc
|
||||||
|
comprLength := len(data) - 7
|
||||||
|
if debug {
|
||||||
|
fmt.Printf(
|
||||||
|
"writeCompressedPacket: comprLength=%v, uncompressedLen=%v, seq=%v\n",
|
||||||
|
comprLength, uncompressedLen, mc.compressSequence)
|
||||||
|
}
|
||||||
|
|
||||||
|
// compression header
|
||||||
|
putUint24(data[0:3], comprLength)
|
||||||
|
data[3] = mc.compressSequence
|
||||||
|
putUint24(data[4:7], uncompressedLen)
|
||||||
|
|
||||||
|
mc.compressSequence++
|
||||||
|
return mc.writeWithTimeout(data)
|
||||||
|
}
|
||||||
+55
@@ -0,0 +1,55 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2019 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
//go:build linux || darwin || dragonfly || freebsd || netbsd || openbsd || solaris || illumos
|
||||||
|
// +build linux darwin dragonfly freebsd netbsd openbsd solaris illumos
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"syscall"
|
||||||
|
)
|
||||||
|
|
||||||
|
var errUnexpectedRead = errors.New("unexpected read from socket")
|
||||||
|
|
||||||
|
func connCheck(conn net.Conn) error {
|
||||||
|
var sysErr error
|
||||||
|
|
||||||
|
sysConn, ok := conn.(syscall.Conn)
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
rawConn, err := sysConn.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
err = rawConn.Read(func(fd uintptr) bool {
|
||||||
|
var buf [1]byte
|
||||||
|
n, err := syscall.Read(int(fd), buf[:])
|
||||||
|
switch {
|
||||||
|
case n == 0 && err == nil:
|
||||||
|
sysErr = io.EOF
|
||||||
|
case n > 0:
|
||||||
|
sysErr = errUnexpectedRead
|
||||||
|
case err == syscall.EAGAIN || err == syscall.EWOULDBLOCK:
|
||||||
|
sysErr = nil
|
||||||
|
default:
|
||||||
|
sysErr = err
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return sysErr
|
||||||
|
}
|
||||||
+18
@@ -0,0 +1,18 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2019 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
//go:build !linux && !darwin && !dragonfly && !freebsd && !netbsd && !openbsd && !solaris && !illumos
|
||||||
|
// +build !linux,!darwin,!dragonfly,!freebsd,!netbsd,!openbsd,!solaris,!illumos
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import "net"
|
||||||
|
|
||||||
|
func connCheck(conn net.Conn) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+721
@@ -0,0 +1,721 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2012 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"runtime"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mysqlConn struct {
|
||||||
|
buf buffer
|
||||||
|
netConn net.Conn
|
||||||
|
rawConn net.Conn // underlying connection when netConn is TLS connection.
|
||||||
|
result mysqlResult // managed by clearResult() and handleOkPacket().
|
||||||
|
compIO *compIO
|
||||||
|
cfg *Config
|
||||||
|
connector *connector
|
||||||
|
maxAllowedPacket int
|
||||||
|
maxWriteSize int
|
||||||
|
flags clientFlag
|
||||||
|
status statusFlag
|
||||||
|
sequence uint8
|
||||||
|
compressSequence uint8
|
||||||
|
parseTime bool
|
||||||
|
compress bool
|
||||||
|
|
||||||
|
// for context support (Go 1.8+)
|
||||||
|
watching bool
|
||||||
|
watcher chan<- context.Context
|
||||||
|
closech chan struct{}
|
||||||
|
finished chan<- struct{}
|
||||||
|
canceled atomicError // set non-nil if conn is canceled
|
||||||
|
closed atomic.Bool // set when conn is closed, before closech is closed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to call per-connection logger.
|
||||||
|
func (mc *mysqlConn) log(v ...any) {
|
||||||
|
_, filename, lineno, ok := runtime.Caller(1)
|
||||||
|
if ok {
|
||||||
|
pos := strings.LastIndexByte(filename, '/')
|
||||||
|
if pos != -1 {
|
||||||
|
filename = filename[pos+1:]
|
||||||
|
}
|
||||||
|
prefix := fmt.Sprintf("%s:%d ", filename, lineno)
|
||||||
|
v = append([]any{prefix}, v...)
|
||||||
|
}
|
||||||
|
|
||||||
|
mc.cfg.Logger.Print(v...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) readWithTimeout(b []byte) (int, error) {
|
||||||
|
to := mc.cfg.ReadTimeout
|
||||||
|
if to > 0 {
|
||||||
|
if err := mc.netConn.SetReadDeadline(time.Now().Add(to)); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return mc.netConn.Read(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) writeWithTimeout(b []byte) (int, error) {
|
||||||
|
to := mc.cfg.WriteTimeout
|
||||||
|
if to > 0 {
|
||||||
|
if err := mc.netConn.SetWriteDeadline(time.Now().Add(to)); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return mc.netConn.Write(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) resetSequence() {
|
||||||
|
mc.sequence = 0
|
||||||
|
mc.compressSequence = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// syncSequence must be called when finished writing some packet and before start reading.
|
||||||
|
func (mc *mysqlConn) syncSequence() {
|
||||||
|
// Syncs compressionSequence to sequence.
|
||||||
|
// This is not documented but done in `net_flush()` in MySQL and MariaDB.
|
||||||
|
// https://github.com/mariadb-corporation/mariadb-connector-c/blob/8228164f850b12353da24df1b93a1e53cc5e85e9/libmariadb/ma_net.c#L170-L171
|
||||||
|
// https://github.com/mysql/mysql-server/blob/824e2b4064053f7daf17d7f3f84b7a3ed92e5fb4/sql-common/net_serv.cc#L293
|
||||||
|
if mc.compress {
|
||||||
|
mc.sequence = mc.compressSequence
|
||||||
|
mc.compIO.reset()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handles parameters set in DSN after the connection is established
|
||||||
|
func (mc *mysqlConn) handleParams() (err error) {
|
||||||
|
var cmdSet strings.Builder
|
||||||
|
|
||||||
|
for param, val := range mc.cfg.Params {
|
||||||
|
if cmdSet.Len() == 0 {
|
||||||
|
// Heuristic: 29 chars for each other key=value to reduce reallocations
|
||||||
|
cmdSet.Grow(4 + len(param) + 3 + len(val) + 30*(len(mc.cfg.Params)-1))
|
||||||
|
cmdSet.WriteString("SET ")
|
||||||
|
} else {
|
||||||
|
cmdSet.WriteString(", ")
|
||||||
|
}
|
||||||
|
cmdSet.WriteString(param)
|
||||||
|
cmdSet.WriteString(" = ")
|
||||||
|
cmdSet.WriteString(val)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cmdSet.Len() > 0 {
|
||||||
|
err = mc.exec(cmdSet.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// markBadConn replaces errBadConnNoWrite with driver.ErrBadConn.
|
||||||
|
// This function is used to return driver.ErrBadConn only when safe to retry.
|
||||||
|
func (mc *mysqlConn) markBadConn(err error) error {
|
||||||
|
if err == errBadConnNoWrite {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) Begin() (driver.Tx, error) {
|
||||||
|
return mc.begin(false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) begin(readOnly bool) (driver.Tx, error) {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
var q string
|
||||||
|
if readOnly {
|
||||||
|
q = "START TRANSACTION READ ONLY"
|
||||||
|
} else {
|
||||||
|
q = "START TRANSACTION"
|
||||||
|
}
|
||||||
|
err := mc.exec(q)
|
||||||
|
if err == nil {
|
||||||
|
return &mysqlTx{mc}, err
|
||||||
|
}
|
||||||
|
return nil, mc.markBadConn(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) Close() (err error) {
|
||||||
|
// Makes Close idempotent
|
||||||
|
if !mc.closed.Load() {
|
||||||
|
err = mc.writeCommandPacket(comQuit)
|
||||||
|
}
|
||||||
|
mc.close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// close closes the network connection and clear results without sending COM_QUIT.
|
||||||
|
func (mc *mysqlConn) close() {
|
||||||
|
mc.cleanup()
|
||||||
|
mc.clearResult()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Closes the network connection and unsets internal variables. Do not call this
|
||||||
|
// function after successfully authentication, call Close instead. This function
|
||||||
|
// is called before auth or on auth failure because MySQL will have already
|
||||||
|
// closed the network connection.
|
||||||
|
func (mc *mysqlConn) cleanup() {
|
||||||
|
if mc.closed.Swap(true) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Makes cleanup idempotent
|
||||||
|
close(mc.closech)
|
||||||
|
conn := mc.rawConn
|
||||||
|
if conn == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
mc.log("closing connection:", err)
|
||||||
|
}
|
||||||
|
// This function can be called from multiple goroutines.
|
||||||
|
// So we can not mc.clearResult() here.
|
||||||
|
// Caller should do it if they are in safe goroutine.
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) error() error {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
if err := mc.canceled.Value(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return ErrInvalidConn
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) Prepare(query string) (driver.Stmt, error) {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
// Send command
|
||||||
|
err := mc.writeCommandPacketStr(comStmtPrepare, query)
|
||||||
|
if err != nil {
|
||||||
|
// STMT_PREPARE is safe to retry. So we can return ErrBadConn here.
|
||||||
|
mc.log(err)
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
stmt := &mysqlStmt{
|
||||||
|
mc: mc,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read Result
|
||||||
|
columnCount, err := stmt.readPrepareResultPacket()
|
||||||
|
if err == nil {
|
||||||
|
if stmt.paramCount > 0 {
|
||||||
|
if err = mc.readUntilEOF(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if columnCount > 0 {
|
||||||
|
err = mc.readUntilEOF()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return stmt, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) interpolateParams(query string, args []driver.Value) (string, error) {
|
||||||
|
// Number of ? should be same to len(args)
|
||||||
|
if strings.Count(query, "?") != len(args) {
|
||||||
|
return "", driver.ErrSkip
|
||||||
|
}
|
||||||
|
|
||||||
|
buf, err := mc.buf.takeCompleteBuffer()
|
||||||
|
if err != nil {
|
||||||
|
// can not take the buffer. Something must be wrong with the connection
|
||||||
|
mc.cleanup()
|
||||||
|
// interpolateParams would be called before sending any query.
|
||||||
|
// So its safe to retry.
|
||||||
|
return "", driver.ErrBadConn
|
||||||
|
}
|
||||||
|
buf = buf[:0]
|
||||||
|
argPos := 0
|
||||||
|
|
||||||
|
for i := 0; i < len(query); i++ {
|
||||||
|
q := strings.IndexByte(query[i:], '?')
|
||||||
|
if q == -1 {
|
||||||
|
buf = append(buf, query[i:]...)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
buf = append(buf, query[i:i+q]...)
|
||||||
|
i += q
|
||||||
|
|
||||||
|
arg := args[argPos]
|
||||||
|
argPos++
|
||||||
|
|
||||||
|
if arg == nil {
|
||||||
|
buf = append(buf, "NULL"...)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := arg.(type) {
|
||||||
|
case int64:
|
||||||
|
buf = strconv.AppendInt(buf, v, 10)
|
||||||
|
case uint64:
|
||||||
|
// Handle uint64 explicitly because our custom ConvertValue emits unsigned values
|
||||||
|
buf = strconv.AppendUint(buf, v, 10)
|
||||||
|
case float64:
|
||||||
|
buf = strconv.AppendFloat(buf, v, 'g', -1, 64)
|
||||||
|
case bool:
|
||||||
|
if v {
|
||||||
|
buf = append(buf, '1')
|
||||||
|
} else {
|
||||||
|
buf = append(buf, '0')
|
||||||
|
}
|
||||||
|
case time.Time:
|
||||||
|
if v.IsZero() {
|
||||||
|
buf = append(buf, "'0000-00-00'"...)
|
||||||
|
} else {
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
buf, err = appendDateTime(buf, v.In(mc.cfg.Loc), mc.cfg.timeTruncate)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
}
|
||||||
|
case json.RawMessage:
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
if mc.status&statusNoBackslashEscapes == 0 {
|
||||||
|
buf = escapeBytesBackslash(buf, v)
|
||||||
|
} else {
|
||||||
|
buf = escapeBytesQuotes(buf, v)
|
||||||
|
}
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
case []byte:
|
||||||
|
if v == nil {
|
||||||
|
buf = append(buf, "NULL"...)
|
||||||
|
} else {
|
||||||
|
buf = append(buf, "_binary'"...)
|
||||||
|
if mc.status&statusNoBackslashEscapes == 0 {
|
||||||
|
buf = escapeBytesBackslash(buf, v)
|
||||||
|
} else {
|
||||||
|
buf = escapeBytesQuotes(buf, v)
|
||||||
|
}
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
}
|
||||||
|
case string:
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
if mc.status&statusNoBackslashEscapes == 0 {
|
||||||
|
buf = escapeStringBackslash(buf, v)
|
||||||
|
} else {
|
||||||
|
buf = escapeStringQuotes(buf, v)
|
||||||
|
}
|
||||||
|
buf = append(buf, '\'')
|
||||||
|
default:
|
||||||
|
return "", driver.ErrSkip
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(buf)+4 > mc.maxAllowedPacket {
|
||||||
|
return "", driver.ErrSkip
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if argPos != len(args) {
|
||||||
|
return "", driver.ErrSkip
|
||||||
|
}
|
||||||
|
return string(buf), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) Exec(query string, args []driver.Value) (driver.Result, error) {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
if len(args) != 0 {
|
||||||
|
if !mc.cfg.InterpolateParams {
|
||||||
|
return nil, driver.ErrSkip
|
||||||
|
}
|
||||||
|
// try to interpolate the parameters to save extra roundtrips for preparing and closing a statement
|
||||||
|
prepared, err := mc.interpolateParams(query, args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
query = prepared
|
||||||
|
}
|
||||||
|
|
||||||
|
err := mc.exec(query)
|
||||||
|
if err == nil {
|
||||||
|
copied := mc.result
|
||||||
|
return &copied, err
|
||||||
|
}
|
||||||
|
return nil, mc.markBadConn(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Internal function to execute commands
|
||||||
|
func (mc *mysqlConn) exec(query string) error {
|
||||||
|
handleOk := mc.clearResult()
|
||||||
|
// Send command
|
||||||
|
if err := mc.writeCommandPacketStr(comQuery, query); err != nil {
|
||||||
|
return mc.markBadConn(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read Result
|
||||||
|
resLen, err := handleOk.readResultSetHeaderPacket()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if resLen > 0 {
|
||||||
|
// columns
|
||||||
|
if err := mc.readUntilEOF(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// rows
|
||||||
|
if err := mc.readUntilEOF(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return handleOk.discardResults()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) Query(query string, args []driver.Value) (driver.Rows, error) {
|
||||||
|
return mc.query(query, args)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) query(query string, args []driver.Value) (*textRows, error) {
|
||||||
|
handleOk := mc.clearResult()
|
||||||
|
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
if len(args) != 0 {
|
||||||
|
if !mc.cfg.InterpolateParams {
|
||||||
|
return nil, driver.ErrSkip
|
||||||
|
}
|
||||||
|
// try client-side prepare to reduce roundtrip
|
||||||
|
prepared, err := mc.interpolateParams(query, args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
query = prepared
|
||||||
|
}
|
||||||
|
// Send command
|
||||||
|
err := mc.writeCommandPacketStr(comQuery, query)
|
||||||
|
if err != nil {
|
||||||
|
return nil, mc.markBadConn(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read Result
|
||||||
|
var resLen int
|
||||||
|
resLen, err = handleOk.readResultSetHeaderPacket()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
rows := new(textRows)
|
||||||
|
rows.mc = mc
|
||||||
|
|
||||||
|
if resLen == 0 {
|
||||||
|
rows.rs.done = true
|
||||||
|
|
||||||
|
switch err := rows.NextResultSet(); err {
|
||||||
|
case nil, io.EOF:
|
||||||
|
return rows, nil
|
||||||
|
default:
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Columns
|
||||||
|
rows.rs.columns, err = mc.readColumns(resLen)
|
||||||
|
return rows, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Gets the value of the given MySQL System Variable
|
||||||
|
// The returned byte slice is only valid until the next read
|
||||||
|
func (mc *mysqlConn) getSystemVar(name string) ([]byte, error) {
|
||||||
|
// Send command
|
||||||
|
handleOk := mc.clearResult()
|
||||||
|
if err := mc.writeCommandPacketStr(comQuery, "SELECT @@"+name); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read Result
|
||||||
|
resLen, err := handleOk.readResultSetHeaderPacket()
|
||||||
|
if err == nil {
|
||||||
|
rows := new(textRows)
|
||||||
|
rows.mc = mc
|
||||||
|
rows.rs.columns = []mysqlField{{fieldType: fieldTypeVarChar}}
|
||||||
|
|
||||||
|
if resLen > 0 {
|
||||||
|
// Columns
|
||||||
|
if err := mc.readUntilEOF(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
dest := make([]driver.Value, resLen)
|
||||||
|
if err = rows.readRow(dest); err == nil {
|
||||||
|
return dest[0].([]byte), mc.readUntilEOF()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// cancel is called when the query has canceled.
|
||||||
|
func (mc *mysqlConn) cancel(err error) {
|
||||||
|
mc.canceled.Set(err)
|
||||||
|
mc.cleanup()
|
||||||
|
}
|
||||||
|
|
||||||
|
// finish is called when the query has succeeded.
|
||||||
|
func (mc *mysqlConn) finish() {
|
||||||
|
if !mc.watching || mc.finished == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case mc.finished <- struct{}{}:
|
||||||
|
mc.watching = false
|
||||||
|
case <-mc.closech:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ping implements driver.Pinger interface
|
||||||
|
func (mc *mysqlConn) Ping(ctx context.Context) (err error) {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = mc.watchCancel(ctx); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer mc.finish()
|
||||||
|
|
||||||
|
handleOk := mc.clearResult()
|
||||||
|
if err = mc.writeCommandPacket(comPing); err != nil {
|
||||||
|
return mc.markBadConn(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return handleOk.readResultOK()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeginTx implements driver.ConnBeginTx interface
|
||||||
|
func (mc *mysqlConn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) {
|
||||||
|
if mc.closed.Load() {
|
||||||
|
return nil, driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer mc.finish()
|
||||||
|
|
||||||
|
if sql.IsolationLevel(opts.Isolation) != sql.LevelDefault {
|
||||||
|
level, err := mapIsolationLevel(opts.Isolation)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
err = mc.exec("SET TRANSACTION ISOLATION LEVEL " + level)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return mc.begin(opts.ReadOnly)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
dargs, err := namedValueToValue(args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := mc.query(query, dargs)
|
||||||
|
if err != nil {
|
||||||
|
mc.finish()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
rows.finish = mc.finish
|
||||||
|
return rows, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) {
|
||||||
|
dargs, err := namedValueToValue(args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer mc.finish()
|
||||||
|
|
||||||
|
return mc.Exec(query, dargs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
|
||||||
|
if err := mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
stmt, err := mc.Prepare(query)
|
||||||
|
mc.finish()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
default:
|
||||||
|
case <-ctx.Done():
|
||||||
|
stmt.Close()
|
||||||
|
return nil, ctx.Err()
|
||||||
|
}
|
||||||
|
return stmt, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (stmt *mysqlStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
|
||||||
|
dargs, err := namedValueToValue(args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := stmt.mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := stmt.query(dargs)
|
||||||
|
if err != nil {
|
||||||
|
stmt.mc.finish()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
rows.finish = stmt.mc.finish
|
||||||
|
return rows, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (stmt *mysqlStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
|
||||||
|
dargs, err := namedValueToValue(args)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := stmt.mc.watchCancel(ctx); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer stmt.mc.finish()
|
||||||
|
|
||||||
|
return stmt.Exec(dargs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) watchCancel(ctx context.Context) error {
|
||||||
|
if mc.watching {
|
||||||
|
// Reach here if canceled,
|
||||||
|
// so the connection is already invalid
|
||||||
|
mc.cleanup()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// When ctx is already cancelled, don't watch it.
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// When ctx is not cancellable, don't watch it.
|
||||||
|
if ctx.Done() == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// When watcher is not alive, can't watch it.
|
||||||
|
if mc.watcher == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
mc.watching = true
|
||||||
|
mc.watcher <- ctx
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) startWatcher() {
|
||||||
|
watcher := make(chan context.Context, 1)
|
||||||
|
mc.watcher = watcher
|
||||||
|
finished := make(chan struct{})
|
||||||
|
mc.finished = finished
|
||||||
|
go func() {
|
||||||
|
for {
|
||||||
|
var ctx context.Context
|
||||||
|
select {
|
||||||
|
case ctx = <-watcher:
|
||||||
|
case <-mc.closech:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
mc.cancel(ctx.Err())
|
||||||
|
case <-finished:
|
||||||
|
case <-mc.closech:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mc *mysqlConn) CheckNamedValue(nv *driver.NamedValue) (err error) {
|
||||||
|
nv.Value, err = converter{}.ConvertValue(nv.Value)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetSession implements driver.SessionResetter.
|
||||||
|
// (From Go 1.10)
|
||||||
|
func (mc *mysqlConn) ResetSession(ctx context.Context) error {
|
||||||
|
if mc.closed.Load() || mc.buf.busy() {
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
|
||||||
|
// Perform a stale connection check. We only perform this check for
|
||||||
|
// the first query on a connection that has been checked out of the
|
||||||
|
// connection pool: a fresh connection from the pool is more likely
|
||||||
|
// to be stale, and it has not performed any previous writes that
|
||||||
|
// could cause data corruption, so it's safe to return ErrBadConn
|
||||||
|
// if the check fails.
|
||||||
|
if mc.cfg.CheckConnLiveness {
|
||||||
|
conn := mc.netConn
|
||||||
|
if mc.rawConn != nil {
|
||||||
|
conn = mc.rawConn
|
||||||
|
}
|
||||||
|
var err error
|
||||||
|
if mc.cfg.ReadTimeout != 0 {
|
||||||
|
err = conn.SetReadDeadline(time.Now().Add(mc.cfg.ReadTimeout))
|
||||||
|
}
|
||||||
|
if err == nil {
|
||||||
|
err = connCheck(conn)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
mc.log("closing bad idle connection: ", err)
|
||||||
|
return driver.ErrBadConn
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsValid implements driver.Validator interface
|
||||||
|
// (From Go 1.15)
|
||||||
|
func (mc *mysqlConn) IsValid() bool {
|
||||||
|
return !mc.closed.Load() && !mc.buf.busy()
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ driver.SessionResetter = &mysqlConn{}
|
||||||
|
var _ driver.Validator = &mysqlConn{}
|
||||||
+227
@@ -0,0 +1,227 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2018 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql/driver"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
type connector struct {
|
||||||
|
cfg *Config // immutable private copy.
|
||||||
|
encodedAttributes string // Encoded connection attributes.
|
||||||
|
}
|
||||||
|
|
||||||
|
func encodeConnectionAttributes(cfg *Config) string {
|
||||||
|
connAttrsBuf := make([]byte, 0)
|
||||||
|
|
||||||
|
// default connection attributes
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrClientName)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrClientNameValue)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrOS)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrOSValue)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrPlatform)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrPlatformValue)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrPid)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, strconv.Itoa(os.Getpid()))
|
||||||
|
serverHost, _, _ := net.SplitHostPort(cfg.Addr)
|
||||||
|
if serverHost != "" {
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, connAttrServerHost)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, serverHost)
|
||||||
|
}
|
||||||
|
|
||||||
|
// user-defined connection attributes
|
||||||
|
for _, connAttr := range strings.Split(cfg.ConnectionAttributes, ",") {
|
||||||
|
k, v, found := strings.Cut(connAttr, ":")
|
||||||
|
if !found {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, k)
|
||||||
|
connAttrsBuf = appendLengthEncodedString(connAttrsBuf, v)
|
||||||
|
}
|
||||||
|
|
||||||
|
return string(connAttrsBuf)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newConnector(cfg *Config) *connector {
|
||||||
|
encodedAttributes := encodeConnectionAttributes(cfg)
|
||||||
|
return &connector{
|
||||||
|
cfg: cfg,
|
||||||
|
encodedAttributes: encodedAttributes,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Connect implements driver.Connector interface.
|
||||||
|
// Connect returns a connection to the database.
|
||||||
|
func (c *connector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||||
|
var err error
|
||||||
|
|
||||||
|
// Invoke beforeConnect if present, with a copy of the configuration
|
||||||
|
cfg := c.cfg
|
||||||
|
if c.cfg.beforeConnect != nil {
|
||||||
|
cfg = c.cfg.Clone()
|
||||||
|
err = c.cfg.beforeConnect(ctx, cfg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// New mysqlConn
|
||||||
|
mc := &mysqlConn{
|
||||||
|
maxAllowedPacket: maxPacketSize,
|
||||||
|
maxWriteSize: maxPacketSize - 1,
|
||||||
|
closech: make(chan struct{}),
|
||||||
|
cfg: cfg,
|
||||||
|
connector: c,
|
||||||
|
}
|
||||||
|
mc.parseTime = mc.cfg.ParseTime
|
||||||
|
|
||||||
|
// Connect to Server
|
||||||
|
dctx := ctx
|
||||||
|
if mc.cfg.Timeout > 0 {
|
||||||
|
var cancel context.CancelFunc
|
||||||
|
dctx, cancel = context.WithTimeout(ctx, c.cfg.Timeout)
|
||||||
|
defer cancel()
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.cfg.DialFunc != nil {
|
||||||
|
mc.netConn, err = c.cfg.DialFunc(dctx, mc.cfg.Net, mc.cfg.Addr)
|
||||||
|
} else {
|
||||||
|
dialsLock.RLock()
|
||||||
|
dial, ok := dials[mc.cfg.Net]
|
||||||
|
dialsLock.RUnlock()
|
||||||
|
if ok {
|
||||||
|
mc.netConn, err = dial(dctx, mc.cfg.Addr)
|
||||||
|
} else {
|
||||||
|
nd := net.Dialer{}
|
||||||
|
mc.netConn, err = nd.DialContext(dctx, mc.cfg.Net, mc.cfg.Addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
mc.rawConn = mc.netConn
|
||||||
|
|
||||||
|
// Enable TCP Keepalives on TCP connections
|
||||||
|
if tc, ok := mc.netConn.(*net.TCPConn); ok {
|
||||||
|
if err := tc.SetKeepAlive(true); err != nil {
|
||||||
|
c.cfg.Logger.Print(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call startWatcher for context support (From Go 1.8)
|
||||||
|
mc.startWatcher()
|
||||||
|
if err := mc.watchCancel(ctx); err != nil {
|
||||||
|
mc.cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer mc.finish()
|
||||||
|
|
||||||
|
mc.buf = newBuffer()
|
||||||
|
|
||||||
|
// Reading Handshake Initialization Packet
|
||||||
|
authData, plugin, err := mc.readHandshakePacket()
|
||||||
|
if err != nil {
|
||||||
|
mc.cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if plugin == "" {
|
||||||
|
plugin = defaultAuthPlugin
|
||||||
|
}
|
||||||
|
|
||||||
|
// Send Client Authentication Packet
|
||||||
|
authResp, err := mc.auth(authData, plugin)
|
||||||
|
if err != nil {
|
||||||
|
// try the default auth plugin, if using the requested plugin failed
|
||||||
|
c.cfg.Logger.Print("could not use requested auth plugin '"+plugin+"': ", err.Error())
|
||||||
|
plugin = defaultAuthPlugin
|
||||||
|
authResp, err = mc.auth(authData, plugin)
|
||||||
|
if err != nil {
|
||||||
|
mc.cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err = mc.writeHandshakeResponsePacket(authResp, plugin); err != nil {
|
||||||
|
mc.cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle response to auth packet, switch methods if possible
|
||||||
|
if err = mc.handleAuthResult(authData, plugin); err != nil {
|
||||||
|
// Authentication failed and MySQL has already closed the connection
|
||||||
|
// (https://dev.mysql.com/doc/internals/en/authentication-fails.html).
|
||||||
|
// Do not send COM_QUIT, just cleanup and return the error.
|
||||||
|
mc.cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if mc.cfg.compress && mc.flags&clientCompress == clientCompress {
|
||||||
|
mc.compress = true
|
||||||
|
mc.compIO = newCompIO(mc)
|
||||||
|
}
|
||||||
|
if mc.cfg.MaxAllowedPacket > 0 {
|
||||||
|
mc.maxAllowedPacket = mc.cfg.MaxAllowedPacket
|
||||||
|
} else {
|
||||||
|
// Get max allowed packet size
|
||||||
|
maxap, err := mc.getSystemVar("max_allowed_packet")
|
||||||
|
if err != nil {
|
||||||
|
mc.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(string(maxap))
|
||||||
|
if err != nil {
|
||||||
|
mc.Close()
|
||||||
|
return nil, fmt.Errorf("invalid max_allowed_packet value (%q): %w", maxap, err)
|
||||||
|
}
|
||||||
|
mc.maxAllowedPacket = n - 1
|
||||||
|
}
|
||||||
|
if mc.maxAllowedPacket < maxPacketSize {
|
||||||
|
mc.maxWriteSize = mc.maxAllowedPacket
|
||||||
|
}
|
||||||
|
|
||||||
|
// Charset: character_set_connection, character_set_client, character_set_results
|
||||||
|
if len(mc.cfg.charsets) > 0 {
|
||||||
|
for _, cs := range mc.cfg.charsets {
|
||||||
|
// ignore errors here - a charset may not exist
|
||||||
|
if mc.cfg.Collation != "" {
|
||||||
|
err = mc.exec("SET NAMES " + cs + " COLLATE " + mc.cfg.Collation)
|
||||||
|
} else {
|
||||||
|
err = mc.exec("SET NAMES " + cs)
|
||||||
|
}
|
||||||
|
if err == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
mc.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle DSN Params
|
||||||
|
err = mc.handleParams()
|
||||||
|
if err != nil {
|
||||||
|
mc.Close()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return mc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Driver implements driver.Connector interface.
|
||||||
|
// Driver returns &MySQLDriver{}.
|
||||||
|
func (c *connector) Driver() driver.Driver {
|
||||||
|
return &MySQLDriver{}
|
||||||
|
}
|
||||||
+192
@@ -0,0 +1,192 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2012 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import "runtime"
|
||||||
|
|
||||||
|
const (
|
||||||
|
debug = false // for debugging. Set true only in development.
|
||||||
|
|
||||||
|
defaultAuthPlugin = "mysql_native_password"
|
||||||
|
defaultMaxAllowedPacket = 64 << 20 // 64 MiB. See https://github.com/go-sql-driver/mysql/issues/1355
|
||||||
|
minProtocolVersion = 10
|
||||||
|
maxPacketSize = 1<<24 - 1
|
||||||
|
timeFormat = "2006-01-02 15:04:05.999999"
|
||||||
|
|
||||||
|
// Connection attributes
|
||||||
|
// See https://dev.mysql.com/doc/refman/8.0/en/performance-schema-connection-attribute-tables.html#performance-schema-connection-attributes-available
|
||||||
|
connAttrClientName = "_client_name"
|
||||||
|
connAttrClientNameValue = "Go-MySQL-Driver"
|
||||||
|
connAttrOS = "_os"
|
||||||
|
connAttrOSValue = runtime.GOOS
|
||||||
|
connAttrPlatform = "_platform"
|
||||||
|
connAttrPlatformValue = runtime.GOARCH
|
||||||
|
connAttrPid = "_pid"
|
||||||
|
connAttrServerHost = "_server_host"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MySQL constants documentation:
|
||||||
|
// http://dev.mysql.com/doc/internals/en/client-server-protocol.html
|
||||||
|
|
||||||
|
const (
|
||||||
|
iOK byte = 0x00
|
||||||
|
iAuthMoreData byte = 0x01
|
||||||
|
iLocalInFile byte = 0xfb
|
||||||
|
iEOF byte = 0xfe
|
||||||
|
iERR byte = 0xff
|
||||||
|
)
|
||||||
|
|
||||||
|
// https://dev.mysql.com/doc/internals/en/capability-flags.html#packet-Protocol::CapabilityFlags
|
||||||
|
type clientFlag uint32
|
||||||
|
|
||||||
|
const (
|
||||||
|
clientLongPassword clientFlag = 1 << iota
|
||||||
|
clientFoundRows
|
||||||
|
clientLongFlag
|
||||||
|
clientConnectWithDB
|
||||||
|
clientNoSchema
|
||||||
|
clientCompress
|
||||||
|
clientODBC
|
||||||
|
clientLocalFiles
|
||||||
|
clientIgnoreSpace
|
||||||
|
clientProtocol41
|
||||||
|
clientInteractive
|
||||||
|
clientSSL
|
||||||
|
clientIgnoreSIGPIPE
|
||||||
|
clientTransactions
|
||||||
|
clientReserved
|
||||||
|
clientSecureConn
|
||||||
|
clientMultiStatements
|
||||||
|
clientMultiResults
|
||||||
|
clientPSMultiResults
|
||||||
|
clientPluginAuth
|
||||||
|
clientConnectAttrs
|
||||||
|
clientPluginAuthLenEncClientData
|
||||||
|
clientCanHandleExpiredPasswords
|
||||||
|
clientSessionTrack
|
||||||
|
clientDeprecateEOF
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
comQuit byte = iota + 1
|
||||||
|
comInitDB
|
||||||
|
comQuery
|
||||||
|
comFieldList
|
||||||
|
comCreateDB
|
||||||
|
comDropDB
|
||||||
|
comRefresh
|
||||||
|
comShutdown
|
||||||
|
comStatistics
|
||||||
|
comProcessInfo
|
||||||
|
comConnect
|
||||||
|
comProcessKill
|
||||||
|
comDebug
|
||||||
|
comPing
|
||||||
|
comTime
|
||||||
|
comDelayedInsert
|
||||||
|
comChangeUser
|
||||||
|
comBinlogDump
|
||||||
|
comTableDump
|
||||||
|
comConnectOut
|
||||||
|
comRegisterSlave
|
||||||
|
comStmtPrepare
|
||||||
|
comStmtExecute
|
||||||
|
comStmtSendLongData
|
||||||
|
comStmtClose
|
||||||
|
comStmtReset
|
||||||
|
comSetOption
|
||||||
|
comStmtFetch
|
||||||
|
)
|
||||||
|
|
||||||
|
// https://dev.mysql.com/doc/internals/en/com-query-response.html#packet-Protocol::ColumnType
|
||||||
|
type fieldType byte
|
||||||
|
|
||||||
|
const (
|
||||||
|
fieldTypeDecimal fieldType = iota
|
||||||
|
fieldTypeTiny
|
||||||
|
fieldTypeShort
|
||||||
|
fieldTypeLong
|
||||||
|
fieldTypeFloat
|
||||||
|
fieldTypeDouble
|
||||||
|
fieldTypeNULL
|
||||||
|
fieldTypeTimestamp
|
||||||
|
fieldTypeLongLong
|
||||||
|
fieldTypeInt24
|
||||||
|
fieldTypeDate
|
||||||
|
fieldTypeTime
|
||||||
|
fieldTypeDateTime
|
||||||
|
fieldTypeYear
|
||||||
|
fieldTypeNewDate
|
||||||
|
fieldTypeVarChar
|
||||||
|
fieldTypeBit
|
||||||
|
)
|
||||||
|
const (
|
||||||
|
fieldTypeVector fieldType = iota + 0xf2
|
||||||
|
fieldTypeInvalid
|
||||||
|
fieldTypeBool
|
||||||
|
fieldTypeJSON
|
||||||
|
fieldTypeNewDecimal
|
||||||
|
fieldTypeEnum
|
||||||
|
fieldTypeSet
|
||||||
|
fieldTypeTinyBLOB
|
||||||
|
fieldTypeMediumBLOB
|
||||||
|
fieldTypeLongBLOB
|
||||||
|
fieldTypeBLOB
|
||||||
|
fieldTypeVarString
|
||||||
|
fieldTypeString
|
||||||
|
fieldTypeGeometry
|
||||||
|
)
|
||||||
|
|
||||||
|
type fieldFlag uint16
|
||||||
|
|
||||||
|
const (
|
||||||
|
flagNotNULL fieldFlag = 1 << iota
|
||||||
|
flagPriKey
|
||||||
|
flagUniqueKey
|
||||||
|
flagMultipleKey
|
||||||
|
flagBLOB
|
||||||
|
flagUnsigned
|
||||||
|
flagZeroFill
|
||||||
|
flagBinary
|
||||||
|
flagEnum
|
||||||
|
flagAutoIncrement
|
||||||
|
flagTimestamp
|
||||||
|
flagSet
|
||||||
|
flagUnknown1
|
||||||
|
flagUnknown2
|
||||||
|
flagUnknown3
|
||||||
|
flagUnknown4
|
||||||
|
)
|
||||||
|
|
||||||
|
// http://dev.mysql.com/doc/internals/en/status-flags.html
|
||||||
|
type statusFlag uint16
|
||||||
|
|
||||||
|
const (
|
||||||
|
statusInTrans statusFlag = 1 << iota
|
||||||
|
statusInAutocommit
|
||||||
|
statusReserved // Not in documentation
|
||||||
|
statusMoreResultsExists
|
||||||
|
statusNoGoodIndexUsed
|
||||||
|
statusNoIndexUsed
|
||||||
|
statusCursorExists
|
||||||
|
statusLastRowSent
|
||||||
|
statusDbDropped
|
||||||
|
statusNoBackslashEscapes
|
||||||
|
statusMetadataChanged
|
||||||
|
statusQueryWasSlow
|
||||||
|
statusPsOutParams
|
||||||
|
statusInTransReadonly
|
||||||
|
statusSessionStateChanged
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
cachingSha2PasswordRequestPublicKey = 2
|
||||||
|
cachingSha2PasswordFastAuthSuccess = 3
|
||||||
|
cachingSha2PasswordPerformFullAuthentication = 4
|
||||||
|
)
|
||||||
+118
@@ -0,0 +1,118 @@
|
|||||||
|
// Copyright 2012 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
// Package mysql provides a MySQL driver for Go's database/sql package.
|
||||||
|
//
|
||||||
|
// The driver should be used via the database/sql package:
|
||||||
|
//
|
||||||
|
// import "database/sql"
|
||||||
|
// import _ "github.com/go-sql-driver/mysql"
|
||||||
|
//
|
||||||
|
// db, err := sql.Open("mysql", "user:password@/dbname")
|
||||||
|
//
|
||||||
|
// See https://github.com/go-sql-driver/mysql#usage for details
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MySQLDriver is exported to make the driver directly accessible.
|
||||||
|
// In general the driver is used via the database/sql package.
|
||||||
|
type MySQLDriver struct{}
|
||||||
|
|
||||||
|
// DialFunc is a function which can be used to establish the network connection.
|
||||||
|
// Custom dial functions must be registered with RegisterDial
|
||||||
|
//
|
||||||
|
// Deprecated: users should register a DialContextFunc instead
|
||||||
|
type DialFunc func(addr string) (net.Conn, error)
|
||||||
|
|
||||||
|
// DialContextFunc is a function which can be used to establish the network connection.
|
||||||
|
// Custom dial functions must be registered with RegisterDialContext
|
||||||
|
type DialContextFunc func(ctx context.Context, addr string) (net.Conn, error)
|
||||||
|
|
||||||
|
var (
|
||||||
|
dialsLock sync.RWMutex
|
||||||
|
dials map[string]DialContextFunc
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterDialContext registers a custom dial function. It can then be used by the
|
||||||
|
// network address mynet(addr), where mynet is the registered new network.
|
||||||
|
// The current context for the connection and its address is passed to the dial function.
|
||||||
|
func RegisterDialContext(net string, dial DialContextFunc) {
|
||||||
|
dialsLock.Lock()
|
||||||
|
defer dialsLock.Unlock()
|
||||||
|
if dials == nil {
|
||||||
|
dials = make(map[string]DialContextFunc)
|
||||||
|
}
|
||||||
|
dials[net] = dial
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeregisterDialContext removes the custom dial function registered with the given net.
|
||||||
|
func DeregisterDialContext(net string) {
|
||||||
|
dialsLock.Lock()
|
||||||
|
defer dialsLock.Unlock()
|
||||||
|
if dials != nil {
|
||||||
|
delete(dials, net)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterDial registers a custom dial function. It can then be used by the
|
||||||
|
// network address mynet(addr), where mynet is the registered new network.
|
||||||
|
// addr is passed as a parameter to the dial function.
|
||||||
|
//
|
||||||
|
// Deprecated: users should call RegisterDialContext instead
|
||||||
|
func RegisterDial(network string, dial DialFunc) {
|
||||||
|
RegisterDialContext(network, func(_ context.Context, addr string) (net.Conn, error) {
|
||||||
|
return dial(addr)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Open new Connection.
|
||||||
|
// See https://github.com/go-sql-driver/mysql#dsn-data-source-name for how
|
||||||
|
// the DSN string is formatted
|
||||||
|
func (d MySQLDriver) Open(dsn string) (driver.Conn, error) {
|
||||||
|
cfg, err := ParseDSN(dsn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c := newConnector(cfg)
|
||||||
|
return c.Connect(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
// This variable can be replaced with -ldflags like below:
|
||||||
|
// go build "-ldflags=-X github.com/go-sql-driver/mysql.driverName=custom"
|
||||||
|
var driverName = "mysql"
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
if driverName != "" {
|
||||||
|
sql.Register(driverName, &MySQLDriver{})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewConnector returns new driver.Connector.
|
||||||
|
func NewConnector(cfg *Config) (driver.Connector, error) {
|
||||||
|
cfg = cfg.Clone()
|
||||||
|
// normalize the contents of cfg so calls to NewConnector have the same
|
||||||
|
// behavior as MySQLDriver.OpenConnector
|
||||||
|
if err := cfg.normalize(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return newConnector(cfg), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenConnector implements driver.DriverContext.
|
||||||
|
func (d MySQLDriver) OpenConnector(dsn string) (driver.Connector, error) {
|
||||||
|
cfg, err := ParseDSN(dsn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return newConnector(cfg), nil
|
||||||
|
}
|
||||||
+701
@@ -0,0 +1,701 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2016 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/rsa"
|
||||||
|
"crypto/tls"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"math/big"
|
||||||
|
"net"
|
||||||
|
"net/url"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errInvalidDSNUnescaped = errors.New("invalid DSN: did you forget to escape a param value?")
|
||||||
|
errInvalidDSNAddr = errors.New("invalid DSN: network address not terminated (missing closing brace)")
|
||||||
|
errInvalidDSNNoSlash = errors.New("invalid DSN: missing the slash separating the database name")
|
||||||
|
errInvalidDSNUnsafeCollation = errors.New("invalid DSN: interpolateParams can not be used with unsafe collations")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Config is a configuration parsed from a DSN string.
|
||||||
|
// If a new Config is created instead of being parsed from a DSN string,
|
||||||
|
// the NewConfig function should be used, which sets default values.
|
||||||
|
type Config struct {
|
||||||
|
// non boolean fields
|
||||||
|
|
||||||
|
User string // Username
|
||||||
|
Passwd string // Password (requires User)
|
||||||
|
Net string // Network (e.g. "tcp", "tcp6", "unix". default: "tcp")
|
||||||
|
Addr string // Address (default: "127.0.0.1:3306" for "tcp" and "/tmp/mysql.sock" for "unix")
|
||||||
|
DBName string // Database name
|
||||||
|
Params map[string]string // Connection parameters
|
||||||
|
ConnectionAttributes string // Connection Attributes, comma-delimited string of user-defined "key:value" pairs
|
||||||
|
Collation string // Connection collation. When set, this will be set in SET NAMES <charset> COLLATE <collation> query
|
||||||
|
Loc *time.Location // Location for time.Time values
|
||||||
|
MaxAllowedPacket int // Max packet size allowed
|
||||||
|
ServerPubKey string // Server public key name
|
||||||
|
TLSConfig string // TLS configuration name
|
||||||
|
TLS *tls.Config // TLS configuration, its priority is higher than TLSConfig
|
||||||
|
Timeout time.Duration // Dial timeout
|
||||||
|
ReadTimeout time.Duration // I/O read timeout
|
||||||
|
WriteTimeout time.Duration // I/O write timeout
|
||||||
|
Logger Logger // Logger
|
||||||
|
// DialFunc specifies the dial function for creating connections
|
||||||
|
DialFunc func(ctx context.Context, network, addr string) (net.Conn, error)
|
||||||
|
|
||||||
|
// boolean fields
|
||||||
|
|
||||||
|
AllowAllFiles bool // Allow all files to be used with LOAD DATA LOCAL INFILE
|
||||||
|
AllowCleartextPasswords bool // Allows the cleartext client side plugin
|
||||||
|
AllowFallbackToPlaintext bool // Allows fallback to unencrypted connection if server does not support TLS
|
||||||
|
AllowNativePasswords bool // Allows the native password authentication method
|
||||||
|
AllowOldPasswords bool // Allows the old insecure password method
|
||||||
|
CheckConnLiveness bool // Check connections for liveness before using them
|
||||||
|
ClientFoundRows bool // Return number of matching rows instead of rows changed
|
||||||
|
ColumnsWithAlias bool // Prepend table alias to column names
|
||||||
|
InterpolateParams bool // Interpolate placeholders into query string
|
||||||
|
MultiStatements bool // Allow multiple statements in one query
|
||||||
|
ParseTime bool // Parse time values to time.Time
|
||||||
|
RejectReadOnly bool // Reject read-only connections
|
||||||
|
|
||||||
|
// unexported fields. new options should be come here.
|
||||||
|
// boolean first. alphabetical order.
|
||||||
|
|
||||||
|
compress bool // Enable zlib compression
|
||||||
|
|
||||||
|
beforeConnect func(context.Context, *Config) error // Invoked before a connection is established
|
||||||
|
pubKey *rsa.PublicKey // Server public key
|
||||||
|
timeTruncate time.Duration // Truncate time.Time values to the specified duration
|
||||||
|
charsets []string // Connection charset. When set, this will be set in SET NAMES <charset> query
|
||||||
|
}
|
||||||
|
|
||||||
|
// Functional Options Pattern
|
||||||
|
// https://dave.cheney.net/2014/10/17/functional-options-for-friendly-apis
|
||||||
|
type Option func(*Config) error
|
||||||
|
|
||||||
|
// NewConfig creates a new Config and sets default values.
|
||||||
|
func NewConfig() *Config {
|
||||||
|
cfg := &Config{
|
||||||
|
Loc: time.UTC,
|
||||||
|
MaxAllowedPacket: defaultMaxAllowedPacket,
|
||||||
|
Logger: defaultLogger,
|
||||||
|
AllowNativePasswords: true,
|
||||||
|
CheckConnLiveness: true,
|
||||||
|
}
|
||||||
|
return cfg
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply applies the given options to the Config object.
|
||||||
|
func (c *Config) Apply(opts ...Option) error {
|
||||||
|
for _, opt := range opts {
|
||||||
|
err := opt(c)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TimeTruncate sets the time duration to truncate time.Time values in
|
||||||
|
// query parameters.
|
||||||
|
func TimeTruncate(d time.Duration) Option {
|
||||||
|
return func(cfg *Config) error {
|
||||||
|
cfg.timeTruncate = d
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BeforeConnect sets the function to be invoked before a connection is established.
|
||||||
|
func BeforeConnect(fn func(context.Context, *Config) error) Option {
|
||||||
|
return func(cfg *Config) error {
|
||||||
|
cfg.beforeConnect = fn
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// EnableCompress sets the compression mode.
|
||||||
|
func EnableCompression(yes bool) Option {
|
||||||
|
return func(cfg *Config) error {
|
||||||
|
cfg.compress = yes
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Charset sets the connection charset and collation.
|
||||||
|
//
|
||||||
|
// charset is the connection charset.
|
||||||
|
// collation is the connection collation. It can be null or empty string.
|
||||||
|
//
|
||||||
|
// When collation is not specified, `SET NAMES <charset>` command is sent when the connection is established.
|
||||||
|
// When collation is specified, `SET NAMES <charset> COLLATE <collation>` command is sent when the connection is established.
|
||||||
|
func Charset(charset, collation string) Option {
|
||||||
|
return func(cfg *Config) error {
|
||||||
|
cfg.charsets = []string{charset}
|
||||||
|
cfg.Collation = collation
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cfg *Config) Clone() *Config {
|
||||||
|
cp := *cfg
|
||||||
|
if cp.TLS != nil {
|
||||||
|
cp.TLS = cfg.TLS.Clone()
|
||||||
|
}
|
||||||
|
if len(cp.Params) > 0 {
|
||||||
|
cp.Params = make(map[string]string, len(cfg.Params))
|
||||||
|
for k, v := range cfg.Params {
|
||||||
|
cp.Params[k] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if cfg.pubKey != nil {
|
||||||
|
cp.pubKey = &rsa.PublicKey{
|
||||||
|
N: new(big.Int).Set(cfg.pubKey.N),
|
||||||
|
E: cfg.pubKey.E,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &cp
|
||||||
|
}
|
||||||
|
|
||||||
|
func (cfg *Config) normalize() error {
|
||||||
|
if cfg.InterpolateParams && cfg.Collation != "" && unsafeCollations[cfg.Collation] {
|
||||||
|
return errInvalidDSNUnsafeCollation
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set default network if empty
|
||||||
|
if cfg.Net == "" {
|
||||||
|
cfg.Net = "tcp"
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set default address if empty
|
||||||
|
if cfg.Addr == "" {
|
||||||
|
switch cfg.Net {
|
||||||
|
case "tcp":
|
||||||
|
cfg.Addr = "127.0.0.1:3306"
|
||||||
|
case "unix":
|
||||||
|
cfg.Addr = "/tmp/mysql.sock"
|
||||||
|
default:
|
||||||
|
return errors.New("default addr for network '" + cfg.Net + "' unknown")
|
||||||
|
}
|
||||||
|
} else if cfg.Net == "tcp" {
|
||||||
|
cfg.Addr = ensureHavePort(cfg.Addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.TLS == nil {
|
||||||
|
switch cfg.TLSConfig {
|
||||||
|
case "false", "":
|
||||||
|
// don't set anything
|
||||||
|
case "true":
|
||||||
|
cfg.TLS = &tls.Config{}
|
||||||
|
case "skip-verify":
|
||||||
|
cfg.TLS = &tls.Config{InsecureSkipVerify: true}
|
||||||
|
case "preferred":
|
||||||
|
cfg.TLS = &tls.Config{InsecureSkipVerify: true}
|
||||||
|
cfg.AllowFallbackToPlaintext = true
|
||||||
|
default:
|
||||||
|
cfg.TLS = getTLSConfigClone(cfg.TLSConfig)
|
||||||
|
if cfg.TLS == nil {
|
||||||
|
return errors.New("invalid value / unknown config name: " + cfg.TLSConfig)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.TLS != nil && cfg.TLS.ServerName == "" && !cfg.TLS.InsecureSkipVerify {
|
||||||
|
host, _, err := net.SplitHostPort(cfg.Addr)
|
||||||
|
if err == nil {
|
||||||
|
cfg.TLS.ServerName = host
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ServerPubKey != "" {
|
||||||
|
cfg.pubKey = getServerPubKey(cfg.ServerPubKey)
|
||||||
|
if cfg.pubKey == nil {
|
||||||
|
return errors.New("invalid value / unknown server pub key name: " + cfg.ServerPubKey)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Logger == nil {
|
||||||
|
cfg.Logger = defaultLogger
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeDSNParam(buf *bytes.Buffer, hasParam *bool, name, value string) {
|
||||||
|
buf.Grow(1 + len(name) + 1 + len(value))
|
||||||
|
if !*hasParam {
|
||||||
|
*hasParam = true
|
||||||
|
buf.WriteByte('?')
|
||||||
|
} else {
|
||||||
|
buf.WriteByte('&')
|
||||||
|
}
|
||||||
|
buf.WriteString(name)
|
||||||
|
buf.WriteByte('=')
|
||||||
|
buf.WriteString(value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatDSN formats the given Config into a DSN string which can be passed to
|
||||||
|
// the driver.
|
||||||
|
//
|
||||||
|
// Note: use [NewConnector] and [database/sql.OpenDB] to open a connection from a [*Config].
|
||||||
|
func (cfg *Config) FormatDSN() string {
|
||||||
|
var buf bytes.Buffer
|
||||||
|
|
||||||
|
// [username[:password]@]
|
||||||
|
if len(cfg.User) > 0 {
|
||||||
|
buf.WriteString(cfg.User)
|
||||||
|
if len(cfg.Passwd) > 0 {
|
||||||
|
buf.WriteByte(':')
|
||||||
|
buf.WriteString(cfg.Passwd)
|
||||||
|
}
|
||||||
|
buf.WriteByte('@')
|
||||||
|
}
|
||||||
|
|
||||||
|
// [protocol[(address)]]
|
||||||
|
if len(cfg.Net) > 0 {
|
||||||
|
buf.WriteString(cfg.Net)
|
||||||
|
if len(cfg.Addr) > 0 {
|
||||||
|
buf.WriteByte('(')
|
||||||
|
buf.WriteString(cfg.Addr)
|
||||||
|
buf.WriteByte(')')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// /dbname
|
||||||
|
buf.WriteByte('/')
|
||||||
|
buf.WriteString(url.PathEscape(cfg.DBName))
|
||||||
|
|
||||||
|
// [?param1=value1&...¶mN=valueN]
|
||||||
|
hasParam := false
|
||||||
|
|
||||||
|
if cfg.AllowAllFiles {
|
||||||
|
hasParam = true
|
||||||
|
buf.WriteString("?allowAllFiles=true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.AllowCleartextPasswords {
|
||||||
|
writeDSNParam(&buf, &hasParam, "allowCleartextPasswords", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.AllowFallbackToPlaintext {
|
||||||
|
writeDSNParam(&buf, &hasParam, "allowFallbackToPlaintext", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !cfg.AllowNativePasswords {
|
||||||
|
writeDSNParam(&buf, &hasParam, "allowNativePasswords", "false")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.AllowOldPasswords {
|
||||||
|
writeDSNParam(&buf, &hasParam, "allowOldPasswords", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !cfg.CheckConnLiveness {
|
||||||
|
writeDSNParam(&buf, &hasParam, "checkConnLiveness", "false")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ClientFoundRows {
|
||||||
|
writeDSNParam(&buf, &hasParam, "clientFoundRows", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if charsets := cfg.charsets; len(charsets) > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "charset", strings.Join(charsets, ","))
|
||||||
|
}
|
||||||
|
|
||||||
|
if col := cfg.Collation; col != "" {
|
||||||
|
writeDSNParam(&buf, &hasParam, "collation", col)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ColumnsWithAlias {
|
||||||
|
writeDSNParam(&buf, &hasParam, "columnsWithAlias", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ConnectionAttributes != "" {
|
||||||
|
writeDSNParam(&buf, &hasParam, "connectionAttributes", url.QueryEscape(cfg.ConnectionAttributes))
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.compress {
|
||||||
|
writeDSNParam(&buf, &hasParam, "compress", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.InterpolateParams {
|
||||||
|
writeDSNParam(&buf, &hasParam, "interpolateParams", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Loc != time.UTC && cfg.Loc != nil {
|
||||||
|
writeDSNParam(&buf, &hasParam, "loc", url.QueryEscape(cfg.Loc.String()))
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.MultiStatements {
|
||||||
|
writeDSNParam(&buf, &hasParam, "multiStatements", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ParseTime {
|
||||||
|
writeDSNParam(&buf, &hasParam, "parseTime", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.timeTruncate > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "timeTruncate", cfg.timeTruncate.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.ReadTimeout > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "readTimeout", cfg.ReadTimeout.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.RejectReadOnly {
|
||||||
|
writeDSNParam(&buf, &hasParam, "rejectReadOnly", "true")
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.ServerPubKey) > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "serverPubKey", url.QueryEscape(cfg.ServerPubKey))
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Timeout > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "timeout", cfg.Timeout.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cfg.TLSConfig) > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "tls", url.QueryEscape(cfg.TLSConfig))
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.WriteTimeout > 0 {
|
||||||
|
writeDSNParam(&buf, &hasParam, "writeTimeout", cfg.WriteTimeout.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.MaxAllowedPacket != defaultMaxAllowedPacket {
|
||||||
|
writeDSNParam(&buf, &hasParam, "maxAllowedPacket", strconv.Itoa(cfg.MaxAllowedPacket))
|
||||||
|
}
|
||||||
|
|
||||||
|
// other params
|
||||||
|
if cfg.Params != nil {
|
||||||
|
var params []string
|
||||||
|
for param := range cfg.Params {
|
||||||
|
params = append(params, param)
|
||||||
|
}
|
||||||
|
sort.Strings(params)
|
||||||
|
for _, param := range params {
|
||||||
|
writeDSNParam(&buf, &hasParam, param, url.QueryEscape(cfg.Params[param]))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return buf.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseDSN parses the DSN string to a Config
|
||||||
|
func ParseDSN(dsn string) (cfg *Config, err error) {
|
||||||
|
// New config with some default values
|
||||||
|
cfg = NewConfig()
|
||||||
|
|
||||||
|
// [user[:password]@][net[(addr)]]/dbname[?param1=value1¶mN=valueN]
|
||||||
|
// Find the last '/' (since the password or the net addr might contain a '/')
|
||||||
|
foundSlash := false
|
||||||
|
for i := len(dsn) - 1; i >= 0; i-- {
|
||||||
|
if dsn[i] == '/' {
|
||||||
|
foundSlash = true
|
||||||
|
var j, k int
|
||||||
|
|
||||||
|
// left part is empty if i <= 0
|
||||||
|
if i > 0 {
|
||||||
|
// [username[:password]@][protocol[(address)]]
|
||||||
|
// Find the last '@' in dsn[:i]
|
||||||
|
for j = i; j >= 0; j-- {
|
||||||
|
if dsn[j] == '@' {
|
||||||
|
// username[:password]
|
||||||
|
// Find the first ':' in dsn[:j]
|
||||||
|
for k = 0; k < j; k++ {
|
||||||
|
if dsn[k] == ':' {
|
||||||
|
cfg.Passwd = dsn[k+1 : j]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cfg.User = dsn[:k]
|
||||||
|
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// [protocol[(address)]]
|
||||||
|
// Find the first '(' in dsn[j+1:i]
|
||||||
|
for k = j + 1; k < i; k++ {
|
||||||
|
if dsn[k] == '(' {
|
||||||
|
// dsn[i-1] must be == ')' if an address is specified
|
||||||
|
if dsn[i-1] != ')' {
|
||||||
|
if strings.ContainsRune(dsn[k+1:i], ')') {
|
||||||
|
return nil, errInvalidDSNUnescaped
|
||||||
|
}
|
||||||
|
return nil, errInvalidDSNAddr
|
||||||
|
}
|
||||||
|
cfg.Addr = dsn[k+1 : i-1]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cfg.Net = dsn[j+1 : k]
|
||||||
|
}
|
||||||
|
|
||||||
|
// dbname[?param1=value1&...¶mN=valueN]
|
||||||
|
// Find the first '?' in dsn[i+1:]
|
||||||
|
for j = i + 1; j < len(dsn); j++ {
|
||||||
|
if dsn[j] == '?' {
|
||||||
|
if err = parseDSNParams(cfg, dsn[j+1:]); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
dbname := dsn[i+1 : j]
|
||||||
|
if cfg.DBName, err = url.PathUnescape(dbname); err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid dbname %q: %w", dbname, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !foundSlash && len(dsn) > 0 {
|
||||||
|
return nil, errInvalidDSNNoSlash
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = cfg.normalize(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseDSNParams parses the DSN "query string"
|
||||||
|
// Values must be url.QueryEscape'ed
|
||||||
|
func parseDSNParams(cfg *Config, params string) (err error) {
|
||||||
|
for _, v := range strings.Split(params, "&") {
|
||||||
|
key, value, found := strings.Cut(v, "=")
|
||||||
|
if !found {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// cfg params
|
||||||
|
switch key {
|
||||||
|
// Disable INFILE allowlist / enable all files
|
||||||
|
case "allowAllFiles":
|
||||||
|
var isBool bool
|
||||||
|
cfg.AllowAllFiles, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use cleartext authentication mode (MySQL 5.5.10+)
|
||||||
|
case "allowCleartextPasswords":
|
||||||
|
var isBool bool
|
||||||
|
cfg.AllowCleartextPasswords, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow fallback to unencrypted connection if server does not support TLS
|
||||||
|
case "allowFallbackToPlaintext":
|
||||||
|
var isBool bool
|
||||||
|
cfg.AllowFallbackToPlaintext, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use native password authentication
|
||||||
|
case "allowNativePasswords":
|
||||||
|
var isBool bool
|
||||||
|
cfg.AllowNativePasswords, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use old authentication mode (pre MySQL 4.1)
|
||||||
|
case "allowOldPasswords":
|
||||||
|
var isBool bool
|
||||||
|
cfg.AllowOldPasswords, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check connections for Liveness before using them
|
||||||
|
case "checkConnLiveness":
|
||||||
|
var isBool bool
|
||||||
|
cfg.CheckConnLiveness, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Switch "rowsAffected" mode
|
||||||
|
case "clientFoundRows":
|
||||||
|
var isBool bool
|
||||||
|
cfg.ClientFoundRows, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// charset
|
||||||
|
case "charset":
|
||||||
|
cfg.charsets = strings.Split(value, ",")
|
||||||
|
|
||||||
|
// Collation
|
||||||
|
case "collation":
|
||||||
|
cfg.Collation = value
|
||||||
|
|
||||||
|
case "columnsWithAlias":
|
||||||
|
var isBool bool
|
||||||
|
cfg.ColumnsWithAlias, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Compression
|
||||||
|
case "compress":
|
||||||
|
var isBool bool
|
||||||
|
cfg.compress, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable client side placeholder substitution
|
||||||
|
case "interpolateParams":
|
||||||
|
var isBool bool
|
||||||
|
cfg.InterpolateParams, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Time Location
|
||||||
|
case "loc":
|
||||||
|
if value, err = url.QueryUnescape(value); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg.Loc, err = time.LoadLocation(value)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// multiple statements in one query
|
||||||
|
case "multiStatements":
|
||||||
|
var isBool bool
|
||||||
|
cfg.MultiStatements, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// time.Time parsing
|
||||||
|
case "parseTime":
|
||||||
|
var isBool bool
|
||||||
|
cfg.ParseTime, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// time.Time truncation
|
||||||
|
case "timeTruncate":
|
||||||
|
cfg.timeTruncate, err = time.ParseDuration(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid timeTruncate value: %v, error: %w", value, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// I/O read Timeout
|
||||||
|
case "readTimeout":
|
||||||
|
cfg.ReadTimeout, err = time.ParseDuration(value)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reject read-only connections
|
||||||
|
case "rejectReadOnly":
|
||||||
|
var isBool bool
|
||||||
|
cfg.RejectReadOnly, isBool = readBool(value)
|
||||||
|
if !isBool {
|
||||||
|
return errors.New("invalid bool value: " + value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server public key
|
||||||
|
case "serverPubKey":
|
||||||
|
name, err := url.QueryUnescape(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid value for server pub key name: %v", err)
|
||||||
|
}
|
||||||
|
cfg.ServerPubKey = name
|
||||||
|
|
||||||
|
// Strict mode
|
||||||
|
case "strict":
|
||||||
|
panic("strict mode has been removed. See https://github.com/go-sql-driver/mysql/wiki/strict-mode")
|
||||||
|
|
||||||
|
// Dial Timeout
|
||||||
|
case "timeout":
|
||||||
|
cfg.Timeout, err = time.ParseDuration(value)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// TLS-Encryption
|
||||||
|
case "tls":
|
||||||
|
boolValue, isBool := readBool(value)
|
||||||
|
if isBool {
|
||||||
|
if boolValue {
|
||||||
|
cfg.TLSConfig = "true"
|
||||||
|
} else {
|
||||||
|
cfg.TLSConfig = "false"
|
||||||
|
}
|
||||||
|
} else if vl := strings.ToLower(value); vl == "skip-verify" || vl == "preferred" {
|
||||||
|
cfg.TLSConfig = vl
|
||||||
|
} else {
|
||||||
|
name, err := url.QueryUnescape(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid value for TLS config name: %v", err)
|
||||||
|
}
|
||||||
|
cfg.TLSConfig = name
|
||||||
|
}
|
||||||
|
|
||||||
|
// I/O write Timeout
|
||||||
|
case "writeTimeout":
|
||||||
|
cfg.WriteTimeout, err = time.ParseDuration(value)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
case "maxAllowedPacket":
|
||||||
|
cfg.MaxAllowedPacket, err = strconv.Atoi(value)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Connection attributes
|
||||||
|
case "connectionAttributes":
|
||||||
|
connectionAttributes, err := url.QueryUnescape(value)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid connectionAttributes value: %v", err)
|
||||||
|
}
|
||||||
|
cfg.ConnectionAttributes = connectionAttributes
|
||||||
|
|
||||||
|
default:
|
||||||
|
// lazy init
|
||||||
|
if cfg.Params == nil {
|
||||||
|
cfg.Params = make(map[string]string)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Params[key], err = url.QueryUnescape(value); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
func ensureHavePort(addr string) string {
|
||||||
|
if _, _, err := net.SplitHostPort(addr); err != nil {
|
||||||
|
return net.JoinHostPort(addr, "3306")
|
||||||
|
}
|
||||||
|
return addr
|
||||||
|
}
|
||||||
+83
@@ -0,0 +1,83 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2013 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"os"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Various errors the driver might return. Can change between driver versions.
|
||||||
|
var (
|
||||||
|
ErrInvalidConn = errors.New("invalid connection")
|
||||||
|
ErrMalformPkt = errors.New("malformed packet")
|
||||||
|
ErrNoTLS = errors.New("TLS requested but server does not support TLS")
|
||||||
|
ErrCleartextPassword = errors.New("this user requires clear text authentication. If you still want to use it, please add 'allowCleartextPasswords=1' to your DSN")
|
||||||
|
ErrNativePassword = errors.New("this user requires mysql native password authentication")
|
||||||
|
ErrOldPassword = errors.New("this user requires old password authentication. If you still want to use it, please add 'allowOldPasswords=1' to your DSN. See also https://github.com/go-sql-driver/mysql/wiki/old_passwords")
|
||||||
|
ErrUnknownPlugin = errors.New("this authentication plugin is not supported")
|
||||||
|
ErrOldProtocol = errors.New("MySQL server does not support required protocol 41+")
|
||||||
|
ErrPktSync = errors.New("commands out of sync. You can't run this command now")
|
||||||
|
ErrPktSyncMul = errors.New("commands out of sync. Did you run multiple statements at once?")
|
||||||
|
ErrPktTooLarge = errors.New("packet for query is too large. Try adjusting the `Config.MaxAllowedPacket`")
|
||||||
|
ErrBusyBuffer = errors.New("busy buffer")
|
||||||
|
|
||||||
|
// errBadConnNoWrite is used for connection errors where nothing was sent to the database yet.
|
||||||
|
// If this happens first in a function starting a database interaction, it should be replaced by driver.ErrBadConn
|
||||||
|
// to trigger a resend. Use mc.markBadConn(err) to do this.
|
||||||
|
// See https://github.com/go-sql-driver/mysql/pull/302
|
||||||
|
errBadConnNoWrite = errors.New("bad connection")
|
||||||
|
)
|
||||||
|
|
||||||
|
var defaultLogger = Logger(log.New(os.Stderr, "[mysql] ", log.Ldate|log.Ltime))
|
||||||
|
|
||||||
|
// Logger is used to log critical error messages.
|
||||||
|
type Logger interface {
|
||||||
|
Print(v ...any)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NopLogger is a nop implementation of the Logger interface.
|
||||||
|
type NopLogger struct{}
|
||||||
|
|
||||||
|
// Print implements Logger interface.
|
||||||
|
func (nl *NopLogger) Print(_ ...any) {}
|
||||||
|
|
||||||
|
// SetLogger is used to set the default logger for critical errors.
|
||||||
|
// The initial logger is os.Stderr.
|
||||||
|
func SetLogger(logger Logger) error {
|
||||||
|
if logger == nil {
|
||||||
|
return errors.New("logger is nil")
|
||||||
|
}
|
||||||
|
defaultLogger = logger
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// MySQLError is an error type which represents a single MySQL error
|
||||||
|
type MySQLError struct {
|
||||||
|
Number uint16
|
||||||
|
SQLState [5]byte
|
||||||
|
Message string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (me *MySQLError) Error() string {
|
||||||
|
if me.SQLState != [5]byte{} {
|
||||||
|
return fmt.Sprintf("Error %d (%s): %s", me.Number, me.SQLState, me.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("Error %d: %s", me.Number, me.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (me *MySQLError) Is(err error) bool {
|
||||||
|
if merr, ok := err.(*MySQLError); ok {
|
||||||
|
return merr.Number == me.Number
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
+224
@@ -0,0 +1,224 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2017 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"reflect"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (mf *mysqlField) typeDatabaseName() string {
|
||||||
|
switch mf.fieldType {
|
||||||
|
case fieldTypeBit:
|
||||||
|
return "BIT"
|
||||||
|
case fieldTypeBLOB:
|
||||||
|
if mf.charSet != binaryCollationID {
|
||||||
|
return "TEXT"
|
||||||
|
}
|
||||||
|
return "BLOB"
|
||||||
|
case fieldTypeDate:
|
||||||
|
return "DATE"
|
||||||
|
case fieldTypeDateTime:
|
||||||
|
return "DATETIME"
|
||||||
|
case fieldTypeDecimal:
|
||||||
|
return "DECIMAL"
|
||||||
|
case fieldTypeDouble:
|
||||||
|
return "DOUBLE"
|
||||||
|
case fieldTypeEnum:
|
||||||
|
return "ENUM"
|
||||||
|
case fieldTypeFloat:
|
||||||
|
return "FLOAT"
|
||||||
|
case fieldTypeGeometry:
|
||||||
|
return "GEOMETRY"
|
||||||
|
case fieldTypeInt24:
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return "UNSIGNED MEDIUMINT"
|
||||||
|
}
|
||||||
|
return "MEDIUMINT"
|
||||||
|
case fieldTypeJSON:
|
||||||
|
return "JSON"
|
||||||
|
case fieldTypeLong:
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return "UNSIGNED INT"
|
||||||
|
}
|
||||||
|
return "INT"
|
||||||
|
case fieldTypeLongBLOB:
|
||||||
|
if mf.charSet != binaryCollationID {
|
||||||
|
return "LONGTEXT"
|
||||||
|
}
|
||||||
|
return "LONGBLOB"
|
||||||
|
case fieldTypeLongLong:
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return "UNSIGNED BIGINT"
|
||||||
|
}
|
||||||
|
return "BIGINT"
|
||||||
|
case fieldTypeMediumBLOB:
|
||||||
|
if mf.charSet != binaryCollationID {
|
||||||
|
return "MEDIUMTEXT"
|
||||||
|
}
|
||||||
|
return "MEDIUMBLOB"
|
||||||
|
case fieldTypeNewDate:
|
||||||
|
return "DATE"
|
||||||
|
case fieldTypeNewDecimal:
|
||||||
|
return "DECIMAL"
|
||||||
|
case fieldTypeNULL:
|
||||||
|
return "NULL"
|
||||||
|
case fieldTypeSet:
|
||||||
|
return "SET"
|
||||||
|
case fieldTypeShort:
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return "UNSIGNED SMALLINT"
|
||||||
|
}
|
||||||
|
return "SMALLINT"
|
||||||
|
case fieldTypeString:
|
||||||
|
if mf.flags&flagEnum != 0 {
|
||||||
|
return "ENUM"
|
||||||
|
} else if mf.flags&flagSet != 0 {
|
||||||
|
return "SET"
|
||||||
|
}
|
||||||
|
if mf.charSet == binaryCollationID {
|
||||||
|
return "BINARY"
|
||||||
|
}
|
||||||
|
return "CHAR"
|
||||||
|
case fieldTypeTime:
|
||||||
|
return "TIME"
|
||||||
|
case fieldTypeTimestamp:
|
||||||
|
return "TIMESTAMP"
|
||||||
|
case fieldTypeTiny:
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return "UNSIGNED TINYINT"
|
||||||
|
}
|
||||||
|
return "TINYINT"
|
||||||
|
case fieldTypeTinyBLOB:
|
||||||
|
if mf.charSet != binaryCollationID {
|
||||||
|
return "TINYTEXT"
|
||||||
|
}
|
||||||
|
return "TINYBLOB"
|
||||||
|
case fieldTypeVarChar:
|
||||||
|
if mf.charSet == binaryCollationID {
|
||||||
|
return "VARBINARY"
|
||||||
|
}
|
||||||
|
return "VARCHAR"
|
||||||
|
case fieldTypeVarString:
|
||||||
|
if mf.charSet == binaryCollationID {
|
||||||
|
return "VARBINARY"
|
||||||
|
}
|
||||||
|
return "VARCHAR"
|
||||||
|
case fieldTypeYear:
|
||||||
|
return "YEAR"
|
||||||
|
case fieldTypeVector:
|
||||||
|
return "VECTOR"
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
scanTypeFloat32 = reflect.TypeOf(float32(0))
|
||||||
|
scanTypeFloat64 = reflect.TypeOf(float64(0))
|
||||||
|
scanTypeInt8 = reflect.TypeOf(int8(0))
|
||||||
|
scanTypeInt16 = reflect.TypeOf(int16(0))
|
||||||
|
scanTypeInt32 = reflect.TypeOf(int32(0))
|
||||||
|
scanTypeInt64 = reflect.TypeOf(int64(0))
|
||||||
|
scanTypeNullFloat = reflect.TypeOf(sql.NullFloat64{})
|
||||||
|
scanTypeNullInt = reflect.TypeOf(sql.NullInt64{})
|
||||||
|
scanTypeNullTime = reflect.TypeOf(sql.NullTime{})
|
||||||
|
scanTypeUint8 = reflect.TypeOf(uint8(0))
|
||||||
|
scanTypeUint16 = reflect.TypeOf(uint16(0))
|
||||||
|
scanTypeUint32 = reflect.TypeOf(uint32(0))
|
||||||
|
scanTypeUint64 = reflect.TypeOf(uint64(0))
|
||||||
|
scanTypeString = reflect.TypeOf("")
|
||||||
|
scanTypeNullString = reflect.TypeOf(sql.NullString{})
|
||||||
|
scanTypeBytes = reflect.TypeOf([]byte{})
|
||||||
|
scanTypeUnknown = reflect.TypeOf(new(any))
|
||||||
|
)
|
||||||
|
|
||||||
|
type mysqlField struct {
|
||||||
|
tableName string
|
||||||
|
name string
|
||||||
|
length uint32
|
||||||
|
flags fieldFlag
|
||||||
|
fieldType fieldType
|
||||||
|
decimals byte
|
||||||
|
charSet uint8
|
||||||
|
}
|
||||||
|
|
||||||
|
func (mf *mysqlField) scanType() reflect.Type {
|
||||||
|
switch mf.fieldType {
|
||||||
|
case fieldTypeTiny:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return scanTypeUint8
|
||||||
|
}
|
||||||
|
return scanTypeInt8
|
||||||
|
}
|
||||||
|
return scanTypeNullInt
|
||||||
|
|
||||||
|
case fieldTypeShort, fieldTypeYear:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return scanTypeUint16
|
||||||
|
}
|
||||||
|
return scanTypeInt16
|
||||||
|
}
|
||||||
|
return scanTypeNullInt
|
||||||
|
|
||||||
|
case fieldTypeInt24, fieldTypeLong:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return scanTypeUint32
|
||||||
|
}
|
||||||
|
return scanTypeInt32
|
||||||
|
}
|
||||||
|
return scanTypeNullInt
|
||||||
|
|
||||||
|
case fieldTypeLongLong:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
if mf.flags&flagUnsigned != 0 {
|
||||||
|
return scanTypeUint64
|
||||||
|
}
|
||||||
|
return scanTypeInt64
|
||||||
|
}
|
||||||
|
return scanTypeNullInt
|
||||||
|
|
||||||
|
case fieldTypeFloat:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
return scanTypeFloat32
|
||||||
|
}
|
||||||
|
return scanTypeNullFloat
|
||||||
|
|
||||||
|
case fieldTypeDouble:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
return scanTypeFloat64
|
||||||
|
}
|
||||||
|
return scanTypeNullFloat
|
||||||
|
|
||||||
|
case fieldTypeBit, fieldTypeTinyBLOB, fieldTypeMediumBLOB, fieldTypeLongBLOB,
|
||||||
|
fieldTypeBLOB, fieldTypeVarString, fieldTypeString, fieldTypeGeometry, fieldTypeVector:
|
||||||
|
if mf.charSet == binaryCollationID {
|
||||||
|
return scanTypeBytes
|
||||||
|
}
|
||||||
|
fallthrough
|
||||||
|
case fieldTypeDecimal, fieldTypeNewDecimal, fieldTypeVarChar,
|
||||||
|
fieldTypeEnum, fieldTypeSet, fieldTypeJSON, fieldTypeTime:
|
||||||
|
if mf.flags&flagNotNULL != 0 {
|
||||||
|
return scanTypeString
|
||||||
|
}
|
||||||
|
return scanTypeNullString
|
||||||
|
|
||||||
|
case fieldTypeDate, fieldTypeNewDate,
|
||||||
|
fieldTypeTimestamp, fieldTypeDateTime:
|
||||||
|
// NullTime is always returned for more consistent behavior as it can
|
||||||
|
// handle both cases of parseTime regardless if the field is nullable.
|
||||||
|
return scanTypeNullTime
|
||||||
|
|
||||||
|
default:
|
||||||
|
return scanTypeUnknown
|
||||||
|
}
|
||||||
|
}
|
||||||
+184
@@ -0,0 +1,184 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2013 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
fileRegister map[string]struct{}
|
||||||
|
fileRegisterLock sync.RWMutex
|
||||||
|
readerRegister map[string]func() io.Reader
|
||||||
|
readerRegisterLock sync.RWMutex
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterLocalFile adds the given file to the file allowlist,
|
||||||
|
// so that it can be used by "LOAD DATA LOCAL INFILE <filepath>".
|
||||||
|
// Alternatively you can allow the use of all local files with
|
||||||
|
// the DSN parameter 'allowAllFiles=true'
|
||||||
|
//
|
||||||
|
// filePath := "/home/gopher/data.csv"
|
||||||
|
// mysql.RegisterLocalFile(filePath)
|
||||||
|
// err := db.Exec("LOAD DATA LOCAL INFILE '" + filePath + "' INTO TABLE foo")
|
||||||
|
// if err != nil {
|
||||||
|
// ...
|
||||||
|
func RegisterLocalFile(filePath string) {
|
||||||
|
fileRegisterLock.Lock()
|
||||||
|
// lazy map init
|
||||||
|
if fileRegister == nil {
|
||||||
|
fileRegister = make(map[string]struct{})
|
||||||
|
}
|
||||||
|
|
||||||
|
fileRegister[strings.Trim(filePath, `"`)] = struct{}{}
|
||||||
|
fileRegisterLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeregisterLocalFile removes the given filepath from the allowlist.
|
||||||
|
func DeregisterLocalFile(filePath string) {
|
||||||
|
fileRegisterLock.Lock()
|
||||||
|
delete(fileRegister, strings.Trim(filePath, `"`))
|
||||||
|
fileRegisterLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterReaderHandler registers a handler function which is used
|
||||||
|
// to receive a io.Reader.
|
||||||
|
// The Reader can be used by "LOAD DATA LOCAL INFILE Reader::<name>".
|
||||||
|
// If the handler returns a io.ReadCloser Close() is called when the
|
||||||
|
// request is finished.
|
||||||
|
//
|
||||||
|
// mysql.RegisterReaderHandler("data", func() io.Reader {
|
||||||
|
// var csvReader io.Reader // Some Reader that returns CSV data
|
||||||
|
// ... // Open Reader here
|
||||||
|
// return csvReader
|
||||||
|
// })
|
||||||
|
// err := db.Exec("LOAD DATA LOCAL INFILE 'Reader::data' INTO TABLE foo")
|
||||||
|
// if err != nil {
|
||||||
|
// ...
|
||||||
|
func RegisterReaderHandler(name string, handler func() io.Reader) {
|
||||||
|
readerRegisterLock.Lock()
|
||||||
|
// lazy map init
|
||||||
|
if readerRegister == nil {
|
||||||
|
readerRegister = make(map[string]func() io.Reader)
|
||||||
|
}
|
||||||
|
|
||||||
|
readerRegister[name] = handler
|
||||||
|
readerRegisterLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeregisterReaderHandler removes the ReaderHandler function with
|
||||||
|
// the given name from the registry.
|
||||||
|
func DeregisterReaderHandler(name string) {
|
||||||
|
readerRegisterLock.Lock()
|
||||||
|
delete(readerRegister, name)
|
||||||
|
readerRegisterLock.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func deferredClose(err *error, closer io.Closer) {
|
||||||
|
closeErr := closer.Close()
|
||||||
|
if *err == nil {
|
||||||
|
*err = closeErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const defaultPacketSize = 16 * 1024 // 16KB is small enough for disk readahead and large enough for TCP
|
||||||
|
|
||||||
|
func (mc *okHandler) handleInFileRequest(name string) (err error) {
|
||||||
|
var rdr io.Reader
|
||||||
|
packetSize := defaultPacketSize
|
||||||
|
if mc.maxWriteSize < packetSize {
|
||||||
|
packetSize = mc.maxWriteSize
|
||||||
|
}
|
||||||
|
|
||||||
|
if idx := strings.Index(name, "Reader::"); idx == 0 || (idx > 0 && name[idx-1] == '/') { // io.Reader
|
||||||
|
// The server might return an an absolute path. See issue #355.
|
||||||
|
name = name[idx+8:]
|
||||||
|
|
||||||
|
readerRegisterLock.RLock()
|
||||||
|
handler, inMap := readerRegister[name]
|
||||||
|
readerRegisterLock.RUnlock()
|
||||||
|
|
||||||
|
if inMap {
|
||||||
|
rdr = handler()
|
||||||
|
if rdr != nil {
|
||||||
|
if cl, ok := rdr.(io.Closer); ok {
|
||||||
|
defer deferredClose(&err, cl)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
err = fmt.Errorf("reader '%s' is <nil>", name)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
err = fmt.Errorf("reader '%s' is not registered", name)
|
||||||
|
}
|
||||||
|
} else { // File
|
||||||
|
name = strings.Trim(name, `"`)
|
||||||
|
fileRegisterLock.RLock()
|
||||||
|
_, exists := fileRegister[name]
|
||||||
|
fileRegisterLock.RUnlock()
|
||||||
|
if mc.cfg.AllowAllFiles || exists {
|
||||||
|
var file *os.File
|
||||||
|
var fi os.FileInfo
|
||||||
|
|
||||||
|
if file, err = os.Open(name); err == nil {
|
||||||
|
defer deferredClose(&err, file)
|
||||||
|
|
||||||
|
// get file size
|
||||||
|
if fi, err = file.Stat(); err == nil {
|
||||||
|
rdr = file
|
||||||
|
if fileSize := int(fi.Size()); fileSize < packetSize {
|
||||||
|
packetSize = fileSize
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
err = fmt.Errorf("local file '%s' is not registered", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// send content packets
|
||||||
|
var data []byte
|
||||||
|
|
||||||
|
// if packetSize == 0, the Reader contains no data
|
||||||
|
if err == nil && packetSize > 0 {
|
||||||
|
data = make([]byte, 4+packetSize)
|
||||||
|
var n int
|
||||||
|
for err == nil {
|
||||||
|
n, err = rdr.Read(data[4:])
|
||||||
|
if n > 0 {
|
||||||
|
if ioErr := mc.conn().writePacket(data[:4+n]); ioErr != nil {
|
||||||
|
return ioErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err == io.EOF {
|
||||||
|
err = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// send empty packet (termination)
|
||||||
|
if data == nil {
|
||||||
|
data = make([]byte, 4)
|
||||||
|
}
|
||||||
|
if ioErr := mc.conn().writePacket(data[:4]); ioErr != nil {
|
||||||
|
return ioErr
|
||||||
|
}
|
||||||
|
mc.conn().syncSequence()
|
||||||
|
|
||||||
|
// read OK packet
|
||||||
|
if err == nil {
|
||||||
|
return mc.readResultOK()
|
||||||
|
}
|
||||||
|
|
||||||
|
mc.conn().readPacket()
|
||||||
|
return err
|
||||||
|
}
|
||||||
+71
@@ -0,0 +1,71 @@
|
|||||||
|
// Go MySQL Driver - A MySQL-Driver for Go's database/sql package
|
||||||
|
//
|
||||||
|
// Copyright 2013 The Go-MySQL-Driver Authors. All rights reserved.
|
||||||
|
//
|
||||||
|
// This Source Code Form is subject to the terms of the Mozilla Public
|
||||||
|
// License, v. 2.0. If a copy of the MPL was not distributed with this file,
|
||||||
|
// You can obtain one at http://mozilla.org/MPL/2.0/.
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"database/sql/driver"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NullTime represents a time.Time that may be NULL.
|
||||||
|
// NullTime implements the Scanner interface so
|
||||||
|
// it can be used as a scan destination:
|
||||||
|
//
|
||||||
|
// var nt NullTime
|
||||||
|
// err := db.QueryRow("SELECT time FROM foo WHERE id=?", id).Scan(&nt)
|
||||||
|
// ...
|
||||||
|
// if nt.Valid {
|
||||||
|
// // use nt.Time
|
||||||
|
// } else {
|
||||||
|
// // NULL value
|
||||||
|
// }
|
||||||
|
//
|
||||||
|
// # This NullTime implementation is not driver-specific
|
||||||
|
//
|
||||||
|
// Deprecated: NullTime doesn't honor the loc DSN parameter.
|
||||||
|
// NullTime.Scan interprets a time as UTC, not the loc DSN parameter.
|
||||||
|
// Use sql.NullTime instead.
|
||||||
|
type NullTime sql.NullTime
|
||||||
|
|
||||||
|
// Scan implements the Scanner interface.
|
||||||
|
// The value type must be time.Time or string / []byte (formatted time-string),
|
||||||
|
// otherwise Scan fails.
|
||||||
|
func (nt *NullTime) Scan(value any) (err error) {
|
||||||
|
if value == nil {
|
||||||
|
nt.Time, nt.Valid = time.Time{}, false
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := value.(type) {
|
||||||
|
case time.Time:
|
||||||
|
nt.Time, nt.Valid = v, true
|
||||||
|
return
|
||||||
|
case []byte:
|
||||||
|
nt.Time, err = parseDateTime(v, time.UTC)
|
||||||
|
nt.Valid = (err == nil)
|
||||||
|
return
|
||||||
|
case string:
|
||||||
|
nt.Time, err = parseDateTime([]byte(v), time.UTC)
|
||||||
|
nt.Valid = (err == nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
nt.Valid = false
|
||||||
|
return fmt.Errorf("can't convert %T to time.Time", value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Value implements the driver Valuer interface.
|
||||||
|
func (nt NullTime) Value() (driver.Value, error) {
|
||||||
|
if !nt.Valid {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return nt.Time, nil
|
||||||
|
}
|
||||||
+1384
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user