Files
kjol/go/dbutil/builder_test.go

738 lines
17 KiB
Go

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
}