Merge pull request #160 from jwetzell/fix/database-module-interface

rework database module to not expose raw DB
This commit is contained in:
Joel Wetzell
2026-05-19 12:38:13 -05:00
committed by GitHub
4 changed files with 33 additions and 45 deletions
+1 -1
View File
@@ -22,5 +22,5 @@ type KeyValueModule interface {
} }
type DatabaseModule interface { type DatabaseModule interface {
Database() (*sql.DB, error) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
} }
+28 -29
View File
@@ -53,51 +53,50 @@ func init() {
}) })
} }
func (t *DbSqlite) Id() string { func (dbs *DbSqlite) Id() string {
return t.config.Id return dbs.config.Id
} }
func (t *DbSqlite) Type() string { func (dbs *DbSqlite) Type() string {
return t.config.Type return dbs.config.Type
} }
func (t *DbSqlite) Start(ctx context.Context, router common.RouteIO) error { func (dbs *DbSqlite) Start(ctx context.Context, router common.RouteIO) error {
t.logger.Debug("running") dbs.logger.Debug("running")
t.router = router dbs.router = router
moduleContext, cancel := context.WithCancel(ctx) moduleContext, cancel := context.WithCancel(ctx)
t.ctx = moduleContext dbs.ctx = moduleContext
t.cancel = cancel dbs.cancel = cancel
db, err := sql.Open("sqlite", t.Dsn) db, err := sql.Open("sqlite", dbs.Dsn)
if err != nil { if err != nil {
return fmt.Errorf("db.sqlite error opening database: %w", err) return fmt.Errorf("db.sqlite error opening database: %w", err)
} }
t.dbMu.Lock() dbs.dbMu.Lock()
t.db = db dbs.db = db
t.dbMu.Unlock() dbs.dbMu.Unlock()
<-t.ctx.Done() <-dbs.ctx.Done()
return nil return nil
} }
func (t *DbSqlite) Stop() { func (dbs *DbSqlite) Stop() {
if t.cancel != nil { if dbs.cancel != nil {
t.cancel() dbs.cancel()
} }
t.dbMu.Lock() dbs.dbMu.Lock()
defer t.dbMu.Unlock() defer dbs.dbMu.Unlock()
if t.db != nil { if dbs.db != nil {
t.db.Close() dbs.db.Close()
t.db = nil dbs.db = nil
} }
t.logger.Debug("done") dbs.logger.Debug("done")
} }
// TODO(jwetzell): get a database module layout that doesn't require handing the DB over func (dbs *DbSqlite) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) {
func (t *DbSqlite) Database() (*sql.DB, error) { dbs.dbMu.Lock()
t.dbMu.Lock() defer dbs.dbMu.Unlock()
defer t.dbMu.Unlock() if dbs.db == nil {
if t.db == nil {
return nil, fmt.Errorf("database not initialized") return nil, fmt.Errorf("database not initialized")
} }
return t.db, nil return dbs.db.QueryContext(ctx, query, args...)
} }
+2 -13
View File
@@ -42,19 +42,8 @@ func (dq *DbQuery) Process(ctx context.Context, wrappedPayload common.WrappedPay
dq.module = dbModule dq.module = dbModule
} }
db, err := dq.module.Database()
if err != nil {
wrappedPayload.End = true
return wrappedPayload, fmt.Errorf("db.query error getting database from module: %w", err)
}
if db == nil {
wrappedPayload.End = true
return wrappedPayload, fmt.Errorf("db.query module with id %s returned nil database", dq.ModuleId)
}
var queryBuffer bytes.Buffer var queryBuffer bytes.Buffer
err = dq.Query.Execute(&queryBuffer, wrappedPayload) err := dq.Query.Execute(&queryBuffer, wrappedPayload)
if err != nil { if err != nil {
wrappedPayload.End = true wrappedPayload.End = true
@@ -62,7 +51,7 @@ func (dq *DbQuery) Process(ctx context.Context, wrappedPayload common.WrappedPay
} }
// support proper parameterized queries // support proper parameterized queries
rows, err := db.QueryContext(ctx, queryBuffer.String()) rows, err := dq.module.QueryContext(ctx, queryBuffer.String())
if err != nil { if err != nil {
wrappedPayload.End = true wrappedPayload.End = true
return wrappedPayload, fmt.Errorf("db.query error executing query: %w", err) return wrappedPayload, fmt.Errorf("db.query error executing query: %w", err)
+2 -2
View File
@@ -79,7 +79,7 @@ func (m *TestDBModule) Start(ctx context.Context, router common.RouteIO) error {
return nil return nil
} }
func (m *TestDBModule) Database() (*sql.DB, error) { func (m *TestDBModule) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) {
if m.db == nil { if m.db == nil {
db, err := sql.Open("sqlite", ":memory:") db, err := sql.Open("sqlite", ":memory:")
if err != nil { if err != nil {
@@ -99,7 +99,7 @@ func (m *TestDBModule) Database() (*sql.DB, error) {
} }
m.db = db m.db = db
} }
return m.db, nil return m.db.QueryContext(ctx, query, args...)
} }
func (m *TestDBModule) Stop() {} func (m *TestDBModule) Stop() {}