package dbutil import ( "testing" "time" "github.com/google/uuid" ) // toSnakeCase func TestToSnakeCase(t *testing.T) { cases := []struct{ in, want string }{ {"AppUser", "app_user"}, {"OrgUser", "org_user"}, {"ID", "id"}, {"IPAddr", "ip_addr"}, {"OrgUserDTO", "org_user_dto"}, {"CreatedBy", "created_by"}, {"EmailNotificationECd", "email_notification_e_cd"}, {"HTMLParser", "html_parser"}, {"Simple", "simple"}, } for _, tc := range cases { got := toSnakeCase(tc.in) if got != tc.want { t.Errorf("toSnakeCase(%q) = %q, want %q", tc.in, got, tc.want) } } } // TableRef func TestTableRef(t *testing.T) { u := T[testUser]("u") if u.ref() != "u" { t.Errorf("ref() = %q, want %q", u.ref(), "u") } if u.fromExpr() != "test_user u" { t.Errorf("fromExpr() = %q, want %q", u.fromExpr(), "test_user u") } noAlias := T[testUser]() if noAlias.ref() != "test_user" { t.Errorf("ref() without alias = %q, want %q", noAlias.ref(), "test_user") } if noAlias.fromExpr() != "test_user" { t.Errorf("fromExpr() without alias = %q, want %q", noAlias.fromExpr(), "test_user") } } func TestTableRefPanicsOnUnregistered(t *testing.T) { defer func() { if r := recover(); r == nil { t.Error("expected panic for unregistered type") } }() type Bogus struct { X string `db:"x"` } T[Bogus]() } func TestTName(t *testing.T) { ref := TName("my_table", "mt") if ref.fromExpr() != "my_table mt" { t.Errorf("fromExpr() = %q, want %q", ref.fromExpr(), "my_table mt") } } // Cols / ColsFlat / AllColNames func TestColsFlat(t *testing.T) { s := T[testSession]("s") flat := s.ColsFlat() if !containsSubstr(flat, "s.id") { t.Errorf("ColsFlat missing s.id: %s", flat) } if containsSubstr(flat, " AS ") { t.Errorf("ColsFlat should not contain AS aliases: %s", flat) } } func TestCols(t *testing.T) { u := T[testUser]("u") cols := u.Cols() if !containsSubstr(cols, `u.id AS "test_user.id"`) { t.Errorf("Cols missing aliased id: %s", cols) } if !containsSubstr(cols, `u.username AS "test_user.username"`) { t.Errorf("Cols missing aliased username: %s", cols) } } func TestColsMapAs(t *testing.T) { cb := T[testUser]("cb").MapAs("created_by") cols := cb.Cols() if !containsSubstr(cols, `cb.id AS "created_by.id"`) { t.Errorf("MapAs Cols missing aliased id: %s", cols) } } func TestAllColNames(t *testing.T) { s := T[testSession]() names := s.AllColNames() want := []string{"id", "key", "user_id", "org_id", "created", "user_agent", "revoked"} if len(names) != len(want) { t.Fatalf("AllColNames len = %d, want %d\ngot: %v\nwant: %v", len(names), len(want), names, want) } for i, n := range names { if n != want[i] { t.Errorf("AllColNames[%d] = %q, want %q", i, n, want[i]) } } } // Col conditions func TestColConditions(t *testing.T) { u := T[testUser]("u") c := u.F(&u.M.ID) id1 := uuid.New() id2 := uuid.New() cases := []struct { name string cond Cond frag string argc int }{ {"Eq", c.Eq(id1), "u.id = ?", 1}, {"Neq", c.Neq(id1), "u.id <> ?", 1}, {"Gt", c.Gt(id1), "u.id > ?", 1}, {"GtEq", c.GtEq(id1), "u.id >= ?", 1}, {"Lt", c.Lt(id1), "u.id < ?", 1}, {"LtEq", c.LtEq(id1), "u.id <= ?", 1}, {"Like", c.Like("%x%"), "u.id LIKE ?", 1}, {"IsNull", c.IsNull(), "u.id IS NULL", 0}, {"IsNotNull", c.IsNotNull(), "u.id IS NOT NULL", 0}, {"Between", c.Between(id1, id2), "u.id BETWEEN ? AND ?", 2}, {"EqCol", c.EqCol(u.C("other")), "u.id = u.other", 0}, } for _, tc := range cases { if tc.cond.fragment != tc.frag { t.Errorf("%s: fragment = %q, want %q", tc.name, tc.cond.fragment, tc.frag) } if len(tc.cond.args) != tc.argc { t.Errorf("%s: args len = %d, want %d", tc.name, len(tc.cond.args), tc.argc) } } } func TestColIn(t *testing.T) { u := T[testUser]("u") id1, id2, id3 := uuid.New(), uuid.New(), uuid.New() // Variadic c := u.F(&u.M.ID).In(id1, id2, id3) if c.fragment != "u.id IN (?, ?, ?)" { t.Errorf("In variadic fragment = %q", c.fragment) } if len(c.args) != 3 { t.Errorf("In variadic args len = %d", len(c.args)) } // Slice expansion ids := []uuid.UUID{uuid.New(), uuid.New()} c2 := u.F(&u.M.ID).In(ids) if c2.fragment != "u.id IN (?, ?)" { t.Errorf("In slice fragment = %q", c2.fragment) } if len(c2.args) != 2 { t.Errorf("In slice args len = %d", len(c2.args)) } } func TestLower(t *testing.T) { u := T[testUser]("u") c := Lower(u.F(&u.M.Username)) if c.expr != "LOWER(u.username)" { t.Errorf("Lower expr = %q", c.expr) } } // Cond composition func TestCondAndOr(t *testing.T) { a := Cond{fragment: "a = ?", args: []any{1}} b := Cond{fragment: "b = ?", args: []any{2}} and := a.And(b) if and.fragment != "(a = ? AND b = ?)" { t.Errorf("And fragment = %q", and.fragment) } if len(and.args) != 2 { t.Errorf("And args len = %d", len(and.args)) } or := a.Or(b) if or.fragment != "(a = ? OR b = ?)" { t.Errorf("Or fragment = %q", or.fragment) } } func TestCondIdentity(t *testing.T) { empty := Cond{} real := Cond{fragment: "x = ?", args: []any{1}} if empty.And(real).fragment != real.fragment { t.Error("empty.And(real) should return real") } if real.And(empty).fragment != real.fragment { t.Error("real.And(empty) should return real") } } func TestCondNot(t *testing.T) { c := Cond{fragment: "a = ?", args: []any{1}} n := c.Not() if n.fragment != "NOT (a = ?)" { t.Errorf("Not fragment = %q", n.fragment) } } // SELECT Build func TestSelectSimple(t *testing.T) { u := T[testUser]("u") id := uuid.New() sql, args := Select(u.ColsFlat()). From(u). Where(u.F(&u.M.ID).Eq(id)). Build() if !containsSubstr(sql, "SELECT u.id") { t.Errorf("missing columns: %s", sql) } if !containsSubstr(sql, "FROM test_user u") { t.Errorf("missing FROM: %s", sql) } if !containsSubstr(sql, "WHERE u.id = $1") { t.Errorf("missing WHERE with $1: %s", sql) } if len(args) != 1 || args[0] != id { t.Errorf("args = %v", args) } } func TestSelectJoin(t *testing.T) { m := T[testMembership]("m") u := T[testUser]("u") orgID := uuid.New() sql, args := Select(m.Cols(), u.Cols()). From(m). InnerJoin(u, u.F(&u.M.ID).EqCol(m.F(&m.M.UserID))). Where(m.F(&m.M.OrgID).Eq(orgID)). OrderBy(u.F(&u.M.LastName).Asc()). Limit(25). Offset(50). Build() if !containsSubstr(sql, "INNER JOIN test_user u ON u.id = m.user_id") { t.Errorf("missing JOIN: %s", sql) } if !containsSubstr(sql, "WHERE m.org_id = $1") { t.Errorf("missing WHERE: %s", sql) } if !containsSubstr(sql, "ORDER BY u.last_name ASC") { t.Errorf("missing ORDER BY: %s", sql) } if !containsSubstr(sql, "LIMIT 25") { t.Errorf("missing LIMIT: %s", sql) } if !containsSubstr(sql, "OFFSET 50") { t.Errorf("missing OFFSET: %s", sql) } if len(args) != 1 { t.Errorf("args = %v", args) } } func TestSelectLeftJoin(t *testing.T) { s := T[testSession]("s") u := T[testUser]("u") sql, _ := Select(s.Cols(), u.Cols()). From(s). LeftJoin(u, u.F(&u.M.ID).EqCol(s.F(&s.M.UserID))). Build() if !containsSubstr(sql, "LEFT JOIN test_user u ON u.id = s.user_id") { t.Errorf("missing LEFT JOIN: %s", sql) } } func TestSelectCount(t *testing.T) { m := T[testMembership]("m") orgID := uuid.New() sql, args := Select("COUNT(*)"). From(m). Where(m.F(&m.M.OrgID).Eq(orgID)). Build() if sql != "SELECT COUNT(*) FROM test_membership m WHERE m.org_id = $1" { t.Errorf("sql = %q", sql) } if len(args) != 1 { t.Errorf("args len = %d", len(args)) } } func TestSelectMultipleWhereParams(t *testing.T) { u := T[testUser]("u") now := time.Now() sql, args := Select(u.ColsFlat()). From(u). Where( u.F(&u.M.LoginCount).Gt(int32(5)). And(u.F(&u.M.Created).GtEq(now)). And(u.F(&u.M.Username).Like("%admin%")), ). Build() if !containsSubstr(sql, "$1") && !containsSubstr(sql, "$2") && !containsSubstr(sql, "$3") { t.Errorf("missing param placeholders: %s", sql) } if len(args) != 3 { t.Errorf("args len = %d, want 3", len(args)) } } func TestSelectSubquery(t *testing.T) { s := T[testSession]("s") sub := Select(s.F(&s.M.ID).String()). From(s). Where(s.F(&s.M.UserID).Eq(uuid.New()).And(s.F(&s.M.Revoked).Eq(false))). OrderBy(s.F(&s.M.Created).Asc()). Limit(3) s2 := T[testSession]() sql, args := Update(s2). Set("revoked", true). Where(s2.F(&s2.M.ID).InQuery(sub)). Build() if !containsSubstr(sql, "IN (SELECT s.id FROM test_session s WHERE") { t.Errorf("missing subquery: %s", sql) } if len(args) != 3 { t.Errorf("args len = %d, want 3, args = %v", len(args), args) } if !containsSubstr(sql, "$1") || !containsSubstr(sql, "$2") || !containsSubstr(sql, "$3") { t.Errorf("params not sequential: %s", sql) } } func TestSelectGroupBy(t *testing.T) { u := T[testUser]("u") sql, _ := Select("u.active", "COUNT(*)"). From(u). GroupBy("u.active"). Build() if !containsSubstr(sql, "GROUP BY u.active") { t.Errorf("missing GROUP BY: %s", sql) } } func TestAndWhere(t *testing.T) { u := T[testUser]("u") q := Select(u.ColsFlat()).From(u) q.AndWhere(u.F(&u.M.ID).Eq(uuid.New())) q.AndWhere(u.F(&u.M.Username).Eq("bob")) sql, args := q.Build() if !containsSubstr(sql, "$1") || !containsSubstr(sql, "$2") { t.Errorf("missing params: %s", sql) } if len(args) != 2 { t.Errorf("args len = %d", len(args)) } } // INSERT Build func TestInsertValues(t *testing.T) { u := T[testUser]() id := uuid.New() sql, args := InsertInto(u). Columns(u.FieldNames(&u.M.ID, &u.M.Username, &u.M.Email)...). Values(id, "alice", "alice@example.com"). Build() if sql != "INSERT INTO test_user (id, username, email) VALUES ($1, $2, $3)" { t.Errorf("sql = %q", sql) } if len(args) != 3 { t.Errorf("args len = %d", len(args)) } if args[1] != "alice" { t.Errorf("args[1] = %v", args[1]) } } func TestInsertModel(t *testing.T) { s := T[testSession]() sess := testSession{ Key: "sess_abc", UserID: uuid.New(), } sql, args := InsertInto(s). Columns(s.FieldNames(&s.M.Key, &s.M.UserID)...). Model(sess). Build() if sql != "INSERT INTO test_session (key, user_id) VALUES ($1, $2)" { t.Errorf("sql = %q", sql) } if args[0] != "sess_abc" { t.Errorf("args[0] = %v", args[0]) } if args[1] != sess.UserID { t.Errorf("args[1] = %v", args[1]) } } func TestInsertModelAllColumns(t *testing.T) { s := T[testSession]() sess := testSession{Key: "k"} sql, args := InsertInto(s).Model(sess).Build() if !containsSubstr(sql, "INSERT INTO test_session (id, key, user_id") { t.Errorf("sql = %q", sql) } if len(args) != 7 { // testSession has 7 fields t.Errorf("args len = %d, want 7", len(args)) } } // UPDATE Build func TestUpdateSet(t *testing.T) { u := T[testUser]() id := uuid.New() sql, args := Update(u). Set("login_count", 42). Set("active", false). Where(u.F(&u.M.ID).Eq(id)). Build() if sql != "UPDATE test_user SET login_count = $1, active = $2 WHERE test_user.id = $3" { t.Errorf("sql = %q", sql) } if len(args) != 3 { t.Errorf("args len = %d", len(args)) } if args[0] != 42 { t.Errorf("args[0] = %v", args[0]) } } func TestUpdateModelSetColumns(t *testing.T) { u := T[testUser]() user := testUser{ ID: uuid.New(), FirstName: "Alice", LastName: "Smith", Email: "alice@test.com", } sql, args := Update(u). SetColumns(u.FieldNames(&u.M.FirstName, &u.M.LastName, &u.M.Email)...). Model(user). Where(u.F(&u.M.ID).Eq(user.ID)). Build() if !containsSubstr(sql, "SET first_name = $1, last_name = $2, email = $3") { t.Errorf("missing SET: %s", sql) } if !containsSubstr(sql, "WHERE test_user.id = $4") { t.Errorf("missing WHERE: %s", sql) } if args[0] != "Alice" || args[1] != "Smith" || args[2] != "alice@test.com" { t.Errorf("args = %v", args) } } // DELETE Build func TestDeleteSimple(t *testing.T) { s := T[testSession]() sql, args := DeleteFrom(s). Where(s.F(&s.M.Key).Eq("sess_xyz")). Build() if sql != "DELETE FROM test_session WHERE test_session.key = $1" { t.Errorf("sql = %q", sql) } if len(args) != 1 || args[0] != "sess_xyz" { t.Errorf("args = %v", args) } } func TestDeleteNoWhere(t *testing.T) { s := T[testSession]() sql, args := DeleteFrom(s).Build() if sql != "DELETE FROM test_session" { t.Errorf("sql = %q", sql) } if len(args) != 0 { t.Errorf("args = %v", args) } } // replaceParams func TestReplaceParams(t *testing.T) { cases := []struct{ in, want string }{ {"x = ?", "x = $1"}, {"a = ? AND b = ?", "a = $1 AND b = $2"}, {"IN (?, ?, ?)", "IN ($1, $2, $3)"}, {"no params", "no params"}, } for _, tc := range cases { got := replaceParams(tc.in) if got != tc.want { t.Errorf("replaceParams(%q) = %q, want %q", tc.in, got, tc.want) } } } // extractModelValues func TestExtractModelValues(t *testing.T) { user := testUser{ Username: "bob", Email: "bob@test.com", FirstName: "Bob", } vals := extractModelValues(user, []string{"username", "email", "first_name"}) if vals[0] != "bob" || vals[1] != "bob@test.com" || vals[2] != "Bob" { t.Errorf("vals = %v", vals) } } func TestExtractModelValuesMissing(t *testing.T) { user := testUser{Username: "bob"} vals := extractModelValues(user, []string{"username", "nonexistent"}) if vals[0] != "bob" { t.Errorf("vals[0] = %v", vals[0]) } if vals[1] != nil { t.Errorf("vals[1] for missing column = %v, want nil", vals[1]) } } // Field references func TestFieldReference(t *testing.T) { u := T[testUser]("u") col := u.F(&u.M.ID) if col.String() != "u.id" { t.Errorf("F(&u.M.ID) = %q, want %q", col.String(), "u.id") } col2 := u.F(&u.M.FirstName) if col2.String() != "u.first_name" { t.Errorf("F(&u.M.FirstName) = %q, want %q", col2.String(), "u.first_name") } col3 := u.F(&u.M.Email) if col3.String() != "u.email" { t.Errorf("F(&u.M.Email) = %q, want %q", col3.String(), "u.email") } } func TestFieldNames(t *testing.T) { u := T[testUser]() names := u.FieldNames(&u.M.FirstName, &u.M.LastName, &u.M.Email) want := []string{"first_name", "last_name", "email"} if len(names) != len(want) { t.Fatalf("FieldNames len = %d, want %d", len(names), len(want)) } for i, n := range names { if n != want[i] { t.Errorf("FieldNames[%d] = %q, want %q", i, n, want[i]) } } } func TestFieldReferenceEqCol(t *testing.T) { u := T[testUser]("u") m := T[testMembership]("m") cond := u.F(&u.M.ID).EqCol(m.F(&m.M.UserID)) if cond.fragment != "u.id = m.user_id" { t.Errorf("EqCol fragment = %q", cond.fragment) } } func TestFieldReferenceInSelect(t *testing.T) { u := T[testUser]("u") sql, args := Select(u.F(&u.M.ID).String(), u.F(&u.M.Username).String()). From(u). Where(u.F(&u.M.Email).Eq("test@test.com")). Build() if sql != "SELECT u.id, u.username FROM test_user u WHERE u.email = $1" { t.Errorf("sql = %q", sql) } if len(args) != 1 || args[0] != "test@test.com" { t.Errorf("args = %v", args) } } func TestFieldNamesWithSetColumns(t *testing.T) { u := T[testUser]() user := testUser{ ID: uuid.New(), FirstName: "Test", LastName: "User", } sql, args := Update(u). SetColumns(u.FieldNames(&u.M.FirstName, &u.M.LastName)...). Model(user). Where(u.F(&u.M.ID).Eq(user.ID)). Build() if !containsSubstr(sql, "SET first_name = $1, last_name = $2") { t.Errorf("missing SET: %s", sql) } if !containsSubstr(sql, "WHERE test_user.id = $3") { t.Errorf("missing WHERE: %s", sql) } if args[0] != "Test" || args[1] != "User" { t.Errorf("args = %v", args) } } func TestFieldReferencePanicsOnBadPointer(t *testing.T) { defer func() { if r := recover(); r == nil { t.Error("expected panic for bad field pointer") } }() u := T[testUser]("u") var unrelated int u.F(&unrelated) } func TestTableGenericAs(t *testing.T) { u := T[testUser]("u") u2 := u.As("u2") if u2.ref() != "u2" { t.Errorf("As ref() = %q, want %q", u2.ref(), "u2") } col := u2.F(&u2.M.ID) if col.String() != "u2.id" { t.Errorf("F after As = %q, want %q", col.String(), "u2.id") } } func TestTableGenericMapAs(t *testing.T) { cb := T[testUser]("cb").MapAs("created_by") cols := cb.Cols() if !containsSubstr(cols, `cb.id AS "created_by.id"`) { t.Errorf("MapAs Cols missing aliased id: %s", cols) } col := cb.F(&cb.M.ID) if col.String() != "cb.id" { t.Errorf("F after MapAs = %q, want %q", col.String(), "cb.id") } } // checkType func TestCheckTypePanicsOnMismatch(t *testing.T) { u := T[testUser]("u") defer func() { r := recover() if r == nil { t.Fatal("expected panic for type mismatch") } msg, ok := r.(string) if !ok { t.Fatalf("panic value is not string: %v", r) } if !containsSubstr(msg, "type mismatch") { t.Errorf("panic message = %q, want it to contain 'type mismatch'", msg) } }() // Active is bool, passing string should panic u.F(&u.M.Active).Eq("true") } func TestCheckTypeSkipsForRawCol(t *testing.T) { u := T[testUser]("u") // C() returns a Col without fieldType — should not panic u.C("active").Eq("anything") } func TestCheckTypeSkipsForNilVal(t *testing.T) { u := T[testUser]("u") // nil should not panic even on typed columns u.F(&u.M.Active).Eq(nil) } // helpers func containsSubstr(s, sub string) bool { return len(s) >= len(sub) && (s == sub || len(s) > 0 && containsIdx(s, sub)) } func containsIdx(s, sub string) bool { for i := 0; i <= len(s)-len(sub); i++ { if s[i:i+len(sub)] == sub { return true } } return false }