508 lines
11 KiB
Go
508 lines
11 KiB
Go
package dbutil
|
|
|
|
import (
|
|
"reflect"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/DATA-DOG/go-sqlmock"
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
// buildMapping (flat struct)
|
|
|
|
func TestBuildMappingFlat(t *testing.T) {
|
|
tm := getMapping(reflect.TypeOf(testUser{}))
|
|
|
|
want := map[string]bool{
|
|
"id": true, "username": true, "email": true, "password": true,
|
|
"first_name": true, "last_name": true,
|
|
"login_count": true, "created": true, "active": true,
|
|
}
|
|
for col := range want {
|
|
if _, ok := tm.columns[col]; !ok {
|
|
t.Errorf("missing mapping for column %q", col)
|
|
}
|
|
}
|
|
if len(tm.columns) != len(want) {
|
|
t.Errorf("column count = %d, want %d", len(tm.columns), len(want))
|
|
}
|
|
}
|
|
|
|
// buildMapping (nested struct / DTO)
|
|
|
|
func TestBuildMappingNested(t *testing.T) {
|
|
type DTO struct {
|
|
Membership testMembership
|
|
User testUser
|
|
}
|
|
|
|
tm := getMapping(reflect.TypeOf(DTO{}))
|
|
|
|
if _, ok := tm.columns["membership.id"]; !ok {
|
|
t.Error("missing membership.id")
|
|
}
|
|
if _, ok := tm.columns["membership.user_id"]; !ok {
|
|
t.Error("missing membership.user_id")
|
|
}
|
|
|
|
if _, ok := tm.columns["user.id"]; !ok {
|
|
t.Error("missing user.id")
|
|
}
|
|
if _, ok := tm.columns["user.username"]; !ok {
|
|
t.Error("missing user.username")
|
|
}
|
|
|
|
// No unprefixed columns should exist
|
|
if _, ok := tm.columns["id"]; ok {
|
|
t.Error("should not have unprefixed 'id'")
|
|
}
|
|
}
|
|
|
|
// buildMapping (alias tag)
|
|
|
|
func TestBuildMappingAlias(t *testing.T) {
|
|
type DTO struct {
|
|
User testUser
|
|
CreatedBy testUser `alias:"created_by"`
|
|
}
|
|
|
|
tm := getMapping(reflect.TypeOf(DTO{}))
|
|
|
|
if _, ok := tm.columns["user.id"]; !ok {
|
|
t.Error("missing user.id")
|
|
}
|
|
if _, ok := tm.columns["created_by.id"]; !ok {
|
|
t.Error("missing created_by.id")
|
|
}
|
|
if _, ok := tm.columns["created_by.username"]; !ok {
|
|
t.Error("missing created_by.username")
|
|
}
|
|
}
|
|
|
|
// buildMapping (pointer-to-struct for LEFT JOINs)
|
|
|
|
func TestBuildMappingPointerStruct(t *testing.T) {
|
|
type DTO struct {
|
|
Session testSession
|
|
User *testUser `alias:"test_user"`
|
|
}
|
|
|
|
tm := getMapping(reflect.TypeOf(DTO{}))
|
|
|
|
if _, ok := tm.columns["session.id"]; !ok {
|
|
t.Error("missing session.id")
|
|
}
|
|
if _, ok := tm.columns["test_user.id"]; !ok {
|
|
t.Error("missing test_user.id (via alias tag)")
|
|
}
|
|
if len(tm.ptrStructs) != 1 {
|
|
t.Errorf("ptrStructs len = %d, want 1", len(tm.ptrStructs))
|
|
}
|
|
}
|
|
|
|
// buildMapping (embedded struct)
|
|
|
|
func TestBuildMappingEmbedded(t *testing.T) {
|
|
type Base struct {
|
|
ID uuid.UUID `db:"id"`
|
|
Created time.Time `db:"created"`
|
|
}
|
|
type Extended struct {
|
|
Base
|
|
Name string `db:"name"`
|
|
}
|
|
|
|
tm := getMapping(reflect.TypeOf(Extended{}))
|
|
|
|
if _, ok := tm.columns["id"]; !ok {
|
|
t.Error("missing flattened id from embedded Base")
|
|
}
|
|
if _, ok := tm.columns["created"]; !ok {
|
|
t.Error("missing flattened created from embedded Base")
|
|
}
|
|
if _, ok := tm.columns["name"]; !ok {
|
|
t.Error("missing name")
|
|
}
|
|
}
|
|
|
|
// isModelStruct
|
|
|
|
func TestIsModelStruct(t *testing.T) {
|
|
if !isModelStruct(reflect.TypeOf(testUser{})) {
|
|
t.Error("testUser should be a model struct")
|
|
}
|
|
if isModelStruct(reflect.TypeOf(time.Time{})) {
|
|
t.Error("time.Time should not be a model struct")
|
|
}
|
|
if isModelStruct(reflect.TypeOf(struct{ X int }{})) {
|
|
t.Error("anonymous struct without db tags should not be a model struct")
|
|
}
|
|
}
|
|
|
|
// DebugMapping
|
|
|
|
func TestDebugMapping(t *testing.T) {
|
|
type DTO struct {
|
|
User testUser
|
|
}
|
|
m := DebugMapping(DTO{})
|
|
|
|
if path, ok := m["user.id"]; !ok || path != "User.ID" {
|
|
t.Errorf("user.id mapping = %q, ok = %v", path, ok)
|
|
}
|
|
if path, ok := m["user.username"]; !ok || path != "User.Username" {
|
|
t.Errorf("user.username mapping = %q, ok = %v", path, ok)
|
|
}
|
|
}
|
|
|
|
// ScanOne (flat struct via sqlmock)
|
|
|
|
func TestScanOneFlat(t *testing.T) {
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
id := uuid.New()
|
|
now := time.Now().Truncate(time.Second)
|
|
|
|
rows := sqlmock.NewRows([]string{"id", "username", "email", "password", "first_name", "last_name", "login_count", "created", "active"}).
|
|
AddRow(id, "alice", "alice@test.com", "hash", "Alice", "Smith", int32(5), now, true)
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var user testUser
|
|
err = ScanOne(sqlRows, &user)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if user.ID != id {
|
|
t.Errorf("ID = %v, want %v", user.ID, id)
|
|
}
|
|
if user.Username != "alice" {
|
|
t.Errorf("Username = %q", user.Username)
|
|
}
|
|
if user.Email != "alice@test.com" {
|
|
t.Errorf("Email = %q", user.Email)
|
|
}
|
|
if user.FirstName != "Alice" {
|
|
t.Errorf("FirstName = %q", user.FirstName)
|
|
}
|
|
if user.LoginCount != 5 {
|
|
t.Errorf("LoginCount = %d", user.LoginCount)
|
|
}
|
|
if user.Active != true {
|
|
t.Errorf("Active = %v", user.Active)
|
|
}
|
|
}
|
|
|
|
// ScanAll (multiple rows)
|
|
|
|
func TestScanAllFlat(t *testing.T) {
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
id1, id2 := uuid.New(), uuid.New()
|
|
now := time.Now().Truncate(time.Second)
|
|
|
|
rows := sqlmock.NewRows([]string{"id", "key", "user_id", "org_id", "created", "user_agent", "revoked"}).
|
|
AddRow(id1, "key1", uuid.New(), nil, now, "Mozilla", false).
|
|
AddRow(id2, "key2", uuid.New(), nil, now, "Chrome", true)
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var sessions []testSession
|
|
err = ScanAll(sqlRows, &sessions)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if len(sessions) != 2 {
|
|
t.Fatalf("len = %d, want 2", len(sessions))
|
|
}
|
|
if sessions[0].Key != "key1" {
|
|
t.Errorf("[0].Key = %q", sessions[0].Key)
|
|
}
|
|
if sessions[1].Revoked != true {
|
|
t.Errorf("[1].Revoked = %v", sessions[1].Revoked)
|
|
}
|
|
}
|
|
|
|
// ScanOne (nested DTO with prefixed columns)
|
|
|
|
func TestScanOneNested(t *testing.T) {
|
|
type DTO struct {
|
|
Membership testMembership
|
|
User testUser
|
|
}
|
|
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
mID, uID, orgID := uuid.New(), uuid.New(), uuid.New()
|
|
now := time.Now().Truncate(time.Second)
|
|
|
|
cols := []string{
|
|
"membership.id", "membership.user_id", "membership.org_id",
|
|
"membership.created_by", "membership.joined",
|
|
"membership.login_count",
|
|
"user.id", "user.username", "user.email",
|
|
}
|
|
rows := sqlmock.NewRows(cols).
|
|
AddRow(
|
|
mID, uID, orgID,
|
|
uuid.New(), now,
|
|
int32(10),
|
|
uID, "alice", "alice@test.com",
|
|
)
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var dto DTO
|
|
err = ScanOne(sqlRows, &dto)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if dto.Membership.ID != mID {
|
|
t.Errorf("Membership.ID = %v, want %v", dto.Membership.ID, mID)
|
|
}
|
|
if dto.Membership.OrgID != orgID {
|
|
t.Errorf("Membership.OrgID = %v, want %v", dto.Membership.OrgID, orgID)
|
|
}
|
|
if dto.User.Username != "alice" {
|
|
t.Errorf("User.Username = %q", dto.User.Username)
|
|
}
|
|
if dto.Membership.LoginCount != 10 {
|
|
t.Errorf("Membership.LoginCount = %v, want 10", dto.Membership.LoginCount)
|
|
}
|
|
}
|
|
|
|
// ScanOne (alias tag)
|
|
|
|
func TestScanOneAlias(t *testing.T) {
|
|
type DTO struct {
|
|
User testUser
|
|
CreatedBy testUser `alias:"created_by"`
|
|
}
|
|
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
userID, cbID := uuid.New(), uuid.New()
|
|
|
|
cols := []string{"user.id", "user.username", "created_by.id", "created_by.username"}
|
|
rows := sqlmock.NewRows(cols).
|
|
AddRow(userID, "alice", cbID, "bob")
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var dto DTO
|
|
err = ScanOne(sqlRows, &dto)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if dto.User.ID != userID {
|
|
t.Errorf("User.ID = %v, want %v", dto.User.ID, userID)
|
|
}
|
|
if dto.User.Username != "alice" {
|
|
t.Errorf("User.Username = %q", dto.User.Username)
|
|
}
|
|
if dto.CreatedBy.ID != cbID {
|
|
t.Errorf("CreatedBy.ID = %v, want %v", dto.CreatedBy.ID, cbID)
|
|
}
|
|
if dto.CreatedBy.Username != "bob" {
|
|
t.Errorf("CreatedBy.Username = %q", dto.CreatedBy.Username)
|
|
}
|
|
}
|
|
|
|
// Pointer-to-struct nil detection (LEFT JOIN)
|
|
|
|
func TestScanOnePtrStructNilDetection(t *testing.T) {
|
|
type DTO struct {
|
|
Session testSession
|
|
User *testUser `alias:"test_user"`
|
|
}
|
|
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
sessID := uuid.New()
|
|
now := time.Now().Truncate(time.Second)
|
|
|
|
cols := []string{
|
|
"session.id", "session.key", "session.user_id", "session.org_id",
|
|
"session.created", "session.user_agent", "session.revoked",
|
|
// All user columns are NULL (LEFT JOIN miss)
|
|
"test_user.id", "test_user.username",
|
|
}
|
|
rows := sqlmock.NewRows(cols).
|
|
AddRow(
|
|
sessID, "key1", uuid.New(), nil,
|
|
now, "Mozilla", false,
|
|
// NULL user
|
|
uuid.Nil, "",
|
|
)
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var dto DTO
|
|
err = ScanOne(sqlRows, &dto)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if dto.Session.ID != sessID {
|
|
t.Errorf("Session.ID = %v, want %v", dto.Session.ID, sessID)
|
|
}
|
|
if dto.User != nil {
|
|
t.Errorf("User should be nil for LEFT JOIN miss, got %+v", dto.User)
|
|
}
|
|
}
|
|
|
|
func TestScanOnePtrStructNonNil(t *testing.T) {
|
|
type DTO struct {
|
|
Session testSession
|
|
User *testUser `alias:"test_user"`
|
|
}
|
|
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
sessID, userID := uuid.New(), uuid.New()
|
|
now := time.Now().Truncate(time.Second)
|
|
|
|
cols := []string{
|
|
"session.id", "session.key", "session.user_id", "session.org_id",
|
|
"session.created", "session.user_agent", "session.revoked",
|
|
"test_user.id", "test_user.username",
|
|
}
|
|
rows := sqlmock.NewRows(cols).
|
|
AddRow(
|
|
sessID, "key1", userID, nil,
|
|
now, "Mozilla", false,
|
|
userID, "alice",
|
|
)
|
|
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
var dto DTO
|
|
err = ScanOne(sqlRows, &dto)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
if dto.User == nil {
|
|
t.Fatal("User should not be nil")
|
|
}
|
|
if dto.User.ID != userID {
|
|
t.Errorf("User.ID = %v, want %v", dto.User.ID, userID)
|
|
}
|
|
if dto.User.Username != "alice" {
|
|
t.Errorf("User.Username = %q", dto.User.Username)
|
|
}
|
|
}
|
|
|
|
// ScanOne: sql.ErrNoRows
|
|
|
|
func TestScanOneNoRows(t *testing.T) {
|
|
db, mock, err := sqlmock.New()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer db.Close()
|
|
|
|
rows := sqlmock.NewRows([]string{"id", "username"})
|
|
mock.ExpectQuery("SELECT").WillReturnRows(rows)
|
|
|
|
sqlRows, err := db.Query("SELECT anything")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sqlRows.Close()
|
|
|
|
type Small struct {
|
|
ID string `db:"id"`
|
|
Username string `db:"username"`
|
|
}
|
|
var s Small
|
|
err = ScanOne(sqlRows, &s)
|
|
if err == nil {
|
|
t.Error("expected error for no rows")
|
|
}
|
|
}
|
|
|
|
// Columns function
|
|
|
|
func TestColumnsFunction(t *testing.T) {
|
|
result := Columns(testSession{}, "s")
|
|
if !containsSubstr(result, `s.id AS "test_session.id"`) {
|
|
t.Errorf("missing s.id alias: %s", result)
|
|
}
|
|
if !containsSubstr(result, `s.key AS "test_session.key"`) {
|
|
t.Errorf("missing s.key alias: %s", result)
|
|
}
|
|
}
|
|
|
|
func TestColumnsWithPrefix(t *testing.T) {
|
|
result := Columns(testUser{}, "cb", "created_by")
|
|
if !containsSubstr(result, `cb.id AS "created_by.id"`) {
|
|
t.Errorf("missing custom prefix alias: %s", result)
|
|
}
|
|
}
|