diff --git a/dest.go b/dest.go index 9f97804..a604a7a 100644 --- a/dest.go +++ b/dest.go @@ -20,14 +20,14 @@ func DestName(destTypeStruct any, path ...string) string { for i, p := range path { if v.Kind() != reflect.Struct { - dbxshared.DBModule.Logger().Error("DestName: path parent is not a struct", "path", destIdent+"."+strings.Join(path[:i+1], ".")) + destNameLogError("DestName: path parent is not a struct", destIdent+"."+strings.Join(path[:i+1], ".")) return "" } v = v.FieldByName(p) if !v.IsValid() { - dbxshared.DBModule.Logger().Error("DestName: field does not exist", "path", destIdent+"."+strings.Join(path[:i+1], ".")) + destNameLogError("DestName: field does not exist", destIdent+"."+strings.Join(path[:i+1], ".")) return "" } @@ -36,3 +36,10 @@ func DestName(destTypeStruct any, path ...string) string { return destIdent } + +func destNameLogError(msg, path string) { + if dbxshared.DBModule == nil { + return + } + dbxshared.DBModule.Logger().Error(msg, "path", path) +} diff --git a/internal/dbxshared/logger.go b/internal/dbxshared/logger.go index 8c60e39..2c757e1 100644 --- a/internal/dbxshared/logger.go +++ b/internal/dbxshared/logger.go @@ -1,5 +1,7 @@ package dbxshared +import "fmt" + type Logger interface { InitLogger() } @@ -20,11 +22,17 @@ func RegisterLogger(dialect dialectString, logger Logger) { } // InitLogger initializes the logger for a specific dialect. -func InitLogger(dialect dialectString) { +func InitLogger(dialect dialectString) error { dialectStr := dialect.String() logger, exists := loggerRegistry[dialectStr] if !exists { - panic("No logger registered for dialect: " + dialectStr) + return fmt.Errorf("no logger registered for dialect %q: blank-import dbxm (MySQL) or dbxp (Postgres)", dialectStr) } logger.InitLogger() + return nil +} + +// ResetLoggerRegistry clears registered loggers. It is intended for tests only. +func ResetLoggerRegistry() { + loggerRegistry = make(map[string]Logger) } diff --git a/internal/dbxshared/logger_test.go b/internal/dbxshared/logger_test.go new file mode 100644 index 0000000..ea821b6 --- /dev/null +++ b/internal/dbxshared/logger_test.go @@ -0,0 +1,32 @@ +package dbxshared + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +type testDialect string + +func (d testDialect) String() string { return string(d) } + +type stubLogger struct{} + +func (stubLogger) InitLogger() {} + +func TestInitLogger_UnregisteredDialect(t *testing.T) { + ResetLoggerRegistry() + t.Cleanup(ResetLoggerRegistry) + + err := InitLogger(testDialect("mysql")) + assert.ErrorContains(t, err, "blank-import dbxm") +} + +func TestInitLogger_RegisteredDialect(t *testing.T) { + ResetLoggerRegistry() + t.Cleanup(ResetLoggerRegistry) + + RegisterLogger(testDialect("mysql"), stubLogger{}) + err := InitLogger(testDialect("mysql")) + assert.NoError(t, err) +} diff --git a/module.go b/module.go index 521a339..44994b6 100644 --- a/module.go +++ b/module.go @@ -45,6 +45,11 @@ var ( state = dbState{} ) +// CurrentDialect returns the SQL dialect configured by [ModuleDB]. +func CurrentDialect() Dialect { + return state.dialect +} + // SQLO returns the current SQL database handle. func SQLO() *sql.DB { dbxshared.DBModule.RequireLoaded("dbx.SQLO requires database module") @@ -118,7 +123,10 @@ func setupDB(m *app.Module) error { ) if state.debugLog { - dbxshared.InitLogger(state.dialect) + if err := dbxshared.InitLogger(state.dialect); err != nil { + m.Logger().Error("Couldn't initialize Jet query debug logger", "err", err) + return err + } } return nil diff --git a/module_test.go b/module_test.go new file mode 100644 index 0000000..8d35d25 --- /dev/null +++ b/module_test.go @@ -0,0 +1,45 @@ +package dbx + +import ( + "testing" + + "gitea.auvem.com/go-toolkit/dbx/internal/dbxshared" + "github.com/stretchr/testify/assert" +) + +func TestDestName_MissingFieldWithoutModule(t *testing.T) { + type Foo struct { + Bar int + } + + name := DestName(Foo{}, "Missing") + assert.Empty(t, name) +} + +func TestDestName_ValidPath(t *testing.T) { + type Inner struct { + Value string + } + type Foo struct { + Inner Inner + } + + name := DestName(Foo{}, "Inner", "Value") + assert.Equal(t, "Foo.Inner.Value", name) +} + +func TestDialect_AfterModuleDB(t *testing.T) { + if dbxshared.DBModule != nil { + t.Skip("database module already initialized") + } + + _ = ModuleDB(DialectPostgres, &DBConfig{ + User: "user", + Password: "pass", + URI: "localhost:5432", + Name: "test", + MaxConn: 1, + }, false) + + assert.Equal(t, DialectPostgres, CurrentDialect()) +}