86 lines
1.9 KiB
Go
86 lines
1.9 KiB
Go
// Applies nl-pms-api/init.sql to MySQL.
|
||
// Usage (from view repo root):
|
||
//
|
||
// go run ./tools/applyinit -dsn "user:pass@tcp(host:3306)/"
|
||
//
|
||
// Or set MYSQL_DSN. Schema source of truth is ../nl-pms-api/init.sql.
|
||
package main
|
||
|
||
import (
|
||
"database/sql"
|
||
"flag"
|
||
"fmt"
|
||
"os"
|
||
"path/filepath"
|
||
"strings"
|
||
|
||
_ "github.com/go-sql-driver/mysql"
|
||
)
|
||
|
||
func main() {
|
||
dsnFlag := flag.String("dsn", "", "MySQL DSN without database (or set MYSQL_DSN)")
|
||
sqlPath := flag.String("sql", "", "path to init.sql (default: ../nl-pms-api/init.sql)")
|
||
flag.Parse()
|
||
dsn := strings.TrimSpace(*dsnFlag)
|
||
if dsn == "" {
|
||
dsn = strings.TrimSpace(os.Getenv("MYSQL_DSN"))
|
||
}
|
||
if dsn == "" {
|
||
fmt.Fprintln(os.Stderr, "缺少 -dsn 或 MYSQL_DSN(例: root:root@tcp(127.0.0.1:3306)/)")
|
||
os.Exit(2)
|
||
}
|
||
if !strings.Contains(dsn, "multiStatements") {
|
||
if strings.Contains(dsn, "?") {
|
||
dsn += "&multiStatements=false&charset=utf8mb4"
|
||
} else {
|
||
dsn += "?multiStatements=false&charset=utf8mb4"
|
||
}
|
||
}
|
||
path := *sqlPath
|
||
if path == "" {
|
||
path = filepath.Join("..", "nl-pms-api", "init.sql")
|
||
}
|
||
raw, err := os.ReadFile(path)
|
||
if err != nil {
|
||
fmt.Fprintln(os.Stderr, err)
|
||
os.Exit(1)
|
||
}
|
||
db, err := sql.Open("mysql", dsn)
|
||
if err != nil {
|
||
panic(err)
|
||
}
|
||
defer db.Close()
|
||
db.SetMaxOpenConns(1)
|
||
|
||
var kept []string
|
||
for _, line := range strings.Split(string(raw), "\n") {
|
||
if t := strings.TrimSpace(line); strings.HasPrefix(t, "--") {
|
||
continue
|
||
}
|
||
kept = append(kept, line)
|
||
}
|
||
for _, stmt := range strings.Split(strings.Join(kept, "\n"), ";") {
|
||
s := strings.TrimSpace(stmt)
|
||
if s == "" {
|
||
continue
|
||
}
|
||
if _, err := db.Exec(s); err != nil {
|
||
fmt.Println("ERR:", err, "stmt:", s[:min(80, len(s))])
|
||
os.Exit(1)
|
||
}
|
||
}
|
||
var n int
|
||
if err := db.QueryRow("SELECT COUNT(*) FROM code_count.users").Scan(&n); err != nil {
|
||
fmt.Println("verify failed:", err)
|
||
os.Exit(1)
|
||
}
|
||
fmt.Println("ok, users rows:", n, "sql:", path)
|
||
}
|
||
|
||
func min(a, b int) int {
|
||
if a < b {
|
||
return a
|
||
}
|
||
return b
|
||
}
|