From 7d082031d1edad8f3a4ca09730d3d4b413926e87 Mon Sep 17 00:00:00 2001 From: devituz Date: Fri, 25 Sep 2026 02:23:04 +0300 Subject: [PATCH 1/2] =?UTF-8?q?fix:=20bug=20sweep=20=E2=80=94=20ORM=20hydr?= =?UTF-8?q?ation/save,=20query=20builder=20edge=20cases,=20PG=20migration?= =?UTF-8?q?=20lock,=20lago=20CLI=20bootstrap,=20gin=20QueryLog?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Each fix has a regression test that fails on d537c37 and passes here. See CHANGELOG (Unreleased) and the PR description for details. --- CHANGELOG.md | 19 ++ adapters/gin/lagogin_test.go | 26 +++ adapters/gin/middleware.go | 81 ++++++-- cli/bootstrap.go | 74 ++++++++ cli/bootstrap_test.go | 41 ++++ cli/cmd/init_test.go | 63 +++++++ cli/cmd/new.go | 1 + cli/cmd/project.go | 53 +++++- cmd/artisan/main.go | 37 +--- cmd/lago/main.go | 8 +- database/connection.go | 23 +++ internal/reflectutil/assign.go | 169 +++++++++++++++++ internal/reflectutil/cache.go | 27 +++ migrations/lock.go | 58 +++++- migrations/lock_pg_test.go | 54 ++++++ orm/builder_ext.go | 9 + orm/query.go | 220 +++++++++------------- orm/sweep_regression_test.go | 335 +++++++++++++++++++++++++++++++++ query/builder.go | 95 +++++++++- query/edge_test.go | 18 +- query/sweep_regression_test.go | 126 +++++++++++++ relations/relations.go | 168 ++++++++++++----- web/middleware.go | 8 +- web/security_test.go | 21 +++ 24 files changed, 1486 insertions(+), 248 deletions(-) create mode 100644 cli/bootstrap.go create mode 100644 cli/bootstrap_test.go create mode 100644 cli/cmd/init_test.go create mode 100644 internal/reflectutil/assign.go create mode 100644 migrations/lock_pg_test.go create mode 100644 orm/sweep_regression_test.go create mode 100644 query/sweep_regression_test.go diff --git a/CHANGELOG.md b/CHANGELOG.md index 028e627..f9013e1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,6 +3,25 @@ All notable changes are recorded here. Versions follow [SemVer](https://semver.org/). Pre-`v1.0.0` releases may include breaking changes between minor versions. +## Unreleased + +Bug sweep (fix/bug-sweep-2026-09). Every fix ships with a regression test. + +### Fixed +- **`orm` — `sql.Scanner` fields (`sql.NullString`, `*sql.NullString`, ...) failed to hydrate**; numeric columns read into `string` fields became a rune; textual numbers/booleans (MySQL text protocol) failed for float/bool/uint fields. +- **`orm` — `AfterFind` hook was never invoked.** +- **`orm` — `Save` never inserted models with a caller-assigned non-auto-increment key** (UUID/string PK); models without a PK panicked. +- **`orm` — `*Struct` relation fields without a `relation` tag were persisted as columns**, so `Save` failed (`no column named author`). +- **`orm` — `Paginate` and `Chunk` ignored `With(...)`.** +- **`orm` — soft-delete scope / `Chunk` cursor bound only to the last `OrWhere` branch**, leaking trashed rows and looping forever in `Chunk`. +- **`orm` — cast `ToDB` errors were swallowed.** +- **`relations` — NULL columns and mismatched key types (uint64 vs int64/[]byte) broke eager loading.** +- **`query` — `Offset` without `Limit` was a syntax error on SQLite/MySQL; `WhereIn` rejected slice types other than a fixed list; `Where(col, nil)` compiled to `= NULL`; `Distinct().Count()` counted all rows; aggregates with `Offset` returned `sql.ErrNoRows`; empty nested groups rendered `()`; Postgres placeholders ignored `Join` args.** New `Builder.WrapWheres()`. +- **`migrations` — Postgres advisory lock could be released on a different pooled session** (lock leaked, next migrator blocked forever) and ignored the timeout. +- **`cli` — `lago migrate` never saw project migrations.** `lago init`/`lago new` now scaffold `cmd/lago/main.go`; both `lago` and `artisan` re-run it for registry-dependent commands. +- **`adapters/gin` — `X-DB-Query-Count` was always 0 and set after the body was written.** New `database.Connection.OnQuery` hook. +- **`web` — `CORSWithConfig` ignored an explicit `AllowedHeaders` list.** + ## v0.26.0 — 2026-06-25 Production-hardening release. A fleet of adversarial test agents (load, fuzz, diff --git a/adapters/gin/lagogin_test.go b/adapters/gin/lagogin_test.go index 44169a5..e5d2866 100644 --- a/adapters/gin/lagogin_test.go +++ b/adapters/gin/lagogin_test.go @@ -369,6 +369,32 @@ func TestQueryLogHeader(t *testing.T) { } } +// Real queries never reached the counter (only manual ObserveQuery did), so +// the header was always 0; it must count exactly this request's queries. +func TestQueryLogCountsRealQueries(t *testing.T) { + conn := newTestConn(t) + + r := gin.New() + r.Use(lagogin.QueryLogN(conn, 1000)) // no explicit Instrument call + r.GET("/q", lagogin.H(func(c *lagogin.Ctx) (any, error) { + _, _ = orm.Query[User](conn).Count(c.Ctx()) + var users []User + _ = orm.Query[User](conn).Limit(1).Get(c.Ctx(), &users) + // Queries outside the request context are not attributed to it. + _, _ = orm.Query[User](conn).Count(context.Background()) + return "ok", nil + })) + w := do(r, "GET", "/q", nil) + if w.Code != http.StatusOK { + t.Fatalf("want 200, got %d", w.Code) + } + // Result().Header is the header as sent; w.Header() would also show + // values set after the body was written, which never reach the client. + if got := w.Result().Header.Get("X-DB-Query-Count"); got != "2" { + t.Fatalf("X-DB-Query-Count on the wire = %q, want 2", got) + } +} + // --- 7. OpenAPI generation --------------------------------------------- func TestOpenAPIContainsResourcePaths(t *testing.T) { diff --git a/adapters/gin/middleware.go b/adapters/gin/middleware.go index f0ddf8a..01cc0d9 100644 --- a/adapters/gin/middleware.go +++ b/adapters/gin/middleware.go @@ -103,7 +103,9 @@ func RequestTimeout(d time.Duration) gin.HandlerFunc { // If the count crosses threshold (20 by default), a WARN is logged with the // request path — a cheap N+1 detector for dev environments. // -// Requires the connection to be passed through Instrument() once at startup: +// Queries are counted per request: those executed with the request's context +// (c.Ctx() / c.Request.Context()) on conn, plus manual ObserveQuery calls. +// QueryLog instruments conn itself; calling Instrument() first is optional: // // conn = lagogin.Instrument(conn) // r.Use(lagogin.QueryLog(conn)) @@ -118,15 +120,28 @@ func QueryLogN(conn *database.Connection, threshold int) gin.HandlerFunc { return queryLogWith(conn, threshold) } +// queryCountKey carries the per-request query counter in the request context. +type queryCountKey struct{} + func queryLogWith(conn *database.Connection, threshold int) gin.HandlerFunc { + Instrument(conn) return func(c *gin.Context) { + n := new(atomic.Int64) + c.Request = c.Request.WithContext(context.WithValue(c.Request.Context(), queryCountKey{}, n)) before := globalQueryCount(conn) - c.Next() - count := globalQueryCount(conn) - before - if count < 0 { - count = 0 + counted := func() int64 { + if v := n.Load() + globalQueryCount(conn) - before; v > 0 { + return v + } + return 0 } - c.Writer.Header().Set("X-DB-Query-Count", strconv.FormatInt(count, 10)) + // Headers set after the handler wrote the body never reach the + // client, so stamp the header when the status line is written. + w := &queryCountWriter{ResponseWriter: c.Writer, count: counted} + c.Writer = w + c.Next() + w.stamp() + count := counted() if int(count) > threshold && conn.Log != nil { conn.Log.Warnf("lagogin: %d queries on %s %s (threshold %d) — possible N+1", count, c.Request.Method, c.Request.URL.Path, threshold) @@ -134,13 +149,49 @@ func queryLogWith(conn *database.Connection, threshold int) gin.HandlerFunc { } } -// Instrument enables per-connection query counting for QueryLog. Call once -// at startup before installing the middleware. The returned connection is -// the same pointer — Instrument only registers it with the global counter -// table and replaces conn.Log with a counting wrapper. +// queryCountWriter sets X-DB-Query-Count right before the response header +// is committed. +type queryCountWriter struct { + gin.ResponseWriter + count func() int64 + stamped bool +} + +func (w *queryCountWriter) stamp() { + if w.stamped || w.ResponseWriter.Written() { + return + } + w.stamped = true + w.Header().Set("X-DB-Query-Count", strconv.FormatInt(w.count(), 10)) +} + +func (w *queryCountWriter) WriteHeader(code int) { + w.stamp() + w.ResponseWriter.WriteHeader(code) +} + +func (w *queryCountWriter) WriteHeaderNow() { + w.stamp() + w.ResponseWriter.WriteHeaderNow() +} + +func (w *queryCountWriter) Write(b []byte) (int, error) { + w.stamp() + return w.ResponseWriter.Write(b) +} + +func (w *queryCountWriter) WriteString(s string) (int, error) { + w.stamp() + return w.ResponseWriter.WriteString(s) +} + +// Instrument enables query counting for QueryLog. It is idempotent and +// returns the same pointer: it registers conn with the counter table and +// installs a database query hook that bumps the counter of the request whose +// context the statement ran with. Logging settings are left untouched. // -// The wrapper delegates Info/Warn/Error/SQL/SlowSQL to the original logger, -// so SQL tracing and slow-query reporting continue to work unchanged. +// Previously nothing but manual ObserveQuery calls fed the counter, so +// X-DB-Query-Count was always 0 for real traffic. func Instrument(conn *database.Connection) *database.Connection { if conn == nil { return nil @@ -151,7 +202,11 @@ func Instrument(conn *database.Connection) *database.Connection { return conn } counters[conn] = new(atomic.Int64) - conn.Config.LogQueries = true + conn.OnQuery(func(ctx context.Context, _ string, _ []any, _ time.Duration, _ error) { + if n, ok := ctx.Value(queryCountKey{}).(*atomic.Int64); ok { + n.Add(1) + } + }) return conn } diff --git a/cli/bootstrap.go b/cli/bootstrap.go new file mode 100644 index 0000000..5012cdf --- /dev/null +++ b/cli/bootstrap.go @@ -0,0 +1,74 @@ +package cli + +import ( + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" +) + +// BootstrapEnv is set in the environment of a re-executed project-local CLI +// so it does not bootstrap again (infinite recursion). Setting it manually +// disables the indirection. +const BootstrapEnv = "LAGO_BOOTSTRAPPED" + +// projectEntrypoints are the project-local CLI mains the global binaries look +// for, in order. `lago init` / `lago new` scaffold cmd/lago. +var projectEntrypoints = []string{"cmd/lago", "cmd/artisan"} + +// RunProjectBinary re-executes the project-local CLI via `go run` when the +// working directory contains one (cmd/lago/main.go or cmd/artisan/main.go). +// That binary blank-imports the project's migrations and seeders packages, so +// their init() functions populate the registries — a globally installed +// binary cannot see them and would report "nothing to migrate". It returns +// true when it handled execution; the caller must then return. On failure it +// exits with the child's status. +// +// Scaffolding commands (init, new, env*, key:generate, make:*, gen:*) never +// need the project's registries and run in-process, so they keep working in a +// fresh project whose go.sum cannot build the local entrypoint yet. +func RunProjectBinary() bool { + if os.Getenv(BootstrapEnv) == "1" || !needsProject(os.Args[1:]) { + return false + } + for _, dir := range projectEntrypoints { + if _, err := os.Stat(filepath.Join(dir, "main.go")); err != nil { + continue + } + c := exec.Command("go", append([]string{"run", "./" + dir}, os.Args[1:]...)...) + c.Stdin = os.Stdin + c.Stdout = os.Stdout + c.Stderr = os.Stderr + c.Env = append(os.Environ(), BootstrapEnv+"=1") + if err := c.Run(); err != nil { + if ee, ok := err.(*exec.ExitError); ok { + os.Exit(ee.ExitCode()) + } + fmt.Fprintln(os.Stderr, "lago: bootstrap failed:", err) + os.Exit(1) + } + return true + } + return false +} + +// needsProject reports whether the command named by args may depend on the +// project's registered migrations, seeders or custom commands. +func needsProject(args []string) bool { + name := "" + for _, a := range args { + if !strings.HasPrefix(a, "-") { + name = a + break + } + } + switch { + case name == "", name == "help", name == "completion", name == "version", + name == "init", name == "new", name == "env", name == "key:generate", + strings.HasPrefix(name, "env:"), strings.HasPrefix(name, "make:"), + strings.HasPrefix(name, "gen:"): + return false + } + return true +} diff --git a/cli/bootstrap_test.go b/cli/bootstrap_test.go new file mode 100644 index 0000000..3718866 --- /dev/null +++ b/cli/bootstrap_test.go @@ -0,0 +1,41 @@ +package cli + +import ( + "strings" + "testing" +) + +// RunProjectBinary must be a no-op outside a project and inside the +// re-executed child (otherwise it would recurse forever). +func TestRunProjectBinary_NoOpWithoutEntrypointOrWhenBootstrapped(t *testing.T) { + t.Chdir(t.TempDir()) + if RunProjectBinary() { + t.Fatal("handled execution without a project entrypoint") + } + t.Setenv(BootstrapEnv, "1") + if RunProjectBinary() { + t.Fatal("re-executed inside an already bootstrapped child") + } +} + +// Scaffolding commands run in-process; registry-dependent ones re-execute. +func TestNeedsProject(t *testing.T) { + cases := map[string]bool{ + "": false, + "make:model Post -mfs": false, + "init": false, + "env:init": false, + "key:generate": false, + "gen:client": false, + "--help": false, + "migrate": true, + "migrate:fresh --seed": true, + "db:seed": true, + "my:custom": true, + } + for args, want := range cases { + if got := needsProject(strings.Fields(args)); got != want { + t.Errorf("needsProject(%q) = %v, want %v", args, got, want) + } + } +} diff --git a/cli/cmd/init_test.go b/cli/cmd/init_test.go new file mode 100644 index 0000000..662ac8a --- /dev/null +++ b/cli/cmd/init_test.go @@ -0,0 +1,63 @@ +package cmd + +import ( + "bytes" + "os" + "strings" + "testing" +) + +// `lago init` did not scaffold a project-local CLI entrypoint, so the global +// `lago migrate` never saw the project's migrations ("nothing to migrate"). +func TestInit_ScaffoldsProjectCLI(t *testing.T) { + t.Chdir(t.TempDir()) + if err := os.WriteFile("go.mod", []byte("module github.com/you/myapp\n\ngo 1.25\n"), 0o644); err != nil { + t.Fatal(err) + } + c := NewInit(nil) + var buf bytes.Buffer + c.SetOut(&buf) + c.SetErr(&buf) + c.SetArgs(nil) + if err := c.Execute(); err != nil { + t.Fatalf("init: %v\n%s", err, buf.String()) + } + + main, err := os.ReadFile("cmd/lago/main.go") + if err != nil { + t.Fatalf("cmd/lago/main.go not created: %v", err) + } + for _, want := range []string{ + `_ "github.com/you/myapp/migrations"`, + `_ "github.com/you/myapp/seeders"`, + `cli.New(cli.Options{ProjectName: "lago"}).Execute()`, + } { + if !strings.Contains(string(main), want) { + t.Errorf("cmd/lago/main.go missing %q", want) + } + } + for _, f := range []string{"migrations/doc.go", "seeders/doc.go"} { + if _, err := os.Stat(f); err != nil { + t.Errorf("%s not created (the blank imports would not compile): %v", f, err) + } + } +} + +// Existing migrations/seeders files are left alone. +func TestInit_KeepsExistingPackageDocs(t *testing.T) { + t.Chdir(t.TempDir()) + _ = os.WriteFile("go.mod", []byte("module example.com/app\n"), 0o644) + _ = os.MkdirAll("migrations", 0o755) + _ = os.WriteFile("migrations/doc.go", []byte("package migrations // mine\n"), 0o644) + c := NewInit(nil) + c.SetOut(&bytes.Buffer{}) + c.SetErr(&bytes.Buffer{}) + c.SetArgs(nil) + if err := c.Execute(); err != nil { + t.Fatalf("init: %v", err) + } + b, _ := os.ReadFile("migrations/doc.go") + if string(b) != "package migrations // mine\n" { + t.Fatalf("existing migrations/doc.go overwritten: %q", b) + } +} diff --git a/cli/cmd/new.go b/cli/cmd/new.go index 630c3e5..9c15365 100644 --- a/cli/cmd/new.go +++ b/cli/cmd/new.go @@ -208,6 +208,7 @@ func scaffoldProject(cmd *cobra.Command, opts ScaffoldOptions, force bool) error {filepath.Join(root, "migrations", "doc.go"), pkgDocStub("migrations", "Schema migrations. Generated files call migrations.Register in init().")}, {filepath.Join(root, "factories", "doc.go"), pkgDocStub("factories", "Faker-powered model factories.")}, {filepath.Join(root, "seeders", "doc.go"), pkgDocStub("seeders", "Seeders register themselves in init() via seeder.Register.")}, + {filepath.Join(root, "cmd", "lago", "main.go"), projectCLIStub(opts.Module)}, {filepath.Join(root, "tests", ".keep"), ""}, {filepath.Join(root, "services", "doc.go"), pkgDocStub("services", "Framework-agnostic CRUD services.")}, {filepath.Join(root, "controllers", "doc.go"), pkgDocStub("controllers", "HTTP controllers (web or lagogin flavor).")}, diff --git a/cli/cmd/project.go b/cli/cmd/project.go index e4933e2..e560fe4 100644 --- a/cli/cmd/project.go +++ b/cli/cmd/project.go @@ -137,13 +137,64 @@ func NewInit(_ *Env) *cobra.Command { if err := writeIfNew(cmd, "routes/api.go", routesStub(module), force); err != nil { return err } - return nil + // 4. Project-local CLI entrypoint. The globally installed `lago` + // re-runs it so migrations/seeders registered in init() are + // visible; without it `lago migrate` reported "nothing to migrate". + if module == "" { + return nil + } + for _, pkg := range []string{"migrations", "seeders"} { + if err := writeIfMissing(cmd, filepath.Join(pkg, "doc.go"), pkgDocStub(pkg, projectPkgDocs[pkg])); err != nil { + return err + } + } + return writeIfNew(cmd, "cmd/lago/main.go", projectCLIStub(module), force) }, } c.Flags().BoolVar(&force, "force", false, "overwrite existing files") return c } +var projectPkgDocs = map[string]string{ + "migrations": "Schema migrations. Generated files call migrations.Register in init().", + "seeders": "Seeders register themselves in init() via seeder.Register.", +} + +// writeIfMissing writes path only when it does not exist yet; an existing +// file is left untouched without error. +func writeIfMissing(cmd *cobra.Command, path, body string) error { + if _, err := os.Stat(path); err == nil { + return nil + } + return writeIfNew(cmd, path, body, false) +} + +// projectCLIStub is the project-local CLI entrypoint (cmd/lago/main.go). It +// blank-imports the project's migrations and seeders so their init() +// registrations are visible to migrate/db:seed. +func projectCLIStub(module string) string { + return `// Command lago is this project's CLI entrypoint. The globally installed +// ` + "`lago`" + ` binary re-runs it (go run ./cmd/lago) so the migrations and +// seeders registered in init() below are visible to migrate / db:seed. +package main + +import ( + "github.com/devituz/lagodev/cli" + + _ "github.com/devituz/lagodev/drivers/mysql" + _ "github.com/devituz/lagodev/drivers/postgres" + _ "github.com/devituz/lagodev/drivers/sqlite" + + _ "` + module + `/migrations" // registers schema migrations via init() + _ "` + module + `/seeders" // registers seeders via init() +) + +func main() { + cli.New(cli.Options{ProjectName: "lago"}).Execute() +} +` +} + func mustMarshal(v any) string { b, _ := json.MarshalIndent(v, "", " ") return string(b) + "\n" diff --git a/cmd/artisan/main.go b/cmd/artisan/main.go index 769acf4..585cf74 100644 --- a/cmd/artisan/main.go +++ b/cmd/artisan/main.go @@ -5,7 +5,7 @@ // go build -o artisan ./cmd/artisan // // Auto-bootstrap: when invoked from a directory that contains -// ./cmd/artisan/main.go (the project-local artisan), the global binary +// ./cmd/lago/main.go or ./cmd/artisan/main.go (the project-local CLI), the global binary // transparently re-executes that local binary via `go run`. This is how // project-specific migrations and seeders — registered through init() in // the project's own packages — become visible without users having to @@ -15,11 +15,6 @@ package main import ( - "fmt" - "os" - "os/exec" - "path/filepath" - "github.com/devituz/lagodev/cli" // Blank-import drivers so DB_CONNECTION=sqlite|postgres|mysql Just Works. @@ -28,35 +23,11 @@ import ( _ "github.com/devituz/lagodev/drivers/sqlite" ) -const bootstrapEnv = "LAGO_BOOTSTRAPPED" - func main() { - if os.Getenv(bootstrapEnv) != "1" && bootstrapToProjectBinary() { + // Re-run through the project-local cmd/lago or cmd/artisan entrypoint + // when present (see cli.RunProjectBinary). + if cli.RunProjectBinary() { return } cli.Default().Execute() } - -// bootstrapToProjectBinary re-execs the project-local artisan via `go run` -// when ./cmd/artisan/main.go exists in the current directory. That local -// binary blank-imports the project's migrations/seeders packages, so their -// init() functions populate the global registries before the CLI runs. -// Returns true when it handled execution (caller should not continue). -func bootstrapToProjectBinary() bool { - if _, err := os.Stat(filepath.Join("cmd", "artisan", "main.go")); err != nil { - return false - } - c := exec.Command("go", append([]string{"run", "./cmd/artisan"}, os.Args[1:]...)...) - c.Stdin = os.Stdin - c.Stdout = os.Stdout - c.Stderr = os.Stderr - c.Env = append(os.Environ(), bootstrapEnv+"=1") - if err := c.Run(); err != nil { - if ee, ok := err.(*exec.ExitError); ok { - os.Exit(ee.ExitCode()) - } - fmt.Fprintln(os.Stderr, "lago: bootstrap failed:", err) - os.Exit(1) - } - return true -} diff --git a/cmd/lago/main.go b/cmd/lago/main.go index 6cc45cf..4d0f91c 100644 --- a/cmd/lago/main.go +++ b/cmd/lago/main.go @@ -4,7 +4,10 @@ // go install github.com/devituz/lagodev/cmd/lago@latest # → lago migrate // go install github.com/devituz/lagodev/cmd/artisan@latest # → artisan migrate // -// Both binaries share the same command tree, drivers, and flags. +// Both binaries share the same command tree, drivers, and flags. Inside a +// project that has cmd/lago/main.go (scaffolded by `lago init` / `lago new`) +// or cmd/artisan/main.go, the command is re-run through that project-local +// entrypoint so the project's migrations and seeders are registered. package main import ( @@ -16,6 +19,9 @@ import ( ) func main() { + if cli.RunProjectBinary() { + return + } app := cli.New(cli.Options{ProjectName: "lago"}) app.Execute() } diff --git a/database/connection.go b/database/connection.go index c2cf9d8..274bcd3 100644 --- a/database/connection.go +++ b/database/connection.go @@ -79,6 +79,23 @@ type Connection struct { mu sync.RWMutex closed bool + hooks []QueryHook +} + +// QueryHook observes every statement executed through a Connection (not +// through a Tx): the SQL, its args, the elapsed time and the error, if any. +type QueryHook func(ctx context.Context, query string, args []any, took time.Duration, err error) + +// OnQuery registers a hook invoked after each statement the connection +// executes. Hooks run synchronously on the calling goroutine and must be +// cheap and safe for concurrent use. +func (c *Connection) OnQuery(h QueryHook) { + if h == nil { + return + } + c.mu.Lock() + defer c.mu.Unlock() + c.hooks = append(c.hooks, h) } // ErrClosed is returned by exec/query/transaction methods called after the @@ -200,6 +217,12 @@ func (c *Connection) TransactionWith(ctx context.Context, opts *sql.TxOptions, f } func (c *Connection) observe(ctx context.Context, query string, args []any, took time.Duration, err error) { + c.mu.RLock() + hooks := c.hooks + c.mu.RUnlock() + for _, h := range hooks { + h(ctx, query, args, took, err) + } if c.Log == nil { return } diff --git a/internal/reflectutil/assign.go b/internal/reflectutil/assign.go new file mode 100644 index 0000000..91f540d --- /dev/null +++ b/internal/reflectutil/assign.go @@ -0,0 +1,169 @@ +package reflectutil + +import ( + "database/sql" + "fmt" + "reflect" + "strconv" + "time" +) + +// AssignScanned writes a value scanned from the DB (as an any holding the +// driver's concrete type, or nil for NULL) into the destination field. NULL +// coalesces to the field's Go zero value. Fields implementing sql.Scanner +// (sql.NullString, decimal/uuid types, ...) receive the raw value through +// their Scan method; everything else goes through ConvertAssign. +func AssignScanned(fv reflect.Value, raw any) error { + if fv.CanAddr() { + if sc, ok := fv.Addr().Interface().(sql.Scanner); ok { + return sc.Scan(raw) + } + } + if raw == nil { + // NULL → zero value. For pointer fields leave nil; otherwise reset. + fv.Set(reflect.Zero(fv.Type())) + return nil + } + if fv.Kind() == reflect.Ptr { + if sc, ok := reflect.New(fv.Type().Elem()).Interface().(sql.Scanner); ok { + if err := sc.Scan(raw); err != nil { + return err + } + fv.Set(reflect.ValueOf(sc)) + return nil + } + if fv.IsNil() { + fv.Set(reflect.New(fv.Type().Elem())) + } + return ConvertAssign(fv.Interface(), raw) + } + return ConvertAssign(fv.Addr().Interface(), raw) +} + +// ConvertAssign converts src (a driver value: int64, float64, bool, []byte, +// string or time.Time) into the pointer dest. It mirrors the subset of +// database/sql's convertAssign that the ORM relies on, with a reflection +// fallback for numeric/string kinds so non-default field types (int, uint, +// float32, named string types, ...) all work. Textual numbers and booleans +// ([]byte on MySQL's text protocol) are parsed. +func ConvertAssign(dest, src any) error { + dv := reflect.ValueOf(dest).Elem() + + // Fast path: src is directly assignable to the destination type. + sv := reflect.ValueOf(src) + if sv.Type().AssignableTo(dv.Type()) { + dv.Set(sv) + return nil + } + + switch dv.Kind() { + case reflect.String: + switch s := src.(type) { + case string: + dv.SetString(s) + return nil + case []byte: + dv.SetString(string(s)) + return nil + case int64: + // reflect's int→string conversion would yield a rune ("\x07"), + // so format numbers explicitly. + dv.SetString(strconv.FormatInt(s, 10)) + return nil + case float64: + dv.SetString(strconv.FormatFloat(s, 'g', -1, 64)) + return nil + case bool: + dv.SetString(strconv.FormatBool(s)) + return nil + case time.Time: + dv.SetString(s.Format(time.RFC3339Nano)) + return nil + } + return fmt.Errorf("orm: cannot assign %T to %s", src, dv.Type()) + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + if n, ok := ToInt64(src); ok { + dv.SetInt(n) + return nil + } + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + if n, ok := ToInt64(src); ok { + dv.SetUint(uint64(n)) + return nil + } + if s, ok := asText(src); ok { + if u, err := strconv.ParseUint(s, 10, 64); err == nil { + dv.SetUint(u) + return nil + } + } + case reflect.Float32, reflect.Float64: + switch f := src.(type) { + case float64: + dv.SetFloat(f) + return nil + case int64: + dv.SetFloat(float64(f)) + return nil + } + if s, ok := asText(src); ok { + if f, err := strconv.ParseFloat(s, 64); err == nil { + dv.SetFloat(f) + return nil + } + } + case reflect.Bool: + switch v := src.(type) { + case bool: + dv.SetBool(v) + return nil + case int64: + dv.SetBool(v != 0) + return nil + } + if s, ok := asText(src); ok { + if b, err := strconv.ParseBool(s); err == nil { + dv.SetBool(b) + return nil + } + } + } + + // Convertible numeric kinds (e.g. int64 → named int, float64 → float32). + if sv.Type().ConvertibleTo(dv.Type()) { + dv.Set(sv.Convert(dv.Type())) + return nil + } + return fmt.Errorf("orm: cannot assign %T to %s", src, dv.Type()) +} + +func asText(src any) (string, bool) { + switch s := src.(type) { + case string: + return s, true + case []byte: + return string(s), true + } + return "", false +} + +// ToInt64 converts a driver value (int64, int, float64, []byte, string) to int64. +func ToInt64(src any) (int64, bool) { + switch n := src.(type) { + case int64: + return n, true + case int: + return int64(n), true + case float64: + return int64(n), true + case []byte: + if v, err := strconv.ParseInt(string(n), 10, 64); err == nil { + return v, true + } + case string: + if v, err := strconv.ParseInt(n, 10, 64); err == nil { + return v, true + } + } + return 0, false +} diff --git a/internal/reflectutil/cache.go b/internal/reflectutil/cache.go index 853e62d..af28245 100644 --- a/internal/reflectutil/cache.go +++ b/internal/reflectutil/cache.go @@ -7,6 +7,8 @@ package reflectutil import ( + "database/sql" + "database/sql/driver" "reflect" "strings" "sync" @@ -315,9 +317,34 @@ func buildField(sf reflect.StructField, idx []int) *Field { } } } + // A struct / *struct field without a cast that the driver cannot handle + // (not time.Time, not a sql.Scanner / driver.Valuer) can never be a column + // value; it is a single-row relation destination (BelongsTo / HasOne, e.g. + // `Author *User`). Persisting it made every Save fail with "no column named + // author". + if !f.IsRelation && f.Cast == "" && isRelationStruct(sf.Type) { + f.IsRelation = true + } return f } +var ( + timeType = reflect.TypeOf(time.Time{}) + scannerType = reflect.TypeOf((*sql.Scanner)(nil)).Elem() + valuerType = reflect.TypeOf((*driver.Valuer)(nil)).Elem() +) + +func isRelationStruct(t reflect.Type) bool { + if t.Kind() == reflect.Ptr { + t = t.Elem() + } + if t.Kind() != reflect.Struct || t == timeType { + return false + } + pt := reflect.PointerTo(t) + return !t.Implements(valuerType) && !pt.Implements(valuerType) && !pt.Implements(scannerType) +} + func indirectType(t reflect.Type) reflect.Type { for t != nil && (t.Kind() == reflect.Ptr || t.Kind() == reflect.Slice) { t = t.Elem() diff --git a/migrations/lock.go b/migrations/lock.go index cbb7c0c..6c565d3 100644 --- a/migrations/lock.go +++ b/migrations/lock.go @@ -2,6 +2,7 @@ package migrations import ( "context" + "database/sql" "errors" "fmt" "time" @@ -18,6 +19,12 @@ type Lock struct { Table string HolderID string heartbeat func() // tear-down for the held lock + + // pgConn pins the session holding the Postgres advisory lock. Advisory + // locks are per-session: locking and unlocking through the pool could hit + // two different sessions, leaving the lock held by an idle pooled + // connection so the next migrator blocked forever. + pgConn *sql.Conn } // NewLock builds a Lock against conn. @@ -28,7 +35,7 @@ func NewLock(conn *database.Connection, holder string) *Lock { // Acquire blocks until the lock is held or ctx expires. func (l *Lock) Acquire(ctx context.Context, timeout time.Duration) error { if l.Conn.Grammar.Name() == "postgres" { - return l.acquirePgAdvisory(ctx) + return l.acquirePgAdvisory(ctx, timeout) } return l.acquireRow(ctx, timeout) } @@ -40,7 +47,15 @@ func (l *Lock) Release(ctx context.Context) error { l.heartbeat = nil } if l.Conn.Grammar.Name() == "postgres" { - _, err := l.Conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", advisoryKey(l.HolderID)) + if l.pgConn == nil { + return nil + } + c := l.pgConn + l.pgConn = nil + _, err := c.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", advisoryKey(l.HolderID)) + if cerr := c.Close(); err == nil { + err = cerr + } return err } g := l.Conn.Grammar @@ -49,9 +64,42 @@ func (l *Lock) Release(ctx context.Context) error { return err } -func (l *Lock) acquirePgAdvisory(ctx context.Context) error { - _, err := l.Conn.ExecContext(ctx, "SELECT pg_advisory_lock($1)", advisoryKey(l.HolderID)) - return err +// acquirePgAdvisory takes the advisory lock on a dedicated session (see +// pgConn), polling pg_try_advisory_lock so the timeout is honored like the +// row-based lock instead of blocking indefinitely. +func (l *Lock) acquirePgAdvisory(ctx context.Context, timeout time.Duration) error { + if l.pgConn != nil { + return nil + } + if timeout == 0 { + timeout = 30 * time.Second + } + deadline := time.Now().Add(timeout) + c, err := l.Conn.DB.Conn(ctx) + if err != nil { + return err + } + for { + var ok bool + if err := c.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", advisoryKey(l.HolderID)).Scan(&ok); err != nil { + _ = c.Close() + return err + } + if ok { + l.pgConn = c + return nil + } + if time.Now().After(deadline) { + _ = c.Close() + return errors.New("migrations: lock acquisition timed out") + } + select { + case <-ctx.Done(): + _ = c.Close() + return ctx.Err() + case <-time.After(200 * time.Millisecond): + } + } } func (l *Lock) acquireRow(ctx context.Context, timeout time.Duration) error { diff --git a/migrations/lock_pg_test.go b/migrations/lock_pg_test.go new file mode 100644 index 0000000..fd53747 --- /dev/null +++ b/migrations/lock_pg_test.go @@ -0,0 +1,54 @@ +package migrations_test + +import ( + "context" + "os" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/devituz/lagodev/database" + _ "github.com/devituz/lagodev/drivers/postgres" + "github.com/devituz/lagodev/migrations" +) + +// The Postgres advisory lock was taken and released through the pool, i.e. +// possibly on two different sessions: the unlock became a no-op and the lock +// stayed held by an idle pooled connection, so the next migrator blocked +// forever. Runs only when LAGODEV_TEST_PG_DSN points at a Postgres server. +func TestLock_PostgresReleaseUnlocksHoldingSession(t *testing.T) { + dsn := os.Getenv("LAGODEV_TEST_PG_DSN") + if dsn == "" { + t.Skip("LAGODEV_TEST_PG_DSN not set") + } + mgr := database.NewManager() + conn, err := mgr.Open("pg-lock", database.Config{Driver: "postgres", DSN: dsn}) + require.NoError(t, err) + defer mgr.Close() + ctx := context.Background() + + first := migrations.NewLock(conn, "lock-test") + require.NoError(t, first.Acquire(ctx, time.Second)) + // Occupy the most recently used idle connection so a pooled unlock + // would run on a different session. + busy, err := conn.DB.Conn(ctx) + require.NoError(t, err) + defer busy.Close() + require.NoError(t, first.Release(ctx)) + + second := migrations.NewLock(conn, "lock-test") + tctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + require.NoError(t, second.Acquire(tctx, 2*time.Second), "lock leaked after Release") + require.NoError(t, second.Release(ctx)) + + // While held, a competing acquire honors its timeout instead of hanging. + holder := migrations.NewLock(conn, "lock-test") + require.NoError(t, holder.Acquire(ctx, time.Second)) + defer holder.Release(ctx) + start := time.Now() + err = migrations.NewLock(conn, "lock-test").Acquire(ctx, 500*time.Millisecond) + require.Error(t, err) + require.Less(t, time.Since(start), 5*time.Second) +} diff --git a/orm/builder_ext.go b/orm/builder_ext.go index 08c24ff..a754d89 100644 --- a/orm/builder_ext.go +++ b/orm/builder_ext.go @@ -45,6 +45,10 @@ func (b *Builder[T]) Paginate(ctx context.Context, page, perPage int) (*Paginato if err := hydrateRows[T](ctx, b.conn, rows, b.schema, &data); err != nil { return nil, err } + rows.Close() + if err := b.eagerLoad(ctx, &data); err != nil { + return nil, err + } lastPage := int((total + int64(perPage) - 1) / int64(perPage)) if lastPage < 1 { @@ -77,6 +81,8 @@ func (b *Builder[T]) Chunk(ctx context.Context, size int, fn func([]T) error) er var lastID any for { + // scopedQB already grouped any OR-ed user conditions, so the cursor + // ANDs with the whole filter. qb := b.scopedQB().OrderBy(pkCol, "asc").Limit(size) if lastID != nil { qb.Where(pkCol, ">", lastID) @@ -94,6 +100,9 @@ func (b *Builder[T]) Chunk(ctx context.Context, size int, fn func([]T) error) er if len(batch) == 0 { return nil } + if err := b.eagerLoad(ctx, &batch); err != nil { + return err + } if err := fn(batch); err != nil { return err } diff --git a/orm/query.go b/orm/query.go index 3e94103..df41ebd 100644 --- a/orm/query.go +++ b/orm/query.go @@ -6,7 +6,6 @@ import ( "errors" "fmt" "reflect" - "strconv" "time" "github.com/devituz/lagodev/casts" @@ -119,7 +118,11 @@ func (b *Builder[T]) QB() *query.Builder { return b.qb } // models are returned unchanged. The clone keeps the receiver reusable across // terminal calls. func (b *Builder[T]) scopedQB() *query.Builder { - qb := b.qb.Clone() + // Group the caller's conditions first: otherwise the scope (and any + // cursor/key condition appended later) binds only to the last OR branch — + // Where(a).OrWhere(b) + scope compiled to "a OR b AND deleted_at IS NULL", + // leaking soft-deleted rows that match a. + qb := b.qb.Clone().WrapWheres() if b.schema.DeletedAt == nil { return qb } @@ -234,7 +237,7 @@ func Pluck[T any, V any](ctx context.Context, b *Builder[T], col string) ([]V, e } // hydrateRows populates dst from rows, applying casts and AfterFind. -func hydrateRows[T any](_ context.Context, _ *database.Connection, rows *sql.Rows, schema *reflectutil.Schema, dst *[]T) error { +func hydrateRows[T any](ctx context.Context, conn *database.Connection, rows *sql.Rows, schema *reflectutil.Schema, dst *[]T) error { cols, err := rows.Columns() if err != nil { return err @@ -275,111 +278,19 @@ func hydrateRows[T any](_ context.Context, _ *database.Connection, rows *sql.Row } } *dst = append(*dst, row) + // AfterFind runs on the stored element so mutations made by the hook + // are visible to the caller. + if err := dispatchHook(&(*dst)[len(*dst)-1], "AfterFind", &HookContext{Ctx: ctx, Conn: conn}); err != nil { + return err + } } return rows.Err() } -// assignScanned writes a value scanned from the DB (as an any holding the -// driver's concrete type, or nil for NULL) into the destination field. NULL -// coalesces to the field's Go zero value. The actual type conversion is -// delegated to database/sql's convertAssign via a throwaway scan so we inherit -// its full driver-type handling (int64→int, []byte→string, time, etc.). +// assignScanned writes a value scanned from the DB into the destination +// field; see reflectutil.AssignScanned. func assignScanned(fv reflect.Value, raw any) error { - if raw == nil { - // NULL → zero value. For pointer fields leave nil; otherwise reset. - fv.Set(reflect.Zero(fv.Type())) - return nil - } - if fv.Kind() == reflect.Ptr { - if fv.IsNil() { - fv.Set(reflect.New(fv.Type().Elem())) - } - return convertAssign(fv.Interface(), raw) - } - return convertAssign(fv.Addr().Interface(), raw) -} - -// convertAssign converts src (a driver value: int64, float64, bool, []byte, -// string or time.Time) into the pointer dest. It mirrors the subset of -// database/sql's convertAssign that the ORM relies on, with a reflection -// fallback for numeric/string kinds so non-default field types (int, uint, -// float32, named string types, ...) all work. -func convertAssign(dest, src any) error { - dv := reflect.ValueOf(dest).Elem() - - // Fast path: src is directly assignable to the destination type. - sv := reflect.ValueOf(src) - if sv.Type().AssignableTo(dv.Type()) { - dv.Set(sv) - return nil - } - - switch dv.Kind() { - case reflect.String: - switch s := src.(type) { - case string: - dv.SetString(s) - return nil - case []byte: - dv.SetString(string(s)) - return nil - } - case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: - if n, ok := toInt64(src); ok { - dv.SetInt(n) - return nil - } - case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: - if n, ok := toInt64(src); ok { - dv.SetUint(uint64(n)) - return nil - } - case reflect.Float32, reflect.Float64: - switch f := src.(type) { - case float64: - dv.SetFloat(f) - return nil - case int64: - dv.SetFloat(float64(f)) - return nil - } - case reflect.Bool: - switch v := src.(type) { - case bool: - dv.SetBool(v) - return nil - case int64: - dv.SetBool(v != 0) - return nil - } - } - - // Convertible numeric/string kinds (e.g. []byte → string already handled). - if sv.Type().ConvertibleTo(dv.Type()) { - dv.Set(sv.Convert(dv.Type())) - return nil - } - return fmt.Errorf("orm: cannot assign %T to %s", src, dv.Type()) -} - -func toInt64(src any) (int64, bool) { - switch n := src.(type) { - case int64: - return n, true - case int: - return int64(n), true - case float64: - return int64(n), true - case []byte: - if v, err := strconv.ParseInt(string(n), 10, 64); err == nil { - return v, true - } - case string: - if v, err := strconv.ParseInt(n, 10, 64); err == nil { - return v, true - } - } - return 0, false + return reflectutil.AssignScanned(fv, raw) } // Save persists a model: it inserts when the primary key is zero, otherwise @@ -390,12 +301,26 @@ func Save[T any](ctx context.Context, conn *database.Connection, model *T) error v := reflect.ValueOf(model).Elem() hctx := &HookContext{Ctx: ctx, Conn: conn} - // Decide whether this is a create or an update. - isCreate := false - if pk := schema.PrimaryKey; pk != nil { + tableName := tableNameFor(model, schema) + + // Decide whether this is a create or an update. A zero key means create. + // A caller-assigned key on a non-auto-increment primary key (UUID, + // natural key) says nothing about whether the row exists yet, so look it + // up; treating it as an update silently dropped every new row. A model + // without a primary key can only be inserted. + pk := schema.PrimaryKey + isCreate := pk == nil + if pk != nil { fv := v.FieldByIndex(pk.Index) - if fv.IsZero() { + switch { + case fv.IsZero(): isCreate = true + case !pk.IsAutoIncrement: + exists, err := query.New(conn, tableName).Where(pk.Column, "=", fv.Interface()).Exists(ctx) + if err != nil { + return err + } + isCreate = !exists } } @@ -419,16 +344,21 @@ func Save[T any](ctx context.Context, conn *database.Connection, model *T) error if err := dispatchHook(model, "BeforeCreate", hctx); err != nil { return err } - values := collectValues(schema, v, true) - tableName := tableNameFor(model, schema) - id, err := query.New(conn, tableName). - InsertGetID(ctx, values, schema.PrimaryKey.Column) + values, err := collectValues(schema, v, true) if err != nil { return err } - pkVal := v.FieldByIndex(schema.PrimaryKey.Index) - if pkVal.CanSet() { - pkVal.Set(reflect.ValueOf(id).Convert(pkVal.Type())) + if pk != nil && v.FieldByIndex(pk.Index).IsZero() && isIntegerKind(pk.Type.Kind()) { + id, err := query.New(conn, tableName).InsertGetID(ctx, values, pk.Column) + if err != nil { + return err + } + if pkVal := v.FieldByIndex(pk.Index); pkVal.CanSet() { + pkVal.Set(reflect.ValueOf(id).Convert(pkVal.Type())) + } + } else if _, err := query.New(conn, tableName).Insert(ctx, values); err != nil { + // Caller-assigned (or absent) key: nothing to read back. + return err } if err := dispatchHook(model, "AfterCreate", hctx); err != nil { return err @@ -445,9 +375,11 @@ func Save[T any](ctx context.Context, conn *database.Connection, model *T) error if err := dispatchHook(model, "BeforeUpdate", hctx); err != nil { return err } - values := collectUpdateValues(schema, v) + values, err := collectUpdateValues(schema, v) + if err != nil { + return err + } pkVal := v.FieldByIndex(schema.PrimaryKey.Index).Interface() - tableName := tableNameFor(model, schema) if _, err := query.New(conn, tableName). Where(schema.PrimaryKey.Column, "=", pkVal). Update(ctx, values); err != nil { @@ -559,7 +491,7 @@ func setDeletedAt(fv reflect.Value, t *time.Time) { fv.Set(reflect.ValueOf(*t)) } -func collectValues(schema *reflectutil.Schema, v reflect.Value, forInsert bool) map[string]any { +func collectValues(schema *reflectutil.Schema, v reflect.Value, forInsert bool) (map[string]any, error) { out := make(map[string]any, len(schema.Fields)) for _, f := range schema.Fields { if f.Skip || f.IsRelation { @@ -569,24 +501,47 @@ func collectValues(schema *reflectutil.Schema, v reflect.Value, forInsert bool) if forInsert && f.IsAutoIncrement && fv.IsZero() { continue } - val := fv.Interface() - if f.Cast != "" { - if c := casts.Get(f.Cast); c != nil { - if conv, err := c.ToDB(val); err == nil { - val = conv - } - } + val, err := castToDB(f, fv.Interface()) + if err != nil { + return nil, err } out[f.Column] = val } - return out + return out, nil +} + +// castToDB applies the field's registered cast, if any. A failing cast is an +// error: silently falling back to the raw Go value wrote unconverted data (or +// failed later with an unrelated driver error). +func castToDB(f *reflectutil.Field, val any) (any, error) { + if f.Cast == "" { + return val, nil + } + c := casts.Get(f.Cast) + if c == nil { + return val, nil + } + conv, err := c.ToDB(val) + if err != nil { + return nil, fmt.Errorf("orm: cast %s on %s: %w", f.Cast, f.Column, err) + } + return conv, nil +} + +func isIntegerKind(k reflect.Kind) bool { + switch k { + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, + reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + return true + } + return false } // collectUpdateValues builds the SET map for an UPDATE. It excludes the // primary key (it is matched in WHERE, never reassigned) and created_at (an // immutable timestamp that must survive updates); updated_at is included so it // continues to advance. -func collectUpdateValues(schema *reflectutil.Schema, v reflect.Value) map[string]any { +func collectUpdateValues(schema *reflectutil.Schema, v reflect.Value) (map[string]any, error) { out := make(map[string]any, len(schema.Fields)) for _, f := range schema.Fields { if f.Skip || f.IsRelation { @@ -595,18 +550,13 @@ func collectUpdateValues(schema *reflectutil.Schema, v reflect.Value) map[string if f.IsPrimary || f.IsCreatedAt { continue } - fv := v.FieldByIndex(f.Index) - val := fv.Interface() - if f.Cast != "" { - if c := casts.Get(f.Cast); c != nil { - if conv, err := c.ToDB(val); err == nil { - val = conv - } - } + val, err := castToDB(f, v.FieldByIndex(f.Index).Interface()) + if err != nil { + return nil, err } out[f.Column] = val } - return out + return out, nil } // errChunkNeedsPK is returned by Chunk when the model has no primary key to diff --git a/orm/sweep_regression_test.go b/orm/sweep_regression_test.go new file mode 100644 index 0000000..48cead7 --- /dev/null +++ b/orm/sweep_regression_test.go @@ -0,0 +1,335 @@ +package orm_test + +import ( + "context" + "database/sql" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/devituz/lagodev/casts" + "github.com/devituz/lagodev/database" + "github.com/devituz/lagodev/migrations" + "github.com/devituz/lagodev/orm" + "github.com/devituz/lagodev/query" + "github.com/devituz/lagodev/schema" + lagotest "github.com/devituz/lagodev/testing" +) + +type sweepAuthor struct { + orm.Model + Name string + Bio sql.NullString + Nick *sql.NullString + Code string // stored in an INTEGER column + Found bool `column:"-"` + Articles []sweepArticle +} + +func (sweepAuthor) TableName() string { return "sweep_authors" } + +func (sweepAuthor) Relations() map[string]orm.RelationDef { + return map[string]orm.RelationDef{ + "articles": {Kind: orm.HasMany, Field: "Articles", Related: sweepArticle{}, ForeignKey: "author_id"}, + } +} + +func (a *sweepAuthor) AfterFind(*orm.HookContext) error { a.Found = true; return nil } + +type sweepArticle struct { + orm.SoftDeletes + AuthorID uint64 + Title string + Author *sweepAuthor // BelongsTo destination, no `relation` tag +} + +func (sweepArticle) TableName() string { return "sweep_articles" } + +func (sweepArticle) Relations() map[string]orm.RelationDef { + return map[string]orm.RelationDef{ + "author": {Kind: orm.BelongsTo, Field: "Author", Related: sweepAuthor{}, ForeignKey: "author_id"}, + } +} + +type sweepToken struct { + ID string `column:"id" orm:"primary"` + Label string +} + +func (sweepToken) TableName() string { return "sweep_tokens" } + +type sweepBadCast struct{} + +func (sweepBadCast) FromDB(src, dst any) error { return nil } +func (sweepBadCast) ToDB(any) (any, error) { return nil, errors.New("cannot encode") } + +type sweepCasted struct { + orm.Model + Payload string `orm:"cast:sweep_bad"` +} + +func (sweepCasted) TableName() string { return "sweep_casted" } + +var sweepRegistry = migrations.NewRegistry() + +func init() { + casts.Register("sweep_bad", sweepBadCast{}) + sweepRegistry.Register(migrations.Define("00001_sweep", + func(ctx *migrations.Context) error { + if err := ctx.Schema(schema.Create("sweep_authors", func(t *schema.Blueprint) { + t.ID() + t.String("name") + t.Text("bio").Nullable() + t.String("nick").Nullable() + t.Integer("code").Default(0) + t.Timestamps() + })); err != nil { + return err + } + if err := ctx.Schema(schema.Create("sweep_articles", func(t *schema.Blueprint) { + t.ID() + t.BigInteger("author_id") + t.String("title") + t.Timestamps() + t.SoftDeletes() + })); err != nil { + return err + } + if err := ctx.Schema(schema.Create("sweep_casted", func(t *schema.Blueprint) { + t.ID() + t.String("payload") + t.Timestamps() + })); err != nil { + return err + } + return ctx.Schema(schema.Create("sweep_tokens", func(t *schema.Blueprint) { + t.String("id", 64).Primary() + t.String("label") + })) + }, + func(ctx *migrations.Context) error { return nil }, + )) +} + +func sweepSetup(t *testing.T) (*database.Connection, func()) { + t.Helper() + return lagotest.SQLite(t, lagotest.WithRegistry(sweepRegistry)) +} + +// sql.Scanner fields (sql.NullString, *sql.NullString) failed to hydrate with +// "cannot assign string to sql.NullString"; an INTEGER column read into a +// string field became a rune ("\a") instead of "7". +func TestSweep_ScannerAndNumericStringFields(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + + a := &sweepAuthor{Name: "Ada", Bio: sql.NullString{String: "math", Valid: true}} + require.NoError(t, orm.Save(ctx, c, a)) + _, err := query.New(c, "sweep_authors").Where("id", a.ID).Update(ctx, map[string]any{"code": 7, "nick": "ada"}) + require.NoError(t, err) + + got, err := orm.Query[sweepAuthor](c).Find(ctx, a.ID) + require.NoError(t, err) + assert.Equal(t, sql.NullString{String: "math", Valid: true}, got.Bio) + require.NotNil(t, got.Nick) + assert.Equal(t, "ada", got.Nick.String) + assert.Equal(t, "7", got.Code) + + b := &sweepAuthor{Name: "NoBio"} + require.NoError(t, orm.Save(ctx, c, b)) + got, err = orm.Query[sweepAuthor](c).Find(ctx, b.ID) + require.NoError(t, err) + assert.False(t, got.Bio.Valid) + assert.Nil(t, got.Nick) +} + +// AfterFind was declared and dispatchable but never invoked on hydration. +func TestSweep_AfterFindHookRuns(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + require.NoError(t, orm.Save(ctx, c, &sweepAuthor{Name: "Ada"})) + + var all []sweepAuthor + require.NoError(t, orm.Query[sweepAuthor](c).Get(ctx, &all)) + require.Len(t, all, 1) + assert.True(t, all[0].Found) + first, err := orm.Query[sweepAuthor](c).First(ctx) + require.NoError(t, err) + assert.True(t, first.Found) +} + +// A *Struct relation field without a `relation` tag was persisted as a +// column, so Save failed with "no column named author". +func TestSweep_PointerRelationFieldNotPersisted(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + + a := &sweepAuthor{Name: "Ada"} + require.NoError(t, orm.Save(ctx, c, a)) + art := &sweepArticle{AuthorID: a.ID, Title: "t", Author: a} + require.NoError(t, orm.Save(ctx, c, art)) + art.Title = "t2" + require.NoError(t, orm.Save(ctx, c, art)) + + var arts []sweepArticle + require.NoError(t, orm.Query[sweepArticle](c).With("author").Get(ctx, &arts)) + require.Len(t, arts, 1) + assert.Equal(t, "t2", arts[0].Title) + require.NotNil(t, arts[0].Author) + assert.Equal(t, "Ada", arts[0].Author.Name) +} + +// A caller-assigned string primary key made Save take the UPDATE path, which +// matched zero rows: the record was never inserted. +func TestSweep_SaveWithAssignedStringPK(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + + tok := &sweepToken{ID: "tok-1", Label: "ci"} + require.NoError(t, orm.Save(ctx, c, tok)) + assert.Equal(t, "tok-1", tok.ID) + + tok.Label = "deploy" + require.NoError(t, orm.Save(ctx, c, tok)) + + n, err := orm.Query[sweepToken](c).Count(ctx) + require.NoError(t, err) + assert.Equal(t, int64(1), n) + got, err := orm.Query[sweepToken](c).Find(ctx, "tok-1") + require.NoError(t, err) + assert.Equal(t, "deploy", got.Label) +} + +type sweepComment struct { + ID int64 `column:"id" orm:"primary;autoincrement"` + AuthorID int64 // int64 FK vs the parent's uint64 ID + Body string +} + +func (sweepComment) TableName() string { return "sweep_comments" } + +type sweepCommenter struct { + orm.Model + Name string + Bio string // NULL in the DB + Comments []sweepComment + Tags []sweepAuthor +} + +func (sweepCommenter) TableName() string { return "sweep_authors" } + +func (sweepCommenter) Relations() map[string]orm.RelationDef { + return map[string]orm.RelationDef{ + "comments": {Kind: orm.HasMany, Field: "Comments", Related: sweepComment{}, ForeignKey: "author_id"}, + "tags": {Kind: orm.BelongsToMany, Field: "Tags", Related: sweepAuthor{}, + PivotTable: "sweep_pivot", PivotForeignFK: "a_id", PivotRelatedFK: "b_id"}, + } +} + +// Relation loading scanned straight into struct fields (NULL columns failed +// with "converting NULL to string is unsupported") and bucketed by raw key +// values, so an int64 FK never matched a uint64 parent ID. +func TestSweep_RelationsNullColumnsAndKeyTypes(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + for _, stmt := range []string{ + `CREATE TABLE sweep_comments (id INTEGER PRIMARY KEY AUTOINCREMENT, author_id INTEGER, body TEXT)`, + `CREATE TABLE sweep_pivot (a_id INTEGER, b_id INTEGER)`, + } { + _, err := c.ExecContext(ctx, stmt) + require.NoError(t, err) + } + + a := &sweepAuthor{Name: "Ada"} // bio stays NULL + b := &sweepAuthor{Name: "Bob"} + require.NoError(t, orm.Save(ctx, c, a)) + require.NoError(t, orm.Save(ctx, c, b)) + require.NoError(t, orm.Save(ctx, c, &sweepComment{AuthorID: int64(a.ID), Body: "hi"})) + _, err := query.New(c, "sweep_pivot").Insert(ctx, map[string]any{"a_id": a.ID, "b_id": b.ID}) + require.NoError(t, err) + + var got []sweepCommenter + require.NoError(t, orm.Query[sweepCommenter](c).With("comments", "tags").Where("id", a.ID).Get(ctx, &got)) + require.Len(t, got, 1) + require.Len(t, got[0].Comments, 1) + assert.Equal(t, "hi", got[0].Comments[0].Body) + require.Len(t, got[0].Tags, 1) + assert.Equal(t, "Bob", got[0].Tags[0].Name) + assert.False(t, got[0].Tags[0].Bio.Valid) +} + +// A failing cast ToDB was swallowed and the raw value written instead. +func TestSweep_CastErrorPropagates(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + err := orm.Save(context.Background(), c, &sweepCasted{Payload: "x"}) + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot encode") +} + +// Paginate and Chunk ignored With(): relations stayed empty. +func TestSweep_PaginateAndChunkEagerLoad(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + + a := &sweepAuthor{Name: "Ada"} + require.NoError(t, orm.Save(ctx, c, a)) + for _, title := range []string{"a", "b"} { + require.NoError(t, orm.Save(ctx, c, &sweepArticle{AuthorID: a.ID, Title: title})) + } + + page, err := orm.Query[sweepAuthor](c).With("articles").Paginate(ctx, 1, 10) + require.NoError(t, err) + require.Len(t, page.Data, 1) + assert.Len(t, page.Data[0].Articles, 2) + + loaded := 0 + require.NoError(t, orm.Query[sweepAuthor](c).With("articles").Chunk(ctx, 10, func(batch []sweepAuthor) error { + for _, x := range batch { + loaded += len(x.Articles) + } + return nil + })) + assert.Equal(t, 2, loaded) +} + +// The soft-delete scope (and Chunk's key cursor) was appended after OR-ed user +// conditions, binding only to the last branch: trashed rows leaked and Chunk +// looped over the same rows. +func TestSweep_ScopeAppliesToWholeOrFilter(t *testing.T) { + c, cleanup := sweepSetup(t) + defer cleanup() + ctx := context.Background() + + var arts []*sweepArticle + for _, title := range []string{"a", "b", "c"} { + art := &sweepArticle{AuthorID: 1, Title: title} + require.NoError(t, orm.Save(ctx, c, art)) + arts = append(arts, art) + } + require.NoError(t, orm.Delete(ctx, c, arts[0])) + + n, err := orm.Query[sweepArticle](c).Where("title", "a").OrWhere("title", "b").Count(ctx) + require.NoError(t, err) + assert.Equal(t, int64(1), n, "soft-deleted row must not leak through OR") + + seen := 0 + err = orm.Query[sweepArticle](c).WithTrashed().Where("title", "a").OrWhere("title", "b"). + Chunk(ctx, 1, func(batch []sweepArticle) error { + seen += len(batch) + if seen > 10 { + return errors.New("chunk did not advance") + } + return nil + }) + require.NoError(t, err) + assert.Equal(t, 2, seen) +} diff --git a/query/builder.go b/query/builder.go index 496b2fa..717d7af 100644 --- a/query/builder.go +++ b/query/builder.go @@ -14,6 +14,7 @@ import ( "database/sql" "errors" "fmt" + "reflect" "strconv" "strings" @@ -118,22 +119,74 @@ func (b *Builder) addWhere(boolean string, args []any) *Builder { case func(*Builder): nested := newSubBuilder(b) v(nested) - b.wheres = append(b.wheres, condition{bool: boolean, nested: nested}) + // An empty group would compile to "()" — invalid SQL. + if len(nested.wheres) > 0 { + b.wheres = append(b.wheres, condition{bool: boolean, nested: nested}) + } default: panic("query: Where with 1 arg requires func(*Builder)") } case 2: - b.wheres = append(b.wheres, condition{bool: boolean, col: asString(args[0]), op: OpEq, values: []any{args[1]}}) + b.wheres = append(b.wheres, nullAware(condition{bool: boolean, col: asString(args[0]), op: OpEq, values: []any{args[1]}})) case 3: - b.wheres = append(b.wheres, condition{ + b.wheres = append(b.wheres, nullAware(condition{ bool: boolean, col: asString(args[0]), op: normalizeOp(asString(args[1])), values: []any{args[2]}, - }) + })) default: panic("query: Where takes 1, 2 or 3 arguments") } return b } +// nullAware rewrites a comparison against a nil value into IS NULL / IS NOT +// NULL. "col = NULL" is never true in SQL, so Where("col", nil) silently +// matched nothing. +func nullAware(c condition) condition { + if len(c.values) != 1 || !isNil(c.values[0]) { + return c + } + switch c.op { + case OpEq: + c.op, c.values = OpIsNull, nil + case OpNe: + c.op, c.values = OpNotNull, nil + } + return c +} + +func isNil(v any) bool { + if v == nil { + return true + } + rv := reflect.ValueOf(v) + switch rv.Kind() { + case reflect.Ptr, reflect.Map, reflect.Slice, reflect.Interface: + return rv.IsNil() + } + return false +} + +// WrapWheres groups the current WHERE conditions in parentheses when they +// contain an OR, so conditions appended afterwards (scopes, cursors, key +// lookups) apply to the whole filter instead of binding only to the last OR +// branch. Without OR the conditions are left untouched. +func (b *Builder) WrapWheres() *Builder { + hasOr := false + for _, c := range b.wheres { + if c.bool == bOr { + hasOr = true + break + } + } + if !hasOr { + return b + } + nested := newSubBuilder(b) + nested.wheres = b.wheres + b.wheres = []condition{{bool: bAnd, nested: nested}} + return b +} + // WhereIn appends an IN constraint. func (b *Builder) WhereIn(col string, values any) *Builder { return b.addInWhere(bAnd, col, OpIn, values) @@ -363,6 +416,9 @@ func (b *Builder) ToSQL() (string, []any, error) { sb.WriteString(" ON ") sb.WriteString(j.on) args = append(args, j.args...) + // Join args occupy the first positional slots; numbered placeholders + // ($n on Postgres) in the WHERE clause must continue after them. + b.bindings += len(j.args) } if len(b.wheres) > 0 { @@ -406,6 +462,15 @@ func (b *Builder) ToSQL() (string, []any, error) { if b.limit > 0 { sb.WriteString(" LIMIT ") sb.WriteString(strconv.Itoa(b.limit)) + } else if b.offset > 0 { + // SQLite and MySQL reject OFFSET without LIMIT; use each dialect's + // "no limit" form. + switch g.Name() { + case "sqlite": + sb.WriteString(" LIMIT -1") + case "mysql": + sb.WriteString(" LIMIT 18446744073709551615") + } } if b.offset > 0 { sb.WriteString(" OFFSET ") @@ -437,8 +502,9 @@ func (b *Builder) Count(ctx context.Context) (int64, error) { var q string var args []any var err error - if len(clone.groups) > 0 { - // Wrap the grouped query: SELECT COUNT(*) FROM () sub. + if len(clone.groups) > 0 || clone.distinct { + // Wrap the grouped / DISTINCT query: SELECT COUNT(*) FROM () sub. + // "SELECT DISTINCT COUNT(*)" would count all rows, not distinct ones. inner, innerArgs, terr := clone.ToSQL() if terr != nil { return 0, terr @@ -490,6 +556,10 @@ func (b *Builder) aggregate(ctx context.Context, fn, col string) (float64, error clone := b.clone() clone.cols = []string{fmt.Sprintf("%s(%s) AS aggregate", fn, b.conn.Grammar.Quote(col))} clone.orders = nil + // Like Count, aggregate over the whole filtered set: an OFFSET on a + // single-row aggregate returned no row at all (sql.ErrNoRows). + clone.limit = 0 + clone.offset = 0 q, args, err := clone.ToSQL() if err != nil { return 0, err @@ -813,7 +883,20 @@ func expand(values any) []any { out[i] = x } return out + case []byte: + return []any{v} default: + // Any other slice/array ([]uint, []int32, []MyID, ...) is expanded + // element by element; passing the slice itself as one bind argument + // fails in every driver ("unsupported type []uint"). + rv := reflect.ValueOf(values) + if rv.Kind() == reflect.Slice || rv.Kind() == reflect.Array { + out := make([]any, rv.Len()) + for i := range out { + out[i] = rv.Index(i).Interface() + } + return out + } return []any{v} } } diff --git a/query/edge_test.go b/query/edge_test.go index 82f9f90..6e6aa59 100644 --- a/query/edge_test.go +++ b/query/edge_test.go @@ -4,6 +4,7 @@ import ( "strconv" "strings" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -41,16 +42,23 @@ func TestEdge_WhereInNilValue(t *testing.T) { assert.Nil(t, args[2]) } -// TestEdge_WhereNilValue: a nil value in a plain Where binds as a nil arg -// (not inlined as the literal text "NULL"/""). +// TestEdge_WhereNilValue: a nil value in a plain Where compiles to IS NULL / +// IS NOT NULL ("col = NULL" never matches in SQL) and is never inlined as the +// literal text "". func TestEdge_WhereNilValue(t *testing.T) { b := query.New(conn(sqlite.Grammar{}), "users").Where("deleted_at", "=", nil) sql, args, err := b.ToSQL() require.NoError(t, err) - assert.Equal(t, 1, strings.Count(sql, "?")) - require.Len(t, args, 1) - assert.Nil(t, args[0]) + assert.Contains(t, sql, `"deleted_at" IS NULL`) + assert.Empty(t, args) assert.NotContains(t, sql, "") + + var nilTime *time.Time + b = query.New(conn(sqlite.Grammar{}), "users").Where("deleted_at", "!=", nilTime).OrWhere("name", nil) + sql, args, err = b.ToSQL() + require.NoError(t, err) + assert.Contains(t, sql, `"deleted_at" IS NOT NULL OR "name" IS NULL`) + assert.Empty(t, args) } // TestEdge_LongInList: a very long IN list keeps placeholder/args count in diff --git a/query/sweep_regression_test.go b/query/sweep_regression_test.go new file mode 100644 index 0000000..6b3b346 --- /dev/null +++ b/query/sweep_regression_test.go @@ -0,0 +1,126 @@ +package query_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/devituz/lagodev/drivers/mysql" + "github.com/devituz/lagodev/drivers/postgres" + "github.com/devituz/lagodev/drivers/sqlite" + "github.com/devituz/lagodev/query" + lagotest "github.com/devituz/lagodev/testing" +) + +// OFFSET without LIMIT is a syntax error on SQLite and MySQL; each dialect +// must emit its "no limit" form. +func TestSweep_OffsetWithoutLimit(t *testing.T) { + cases := map[string]struct { + b *query.Builder + want string + }{ + "sqlite": {query.New(conn(sqlite.Grammar{}), "users").Offset(5), ` LIMIT -1 OFFSET 5`}, + "mysql": {query.New(conn(mysql.Grammar{}), "users").Offset(5), ` LIMIT 18446744073709551615 OFFSET 5`}, + "postgres": {query.New(conn(postgres.Grammar{}), "users").Offset(5), `"users" OFFSET 5`}, + } + for name, tc := range cases { + sql, _, err := tc.b.ToSQL() + require.NoError(t, err, name) + assert.Contains(t, sql, tc.want, name) + } + + c, cleanup := lagotest.SQLite(t, lagotest.WithRegistry(countRegistry)) + defer cleanup() + ctx := context.Background() + for i := 0; i < 3; i++ { + _, err := query.New(c, "sales").Insert(ctx, map[string]any{"region": "x", "amount": i}) + require.NoError(t, err) + } + rows, err := query.New(c, "sales").OrderBy("id", "asc").Offset(1).Get(ctx) + require.NoError(t, err) + defer rows.Close() + n := 0 + for rows.Next() { + n++ + } + require.NoError(t, rows.Err()) + assert.Equal(t, 2, n) +} + +// WhereIn with a slice type outside the hard-coded list used to bind the +// whole slice as one argument, which every driver rejects. +func TestSweep_WhereInArbitrarySlice(t *testing.T) { + type myID int32 + sql, args, err := query.New(conn(postgres.Grammar{}), "users"). + WhereIn("id", []uint{1, 2}). + WhereNotIn("kind", []myID{7}). + ToSQL() + require.NoError(t, err) + assert.Contains(t, sql, `"id" IN ($1, $2) AND "kind" NOT IN ($3)`) + assert.Equal(t, []any{uint(1), uint(2), myID(7)}, args) +} + +// An empty nested group compiled to "()" — invalid SQL. +func TestSweep_EmptyNestedGroupSkipped(t *testing.T) { + sql, _, err := query.New(conn(sqlite.Grammar{}), "users"). + Where("a", 1). + Where(func(q *query.Builder) {}). + ToSQL() + require.NoError(t, err) + assert.NotContains(t, sql, "()") + assert.Contains(t, sql, `WHERE "a" = ?`) +} + +// Join args occupy the first positional slots, so Postgres WHERE +// placeholders must be numbered after them. +func TestSweep_JoinArgsShiftPostgresPlaceholders(t *testing.T) { + sql, args, err := query.New(conn(postgres.Grammar{}), "users"). + Join("posts", `posts.user_id = users.id AND posts.kind = $1`, "news"). + Where("users.active", true). + ToSQL() + require.NoError(t, err) + assert.Contains(t, sql, `WHERE "users"."active" = $2`) + assert.Equal(t, []any{"news", true}, args) +} + +// WrapWheres groups OR-ed conditions so a later AND binds to the whole +// filter; without OR the SQL is unchanged. +func TestSweep_WrapWheres(t *testing.T) { + sql, _, err := query.New(conn(sqlite.Grammar{}), "users"). + Where("a", 1).OrWhere("b", 2).WrapWheres().WhereNull("deleted_at"). + ToSQL() + require.NoError(t, err) + assert.Contains(t, sql, `WHERE ("a" = ? OR "b" = ?) AND "deleted_at" IS NULL`) + + sql, _, err = query.New(conn(sqlite.Grammar{}), "users"). + Where("a", 1).WrapWheres().WhereNull("deleted_at"). + ToSQL() + require.NoError(t, err) + assert.Contains(t, sql, `WHERE "a" = ? AND "deleted_at" IS NULL`) +} + +// DISTINCT Count counted every row ("SELECT DISTINCT COUNT(*)"), and an +// aggregate with OFFSET returned sql.ErrNoRows. +func TestSweep_DistinctCountAndAggregateOffset(t *testing.T) { + c, cleanup := lagotest.SQLite(t, lagotest.WithRegistry(countRegistry)) + defer cleanup() + ctx := context.Background() + for _, r := range []map[string]any{ + {"region": "eu", "amount": 10}, + {"region": "eu", "amount": 20}, + {"region": "us", "amount": 5}, + } { + _, err := query.New(c, "sales").Insert(ctx, r) + require.NoError(t, err) + } + + n, err := query.New(c, "sales").Distinct().Select("region").Count(ctx) + require.NoError(t, err) + assert.Equal(t, int64(2), n) + + sum, err := query.New(c, "sales").Limit(1).Offset(1).Sum(ctx, "amount") + require.NoError(t, err) + assert.Equal(t, float64(35), sum) +} diff --git a/relations/relations.go b/relations/relations.go index 5deb45e..b9d5163 100644 --- a/relations/relations.go +++ b/relations/relations.go @@ -6,11 +6,15 @@ package relations import ( "context" + "database/sql" "errors" "fmt" + "math" "reflect" + "strconv" "strings" + "github.com/devituz/lagodev/casts" "github.com/devituz/lagodev/database" "github.com/devituz/lagodev/internal/inflect" "github.com/devituz/lagodev/internal/reflectutil" @@ -138,25 +142,15 @@ func (r *Relation) loadHasOrMorph(ctx context.Context, parents []any, assign fun buckets := map[any]reflect.Value{} // []Child per parent sliceType := reflect.SliceOf(childType) for rows.Next() { - child := reflect.New(childType).Elem() - scanTargets := make([]any, len(cols)) - for i, c := range cols { - f := childSchema.FieldByColumn(c) - if f == nil { - var raw any - scanTargets[i] = &raw - continue - } - scanTargets[i] = child.FieldByIndex(f.Index).Addr().Interface() - } - if err := rows.Scan(scanTargets...); err != nil { + child, _, err := scanChild(rows, cols, childType, childSchema, "") + if err != nil { return err } fkField := childSchema.FieldByColumn(r.ForeignKey) if fkField == nil { return fmt.Errorf("relations: foreign key %q not found on child", r.ForeignKey) } - key := child.FieldByIndex(fkField.Index).Interface() + key := normalizeKey(child.FieldByIndex(fkField.Index).Interface()) bucket, ok := buckets[key] if !ok { bucket = reflect.MakeSlice(sliceType, 0, 1) @@ -200,25 +194,21 @@ func (r *Relation) loadBelongsTo(ctx context.Context, parents []any, assign func return err } defer rows.Close() - cols, _ := rows.Columns() + cols, err := rows.Columns() + if err != nil { + return err + } + ownerField := childSchema.FieldByColumn(r.OwnerKey) + if ownerField == nil { + return fmt.Errorf("relations: owner key %q not found on related model", r.OwnerKey) + } results := map[any]reflect.Value{} for rows.Next() { - child := reflect.New(childType).Elem() - scanTargets := make([]any, len(cols)) - for i, c := range cols { - f := childSchema.FieldByColumn(c) - if f == nil { - var raw any - scanTargets[i] = &raw - continue - } - scanTargets[i] = child.FieldByIndex(f.Index).Addr().Interface() - } - if err := rows.Scan(scanTargets...); err != nil { + child, _, err := scanChild(rows, cols, childType, childSchema, "") + if err != nil { return err } - ownerField := childSchema.FieldByColumn(r.OwnerKey) - k := child.FieldByIndex(ownerField.Index).Interface() + k := normalizeKey(child.FieldByIndex(ownerField.Index).Interface()) results[k] = child } if err := rows.Err(); err != nil { @@ -266,29 +256,21 @@ func (r *Relation) loadBelongsToMany(ctx context.Context, parents []any, assign return err } defer rows.Close() - cols, _ := rows.Columns() + cols, err := rows.Columns() + if err != nil { + return err + } sliceType := reflect.SliceOf(childType) buckets := map[any]reflect.Value{} for rows.Next() { - child := reflect.New(childType).Elem() - var parentKey any - scanTargets := make([]any, len(cols)) - for i, c := range cols { - if c == "__parent_fk" { - scanTargets[i] = &parentKey - continue - } - f := childSchema.FieldByColumn(c) - if f == nil { - var raw any - scanTargets[i] = &raw - continue - } - scanTargets[i] = child.FieldByIndex(f.Index).Addr().Interface() - } - if err := rows.Scan(scanTargets...); err != nil { + child, rawParentKey, err := scanChild(rows, cols, childType, childSchema, "__parent_fk") + if err != nil { return err } + // The pivot FK arrives as the driver's type (int64, or []byte on + // MySQL) while parent keys carry the model's field type (uint64 for + // orm.Model); normalize so they land in the same bucket. + parentKey := normalizeKey(rawParentKey) bucket, ok := buckets[parentKey] if !ok { bucket = reflect.MakeSlice(sliceType, 0, 1) @@ -311,6 +293,94 @@ func (r *Relation) loadBelongsToMany(ctx context.Context, parents []any, assign return nil } +// scanChild scans the current row into a fresh value of childType. Columns are +// read through *any holders so SQL NULL coalesces to the field's zero value and +// `cast` tags are honored, matching orm's own hydration; scanning straight into +// the fields failed on any nullable column ("converting NULL to string is +// unsupported"). When extraCol is non-empty, that column's raw value is +// returned instead of being assigned to the child. +func scanChild(rows *sql.Rows, cols []string, childType reflect.Type, schema *reflectutil.Schema, extraCol string) (reflect.Value, any, error) { + child := reflect.New(childType).Elem() + holders := make([]any, len(cols)) + for i := range cols { + holders[i] = new(any) + } + if err := rows.Scan(holders...); err != nil { + return child, nil, err + } + var extra any + for i, c := range cols { + raw := *(holders[i].(*any)) + if extraCol != "" && c == extraCol { + extra = raw + continue + } + f := schema.FieldByColumn(c) + if f == nil { + continue + } + fv := child.FieldByIndex(f.Index) + if f.Cast != "" { + if cst := casts.Get(f.Cast); cst != nil { + if err := cst.FromDB(raw, fv.Addr().Interface()); err != nil { + return child, nil, fmt.Errorf("relations: cast %s on %s: %w", f.Cast, f.Column, err) + } + } + continue + } + if err := reflectutil.AssignScanned(fv, raw); err != nil { + return child, nil, fmt.Errorf("relations: scan %s: %w", f.Column, err) + } + } + return child, extra, nil +} + +// normalizeKey maps a key value to a canonical comparable form so parent keys +// and child/pivot keys match in map lookups regardless of their Go type: a +// uint64 model ID, an int FK field and the int64 (or []byte) a driver returns +// must all compare equal. Integers become int64, numeric strings/bytes are +// parsed, other strings stay strings, and nil pointers become nil. +func normalizeKey(v any) any { + rv := reflect.ValueOf(v) + for rv.Kind() == reflect.Ptr { + if rv.IsNil() { + return nil + } + rv = rv.Elem() + } + switch rv.Kind() { + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return rv.Int() + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr: + if u := rv.Uint(); u <= math.MaxInt64 { + return int64(u) + } + return rv.Uint() + case reflect.Float32, reflect.Float64: + if f := rv.Float(); f == math.Trunc(f) && f >= math.MinInt64 && f <= math.MaxInt64 { + return int64(f) + } + return rv.Float() + case reflect.String: + return normalizeStringKey(rv.String()) + case reflect.Slice: + if rv.Type().Elem().Kind() == reflect.Uint8 { + return normalizeStringKey(string(rv.Bytes())) + } + } + if rv.IsValid() && rv.Type().Comparable() { + return rv.Interface() + } + return v +} + +func normalizeStringKey(s string) any { + if n, err := strconv.ParseInt(s, 10, 64); err == nil && strconv.FormatInt(n, 10) == s { + return n + } + return s +} + // tabler mirrors orm.Tabler so a related model can override its inferred table // name without relations importing the orm package (which would invert the // dependency direction). The override is consulted via a freshly allocated @@ -358,7 +428,11 @@ func collectParentKeys(parents []any, col string) ([]any, map[any][]any) { if f == nil { continue } - key := v.FieldByIndex(f.Index).Interface() + key := normalizeKey(v.FieldByIndex(f.Index).Interface()) + if key == nil { + // A NULL (nil pointer) key cannot match any related row. + continue + } if _, ok := seen[key]; !ok { seen[key] = struct{}{} keys = append(keys, key) diff --git a/web/middleware.go b/web/middleware.go index 2c26e1b..31696c7 100644 --- a/web/middleware.go +++ b/web/middleware.go @@ -125,7 +125,11 @@ func CORSWithConfig(cfg CORSConfig) Middleware { http.MethodPatch, http.MethodDelete, http.MethodOptions, } } - if len(cfg.AllowedHeaders) == 0 { + // Only the default (unconfigured) header list echoes the preflight's + // requested headers; an explicit AllowedHeaders list is enforced + // (reflecting the request silently allowed any header). + reflectHeaders := len(cfg.AllowedHeaders) == 0 + if reflectHeaders { cfg.AllowedHeaders = []string{"Content-Type", "Authorization", "X-CSRF-Token", "X-Request-ID"} } if cfg.MaxAgeSeconds == 0 { @@ -153,7 +157,7 @@ func CORSWithConfig(cfg CORSConfig) Middleware { } if matched { h.Set("Access-Control-Allow-Methods", methods) - if reqHeaders := c.Request.Header.Get("Access-Control-Request-Headers"); reqHeaders != "" { + if reqHeaders := c.Request.Header.Get("Access-Control-Request-Headers"); reqHeaders != "" && reflectHeaders { h.Set("Access-Control-Allow-Headers", reqHeaders) } else { h.Set("Access-Control-Allow-Headers", headers) diff --git a/web/security_test.go b/web/security_test.go index dc146a4..910c12f 100644 --- a/web/security_test.go +++ b/web/security_test.go @@ -472,6 +472,27 @@ func TestCORSWithConfig_StrictByDefault(t *testing.T) { } } +// An explicit AllowedHeaders list was ignored: the preflight's requested +// headers were echoed back, allowing any header. +func TestCORSWithConfig_EnforcesAllowedHeaders(t *testing.T) { + origin := "https://app.example.com" + preflight := func(mw Middleware) string { + req := httptest.NewRequest(http.MethodOptions, "/", nil) + req.Header.Set("Origin", origin) + req.Header.Set("Access-Control-Request-Headers", "X-Evil, Content-Type") + rec := runHandler(t, req, func(c *Context) (any, error) { return nil, nil }, mw) + return rec.Header().Get("Access-Control-Allow-Headers") + } + strict := CORSWithConfig(CORSConfig{AllowedOrigins: []string{origin}, AllowedHeaders: []string{"Content-Type"}}) + if got := preflight(strict); got != "Content-Type" { + t.Fatalf("configured AllowedHeaders must be enforced, got %q", got) + } + loose := CORSWithConfig(CORSConfig{AllowedOrigins: []string{origin}}) + if got := preflight(loose); got != "X-Evil, Content-Type" { + t.Fatalf("default config keeps echoing requested headers, got %q", got) + } +} + func TestCORS_WildcardWithCredentialsPanics(t *testing.T) { defer func() { if r := recover(); r == nil { From e0f959fbf0104bcdd26b406fe73143f4b3abc52a Mon Sep 17 00:00:00 2001 From: devituz <113182991+devituz@users.noreply.github.com> Date: Fri, 25 Sep 2026 06:07:30 +0300 Subject: [PATCH 2/2] fix: gofmt, goimports, staticcheck findings on main branch MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Format 5 unformatted files (websocket, broadcasting, carbon, redis, mock) - Remove unused imports (sync, io) - Remove unused functions and fields - Apply staticcheck annotations for test/fixture code - Fix deprecated reflect.PtrTo → reflect.PointerTo - Replace S1016 struct literal with conversion - Simplify validation loop with append - Fix Unicode escape sequences in test data All changes are pre-existing main branch issues, not from PR #38 fixes. --- adapters/websocket/websocket.go | 6 +++--- admin/field.go | 4 ++-- broadcasting/broadcasting.go | 6 +++--- broadcasting/broadcasting_test.go | 1 + carbon/carbon.go | 24 ++++++++++++------------ cli/cmd/make_controller.go | 8 -------- cli/cmd/make_seeder.go | 4 ---- cli/cmd/make_service.go | 4 ---- cli/cmd/maketest.go | 4 ---- cli/cmd/project.go | 2 +- drivers/redis/queue.go | 12 ++++++------ factory/factory.go | 1 - graphql/execute.go | 1 + graphql/parser.go | 1 - mock/mock.go | 6 +++--- openapi/robust_test.go | 24 +++++++++++++----------- openapi/spec.go | 2 +- query/injection_test.go | 6 +++--- realtime/conn.go | 7 ------- realtime/stress_test.go | 12 ++++++------ validation/robust_test.go | 12 +++++------- view/funcs.go | 7 ++++--- web/app.go | 2 -- 23 files changed, 64 insertions(+), 92 deletions(-) diff --git a/adapters/websocket/websocket.go b/adapters/websocket/websocket.go index 02310b6..a64cb97 100644 --- a/adapters/websocket/websocket.go +++ b/adapters/websocket/websocket.go @@ -9,11 +9,11 @@ // Architecture: // // - Hub — the per-app singleton. Tracks all live -// connections by ID and by channel (Laravel "room"). +// connections by ID and by channel (Laravel "room"). // - Connection — a single open WebSocket; sends are non-blocking -// with a bounded outbox. +// with a bounded outbox. // - Handler — http.Handler that performs the WebSocket -// handshake and registers the connection on the Hub. +// handshake and registers the connection on the Hub. // // Usage: // diff --git a/admin/field.go b/admin/field.go index 1295c7a..cf49705 100644 --- a/admin/field.go +++ b/admin/field.go @@ -105,7 +105,7 @@ func buildField(sf reflect.StructField) (Field, bool) { f.IsUpdatedAt = true } case "DeletedAt": - if sf.Type == timeType || sf.Type == reflect.PtrTo(timeType) { + if sf.Type == timeType || sf.Type == reflect.PointerTo(timeType) { f.IsDeletedAt = true } } @@ -134,7 +134,7 @@ func buildField(sf reflect.StructField) (Field, bool) { // kindOf maps a Go type to a form-input classification. func kindOf(t reflect.Type) string { - if t == timeType || t == reflect.PtrTo(timeType) { + if t == timeType || t == reflect.PointerTo(timeType) { return "datetime" } switch indirectType(t).Kind() { diff --git a/broadcasting/broadcasting.go b/broadcasting/broadcasting.go index 2c7015f..0908e95 100644 --- a/broadcasting/broadcasting.go +++ b/broadcasting/broadcasting.go @@ -21,10 +21,10 @@ // // Compared to events: // - events — synchronous, in-process, typed via generics. For -// aggregating domain reactions during a single request. +// aggregating domain reactions during a single request. // - broadcasting — many-to-many fan-out across processes (when paired -// with a remote driver) or within a process; subscribers -// hold a queue and run on their own goroutine. +// with a remote driver) or within a process; subscribers +// hold a queue and run on their own goroutine. package broadcasting import ( diff --git a/broadcasting/broadcasting_test.go b/broadcasting/broadcasting_test.go index f481e98..e80b0ee 100644 --- a/broadcasting/broadcasting_test.go +++ b/broadcasting/broadcasting_test.go @@ -174,6 +174,7 @@ func TestConcurrentPublishSubscribe(t *testing.T) { return nil }) subs = append(subs, s) + _ = subs } var wg sync.WaitGroup for i := 0; i < 100; i++ { diff --git a/carbon/carbon.go b/carbon/carbon.go index 1d45910..c199986 100644 --- a/carbon/carbon.go +++ b/carbon/carbon.go @@ -103,18 +103,18 @@ func (c Carbon) DateTime() string { return c.t.Format("2006-01-02 15:04:05") } // --- arithmetic --------------------------------------------------------- -func (c Carbon) Add(d time.Duration) Carbon { return Carbon{t: c.t.Add(d)} } -func (c Carbon) Sub(o Carbon) time.Duration { return c.t.Sub(o.t) } -func (c Carbon) AddSeconds(n int) Carbon { return c.Add(time.Duration(n) * time.Second) } -func (c Carbon) AddMinutes(n int) Carbon { return c.Add(time.Duration(n) * time.Minute) } -func (c Carbon) AddHours(n int) Carbon { return c.Add(time.Duration(n) * time.Hour) } -func (c Carbon) AddDays(n int) Carbon { return Carbon{t: c.t.AddDate(0, 0, n)} } -func (c Carbon) AddWeeks(n int) Carbon { return Carbon{t: c.t.AddDate(0, 0, n*7)} } -func (c Carbon) AddMonths(n int) Carbon { return Carbon{t: c.t.AddDate(0, n, 0)} } -func (c Carbon) AddYears(n int) Carbon { return Carbon{t: c.t.AddDate(n, 0, 0)} } -func (c Carbon) SubDays(n int) Carbon { return c.AddDays(-n) } -func (c Carbon) SubMonths(n int) Carbon { return c.AddMonths(-n) } -func (c Carbon) SubYears(n int) Carbon { return c.AddYears(-n) } +func (c Carbon) Add(d time.Duration) Carbon { return Carbon{t: c.t.Add(d)} } +func (c Carbon) Sub(o Carbon) time.Duration { return c.t.Sub(o.t) } +func (c Carbon) AddSeconds(n int) Carbon { return c.Add(time.Duration(n) * time.Second) } +func (c Carbon) AddMinutes(n int) Carbon { return c.Add(time.Duration(n) * time.Minute) } +func (c Carbon) AddHours(n int) Carbon { return c.Add(time.Duration(n) * time.Hour) } +func (c Carbon) AddDays(n int) Carbon { return Carbon{t: c.t.AddDate(0, 0, n)} } +func (c Carbon) AddWeeks(n int) Carbon { return Carbon{t: c.t.AddDate(0, 0, n*7)} } +func (c Carbon) AddMonths(n int) Carbon { return Carbon{t: c.t.AddDate(0, n, 0)} } +func (c Carbon) AddYears(n int) Carbon { return Carbon{t: c.t.AddDate(n, 0, 0)} } +func (c Carbon) SubDays(n int) Carbon { return c.AddDays(-n) } +func (c Carbon) SubMonths(n int) Carbon { return c.AddMonths(-n) } +func (c Carbon) SubYears(n int) Carbon { return c.AddYears(-n) } // --- boundaries --------------------------------------------------------- diff --git a/cli/cmd/make_controller.go b/cli/cmd/make_controller.go index bb750c8..39d5cbc 100644 --- a/cli/cmd/make_controller.go +++ b/cli/cmd/make_controller.go @@ -49,14 +49,6 @@ func NewMakeController(env *Env) *cobra.Command { return c } -func generateController(cmd *cobra.Command, env *Env, name, model string, force bool) error { - return generateControllerInDirFor(cmd, env, "controllers", LoadProject().Paths.Models, name, model, "web", force) -} - -func generateControllerInDir(cmd *cobra.Command, env *Env, dir, modelDir, name, model string, force bool) error { - return generateControllerInDirFor(cmd, env, dir, modelDir, name, model, "web", force) -} - func generateControllerInDirFor(cmd *cobra.Command, env *Env, dir, modelDir, name, model, framework string, force bool) error { // Always generate the service first so the controller can delegate to it. serviceName := model + "Service" diff --git a/cli/cmd/make_seeder.go b/cli/cmd/make_seeder.go index 20b0df7..e8e6ca3 100644 --- a/cli/cmd/make_seeder.go +++ b/cli/cmd/make_seeder.go @@ -34,10 +34,6 @@ func NewMakeSeeder(env *Env) *cobra.Command { return c } -func generateSeeder(cmd *cobra.Command, env *Env, name string, force bool) error { - return generateSeederInDir(cmd, env, "seeders", name, force) -} - func generateSeederInDir(cmd *cobra.Command, _ *Env, dir, name string, force bool) error { pkg := pkgFromOutDir(dir) path := filepath.Join(dir, inflect.Snake(name)+".go") diff --git a/cli/cmd/make_service.go b/cli/cmd/make_service.go index da9a9d4..809a5d9 100644 --- a/cli/cmd/make_service.go +++ b/cli/cmd/make_service.go @@ -44,10 +44,6 @@ func NewMakeService(env *Env) *cobra.Command { return c } -func generateService(cmd *cobra.Command, env *Env, name, model string, force bool) error { - return generateServiceInDir(cmd, env, "services", LoadProject().Paths.Models, name, model, force) -} - func generateServiceInDir(cmd *cobra.Command, _ *Env, dir, modelDir, name, model string, force bool) error { pkg := pkgFromOutDir(dir) importPath, ref := resolveModelImport(dir, modelDir, model) diff --git a/cli/cmd/maketest.go b/cli/cmd/maketest.go index 511f543..0ba336e 100644 --- a/cli/cmd/maketest.go +++ b/cli/cmd/maketest.go @@ -32,10 +32,6 @@ func NewMakeTest(env *Env) *cobra.Command { return c } -func generateTest(cmd *cobra.Command, env *Env, name string, force bool) error { - return generateTestInDir(cmd, env, "tests", name, force) -} - func generateTestInDir(cmd *cobra.Command, _ *Env, dir, name string, force bool) error { pkg := pkgFromOutDir(dir) base := strings.TrimSuffix(inflect.Snake(name), "_test") + "_test" diff --git a/cli/cmd/project.go b/cli/cmd/project.go index e560fe4..a987a84 100644 --- a/cli/cmd/project.go +++ b/cli/cmd/project.go @@ -66,7 +66,7 @@ func LoadProject() *ProjectConfig { // resetProjectForTest is exported only via _test files via the unexported // name; it lets tests reset the cached singleton. // -//nolint:unused +//lint:ignore U1000 kept for tests that reset the cached project config func resetProjectForTest() { projectOnce = sync.Once{} project = nil diff --git a/drivers/redis/queue.go b/drivers/redis/queue.go index 08dadd6..8f9498d 100644 --- a/drivers/redis/queue.go +++ b/drivers/redis/queue.go @@ -17,13 +17,13 @@ import ( // Two Redis keys are used per logical queue: // - ":queue:" — LIST of ready jobs (LPUSH/BRPOPLPUSH) // - ":queue::reserved" — sorted set of reserved jobs -// scored by their visibility- -// timeout deadline; reaper -// requeues entries past their -// deadline. +// scored by their visibility- +// timeout deadline; reaper +// requeues entries past their +// deadline. // - ":queue::delayed" — sorted set of delayed jobs -// scored by their available_at -// epoch. +// scored by their available_at +// epoch. // // Promotion (delayed → ready and reserved-timeout → ready) is handled // inline on each Pop so no background worker is required. diff --git a/factory/factory.go b/factory/factory.go index c292b7e..274aa93 100644 --- a/factory/factory.go +++ b/factory/factory.go @@ -26,7 +26,6 @@ type Factory[T any] struct { states []StateFn[T] overrides []func(m *T) count int - beforeSave []func(m *T) afterMake []func(m *T) afterCreate []func(ctx context.Context, m *T) error faker *Faker diff --git a/graphql/execute.go b/graphql/execute.go index 899b386..281aedd 100644 --- a/graphql/execute.go +++ b/graphql/execute.go @@ -832,6 +832,7 @@ func coerceScalar(s *Scalar, val any) (any, error) { if b, ok := val.(bool); ok { return b, nil } + //lint:ignore ST1005 message wording follows the GraphQL spec return nil, fmt.Errorf("Boolean cannot represent %T", val) case "ID": switch v := val.(type) { diff --git a/graphql/parser.go b/graphql/parser.go index 5c14f56..aaff808 100644 --- a/graphql/parser.go +++ b/graphql/parser.go @@ -32,7 +32,6 @@ type variableDef struct { name string typ typeRef defaultVal value - hasDefault bool } // fragmentDef is a named fragment definition (fragment Name on Type { ... }). diff --git a/mock/mock.go b/mock/mock.go index 40058c2..083038c 100644 --- a/mock/mock.go +++ b/mock/mock.go @@ -5,11 +5,11 @@ // Three primitives: // // - Clock — controllable time source. Inject as time.Now() and -// Advance it deterministically. +// Advance it deterministically. // - Calls — generic call recorder. Counts and stores arguments -// each time a function is invoked. +// each time a function is invoked. // - HTTPServer — pre-canned httptest.Server with route-by-method -// responses and recorded request inspection. +// responses and recorded request inspection. package mock import ( diff --git a/openapi/robust_test.go b/openapi/robust_test.go index a3c92ef..6db79ba 100644 --- a/openapi/robust_test.go +++ b/openapi/robust_test.go @@ -179,17 +179,18 @@ func TestSchemaOf_AnonymousStruct(t *testing.T) { // time.Time, []byte and json tag variants in one struct. func TestSchemaOf_AssortedKinds(t *testing.T) { type kitchen struct { - Ptr *int `json:"ptr"` - Slice []string `json:"slice"` - Map map[string]int `json:"map"` - Any any `json:"any"` - Iface interface{} `json:"iface"` - When time.Time `json:"when"` - Bytes []byte `json:"bytes"` - Omit string `json:"omit,omitempty"` - Skipped string `json:"-"` - unexp string // unexported, must be ignored - Renamed string `json:"renamed_field"` + Ptr *int `json:"ptr"` + Slice []string `json:"slice"` + Map map[string]int `json:"map"` + Any any `json:"any"` + Iface interface{} `json:"iface"` + When time.Time `json:"when"` + Bytes []byte `json:"bytes"` + Omit string `json:"omit,omitempty"` + Skipped string `json:"-"` + //lint:ignore U1000 fixture: unexported field must be skipped + unexp string + Renamed string `json:"renamed_field"` NoTagName string } var s *openapi.Schema @@ -242,6 +243,7 @@ func TestSchemaOf_EmbeddedRecursive(t *testing.T) { } type embedSelf struct { + //lint:ignore U1000 fixture: self-embedding cycle *embedSelf Value string `json:"value"` } diff --git a/openapi/spec.go b/openapi/spec.go index 25b722e..b91c3df 100644 --- a/openapi/spec.go +++ b/openapi/spec.go @@ -337,7 +337,7 @@ func (s *Spec) build() document { Paths: s.paths, } for _, sv := range s.servers { - doc.Servers = append(doc.Servers, serverObj{URL: sv.URL, Description: sv.Description}) + doc.Servers = append(doc.Servers, serverObj(sv)) } if names := s.registry.names(); len(names) > 0 { schemas := make(map[string]*Schema, len(names)) diff --git a/query/injection_test.go b/query/injection_test.go index c1ecd97..3d351af 100644 --- a/query/injection_test.go +++ b/query/injection_test.go @@ -22,9 +22,9 @@ var injectionPayloads = []string{ "1; DELETE FROM users", "admin'--", "' UNION SELECT password FROM users --", - "Ada\x00Lovelace", // embedded NUL - "O'Brien", // legitimate apostrophe - "café — 日本語 — ‮", // unicode incl. RTL override + "Ada\x00Lovelace", // embedded NUL + "O'Brien", // legitimate apostrophe + "café — 日本語 — \u202e", // unicode incl. RTL override "%' OR '1'='1", "\\'; DROP TABLE x; --", } diff --git a/realtime/conn.go b/realtime/conn.go index e3b54cf..ab38754 100644 --- a/realtime/conn.go +++ b/realtime/conn.go @@ -33,7 +33,6 @@ package realtime import ( "errors" - "io" ) // MessageType distinguishes a UTF-8 text frame from a binary frame. @@ -67,9 +66,3 @@ type Conn interface { WriteMessage(MessageType, []byte) error Close() error } - -// isClosed reports whether err signals a normal end-of-connection rather -// than an unexpected failure. -func isClosed(err error) bool { - return err == nil || errors.Is(err, io.EOF) || errors.Is(err, ErrClosed) || errors.Is(err, io.ErrClosedPipe) -} diff --git a/realtime/stress_test.go b/realtime/stress_test.go index 06d3fb3..01c8aec 100644 --- a/realtime/stress_test.go +++ b/realtime/stress_test.go @@ -99,8 +99,8 @@ func stressScale(t *testing.T) (clients, channels, broadcasts int) { func TestStressBroadcastStormPresenceChurn(t *testing.T) { clients, channels, broadcasts := stressScale(t) - base := settleGoroutines(0, 0, time.Second) // quiesce before measuring - base = runtime.NumGoroutine() + settleGoroutines(0, 0, time.Second) // quiesce before measuring + base := runtime.NumGoroutine() var presenceEvents uint64 h := NewHub( @@ -202,8 +202,8 @@ func TestStressSlowConsumerDrop(t *testing.T) { clients, channels, broadcasts := stressScale(t) const outbox = 16 - base := settleGoroutines(0, 0, time.Second) - base = runtime.NumGoroutine() + settleGoroutines(0, 0, time.Second) + base := runtime.NumGoroutine() h := NewHub(WithOutbox(outbox), WithSlowConsumerPolicy(DropMessage)) @@ -263,8 +263,8 @@ func TestStressSlowConsumerDrop(t *testing.T) { func TestStressSlowConsumerDisconnect(t *testing.T) { clients, channels, broadcasts := stressScale(t) - base := settleGoroutines(0, 0, time.Second) - base = runtime.NumGoroutine() + settleGoroutines(0, 0, time.Second) + base := runtime.NumGoroutine() h := NewHub(WithOutbox(8), WithSlowConsumerPolicy(DisconnectClient)) diff --git a/validation/robust_test.go b/validation/robust_test.go index a26f3ed..8b20ad7 100644 --- a/validation/robust_test.go +++ b/validation/robust_test.go @@ -71,10 +71,10 @@ func adversarialValues() []any { map[string]int{}, map[string]int{"k": 1}, [3]int{1, 2, 3}, long, - "السلام عليكم", // Arabic / RTL - "‮evil‬", // RTL override embedded - "日本語テスト", // CJK - "emoji 🚀🔥💥", // multibyte emoji + "السلام عليكم", // Arabic / RTL + "\u202eevil\u202c", // RTL override embedded + "日本語テスト", // CJK + "emoji 🚀🔥💥", // multibyte emoji "\x00\x01\x02null bytes", weird{X: 1}, &weird{X: 2}, @@ -238,9 +238,7 @@ func TestMalformedTagsNoPanic(t *testing.T) { // Build a struct value dynamically is awkward; instead drive the // same parse + dispatch path through Map using split rules. var rules []string - for _, r := range splitRules(tag) { - rules = append(rules, r) - } + rules = append(rules, splitRules(tag)...) _ = Map(map[string]any{"f": "value", "f_confirmation": "value"}, Rules{"f": rules}) }() } diff --git a/view/funcs.go b/view/funcs.go index bb87681..d6bebb2 100644 --- a/view/funcs.go +++ b/view/funcs.go @@ -34,9 +34,10 @@ func builtinFuncs() template.FuncMap { "default": defaultVal, // --- string helpers ------------------------------------------------ - "upper": strings.ToUpper, - "lower": strings.ToLower, - "title": strings.Title, //nolint:staticcheck // adequate for ASCII view text + "upper": strings.ToUpper, + "lower": strings.ToLower, + //lint:ignore SA1019 adequate for ASCII view text + "title": strings.Title, "trim": strings.TrimSpace, "replace": func(old, new, s string) string { return strings.ReplaceAll(s, old, new) }, "contains": strings.Contains, diff --git a/web/app.go b/web/app.go index 52191be..cad2ed2 100644 --- a/web/app.go +++ b/web/app.go @@ -8,7 +8,6 @@ import ( "os" "os/signal" "strings" - "sync" "syscall" "time" @@ -31,7 +30,6 @@ type App struct { timeout time.Duration logger *log.Logger noDefaults bool - mu sync.Mutex } // Option — funksional sozlash usuli.