diff --git a/internal/common/module.go b/internal/common/module.go index cbf89f9..5f0ca55 100644 --- a/internal/common/module.go +++ b/internal/common/module.go @@ -22,5 +22,5 @@ type KeyValueModule interface { } type DatabaseModule interface { - Database() *sql.DB + Database() (*sql.DB, error) } diff --git a/internal/module/db-sqlite.go b/internal/module/db-sqlite.go index b79e44e..fe0a83d 100644 --- a/internal/module/db-sqlite.go +++ b/internal/module/db-sqlite.go @@ -79,6 +79,9 @@ func (t *DbSqlite) Stop() { } } -func (t *DbSqlite) Database() *sql.DB { - return t.db +func (t *DbSqlite) Database() (*sql.DB, error) { + if t.db == nil { + return nil, fmt.Errorf("database not initialized") + } + return t.db, nil } diff --git a/internal/processor/db-query.go b/internal/processor/db-query.go index b166014..c5893fe 100644 --- a/internal/processor/db-query.go +++ b/internal/processor/db-query.go @@ -41,14 +41,19 @@ func (dq *DbQuery) Process(ctx context.Context, wrappedPayload common.WrappedPay dq.module = dbModule } - db := dq.module.Database() + 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 - err := dq.Query.Execute(&queryBuffer, wrappedPayload) + err = dq.Query.Execute(&queryBuffer, wrappedPayload) if err != nil { wrappedPayload.End = true diff --git a/internal/test/module.go b/internal/test/module.go index c95e616..60d0e35 100644 --- a/internal/test/module.go +++ b/internal/test/module.go @@ -79,11 +79,14 @@ func (m *TestDBModule) Start(ctx context.Context, router common.RouteIO) error { return nil } -func (m *TestDBModule) Database() *sql.DB { +func (m *TestDBModule) Database() (*sql.DB, error) { if m.db == nil { - db, _ := sql.Open("sqlite", ":memory:") + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + return nil, err + } - db.Exec(` + _, err = db.Exec(` CREATE TABLE test ( id INTEGER PRIMARY KEY, value TEXT @@ -91,9 +94,12 @@ func (m *TestDBModule) Database() *sql.DB { INSERT INTO test (id, value) VALUES (1, 'test-1'), (2, 'test-2'); `) + if err != nil { + return nil, err + } m.db = db } - return m.db + return m.db, nil } func (m *TestDBModule) Stop() {}