diff --git a/.golangci.yml b/.golangci.yml index 1035e02..86f48f1 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -3,6 +3,7 @@ version: "2" linters: default: none enable: + - errcheck - govet - ineffassign - staticcheck diff --git a/controlplane/cleanup.go b/controlplane/cleanup.go new file mode 100644 index 0000000..d6c2571 --- /dev/null +++ b/controlplane/cleanup.go @@ -0,0 +1,14 @@ +package controlplane + +import ( + "errors" + "fmt" +) + +// closeWithError joins a cleanup failure onto a function's named error return +// without discarding the primary error. +func closeWithError(errp *error, context string, closeFn func() error) { + if err := closeFn(); err != nil { + *errp = errors.Join(*errp, fmt.Errorf("%s: %w", context, err)) + } +} diff --git a/controlplane/controlplane.go b/controlplane/controlplane.go index e8dc591..f321f55 100644 --- a/controlplane/controlplane.go +++ b/controlplane/controlplane.go @@ -215,7 +215,7 @@ func (s *Syncer) pushObservations(ctx context.Context) error { } } -func (s *Syncer) pushBatch(ctx context.Context, batch observationBatch) error { +func (s *Syncer) pushBatch(ctx context.Context, batch observationBatch) (err error) { body, err := json.Marshal(batch) if err != nil { return err @@ -234,7 +234,7 @@ func (s *Syncer) pushBatch(ctx context.Context, batch observationBatch) error { if err != nil { return err } - defer resp.Body.Close() + defer closeWithError(&err, "close response body", resp.Body.Close) if resp.StatusCode == http.StatusForbidden { return fmt.Errorf("POST observations returned %s (instance %q not registered with gatehub?)", resp.Status, s.cfg.InstanceID) } @@ -247,7 +247,7 @@ func (s *Syncer) pushBatch(ctx context.Context, batch observationBatch) error { // PullPolicy fetches verdicts since the last cursor and applies them locally. func (s *Syncer) PullPolicy() error { return s.pullPolicy(context.Background()) } -func (s *Syncer) pullPolicy(ctx context.Context) error { +func (s *Syncer) pullPolicy(ctx context.Context) (err error) { endpoint, err := endpointURL(s.cfg.URL, "/v1/policy", s.cfg.InstanceID, s.cursor) if err != nil { return err @@ -261,7 +261,7 @@ func (s *Syncer) pullPolicy(ctx context.Context) error { if err != nil { return err } - defer resp.Body.Close() + defer closeWithError(&err, "close response body", resp.Body.Close) if resp.StatusCode == http.StatusForbidden { return fmt.Errorf("GET policy returned %s (instance %q not registered with gatehub?)", resp.Status, s.cfg.InstanceID) } diff --git a/controlplane/controlplane_test.go b/controlplane/controlplane_test.go index 0ed10e0..c1ce44d 100644 --- a/controlplane/controlplane_test.go +++ b/controlplane/controlplane_test.go @@ -20,7 +20,11 @@ func openStore(t *testing.T) *store.Store { if err != nil { t.Fatalf("open store: %v", err) } - t.Cleanup(func() { s.Close() }) + t.Cleanup(func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }) return s } @@ -163,7 +167,7 @@ func TestPullPolicyAppliesDecisions(t *testing.T) { var sawCursor string srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { sawCursor = r.URL.Query().Get("since") - json.NewEncoder(w).Encode(policyResponse{ + if err := json.NewEncoder(w).Encode(policyResponse{ Cursor: "cursor-2", Decisions: []decision{ {Fingerprint: "known", Status: store.StatusBlocked, Label: "bad"}, @@ -171,7 +175,9 @@ func TestPullPolicyAppliesDecisions(t *testing.T) { {Fingerprint: "", Status: store.StatusApproved}, {Fingerprint: "junk", Status: store.Status("nonsense")}, }, - }) + }); err != nil { + t.Errorf("encode policy response: %v", err) + } })) defer srv.Close() @@ -217,7 +223,9 @@ func TestPullPolicyAppliesTrustedRangesWhenPresent(t *testing.T) { ranges := []string{"192.0.2.4/32", "2001:db8:1234::/64"} var applied []string srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - json.NewEncoder(w).Encode(policyResponse{TrustedRanges: &ranges}) + if err := json.NewEncoder(w).Encode(policyResponse{TrustedRanges: &ranges}); err != nil { + t.Errorf("encode policy response: %v", err) + } })) defer srv.Close() @@ -240,7 +248,9 @@ func TestPullPolicyOmittedTrustedRangesPreservesLocalState(t *testing.T) { st := openStore(t) called := false srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - json.NewEncoder(w).Encode(policyResponse{}) + if err := json.NewEncoder(w).Encode(policyResponse{}); err != nil { + t.Errorf("encode policy response: %v", err) + } })) defer srv.Close() diff --git a/proxy/proxy.go b/proxy/proxy.go index 04a2308..1b64c2f 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -121,7 +121,13 @@ func (s *Server) serve(ln net.Listener, route Route, handler Handler) { go func() { defer s.connWG.Done() defer s.sem.Release() - defer conn.Close() + // Nothing to return the error to on a served connection, so log it. + // A peer that already went away is the normal case, not a fault. + defer func() { + if err := conn.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + s.log("[%s] CLOSE %s: %v", clientIP, route.Listen, err) + } + }() handler(conn, route) }() } diff --git a/proxy/proxy_test.go b/proxy/proxy_test.go index 72520e4..9677930 100644 --- a/proxy/proxy_test.go +++ b/proxy/proxy_test.go @@ -63,12 +63,12 @@ func TestServerBoundsConnectionsAndDrains(t *testing.T) { }) serverOne, clientOne := net.Pipe() - defer clientOne.Close() + defer func() { _ = clientOne.Close() }() ln.send(serverOne) <-started serverTwo, clientTwo := net.Pipe() - defer clientTwo.Close() + defer func() { _ = clientTwo.Close() }() ln.send(serverTwo) _ = clientTwo.SetReadDeadline(time.Now().Add(time.Second)) if _, err := clientTwo.Read(make([]byte, 1)); err == nil { @@ -101,7 +101,7 @@ func TestDrainTimesOutForActiveHandler(t *testing.T) { <-release }) serverConn, clientConn := net.Pipe() - defer clientConn.Close() + defer func() { _ = clientConn.Close() }() ln.send(serverConn) <-started _ = ln.Close() diff --git a/sdnotify/cleanup.go b/sdnotify/cleanup.go new file mode 100644 index 0000000..adaf34e --- /dev/null +++ b/sdnotify/cleanup.go @@ -0,0 +1,14 @@ +package sdnotify + +import ( + "errors" + "fmt" +) + +// closeWithError joins a cleanup failure onto a function's named error return +// without discarding the primary error. +func closeWithError(errp *error, context string, closeFn func() error) { + if err := closeFn(); err != nil { + *errp = errors.Join(*errp, fmt.Errorf("%s: %w", context, err)) + } +} diff --git a/sdnotify/sdnotify.go b/sdnotify/sdnotify.go index 6bea3c6..30eb807 100644 --- a/sdnotify/sdnotify.go +++ b/sdnotify/sdnotify.go @@ -17,7 +17,7 @@ func Ready() error { // Notify sends state to NOTIFY_SOCKET. It is a no-op when the variable is // unset. A leading '@' denotes Linux's abstract Unix socket namespace. -func Notify(state string) error { +func Notify(state string) (err error) { socket := os.Getenv("NOTIFY_SOCKET") if socket == "" { return nil @@ -29,7 +29,7 @@ func Notify(state string) error { if err != nil { return fmt.Errorf("dial notification socket: %w", err) } - defer conn.Close() + defer closeWithError(&err, "close notification socket", conn.Close) if _, err := conn.Write([]byte(state)); err != nil { return fmt.Errorf("write notification: %w", err) } diff --git a/sdnotify/sdnotify_test.go b/sdnotify/sdnotify_test.go index 25cdaf7..1865424 100644 --- a/sdnotify/sdnotify_test.go +++ b/sdnotify/sdnotify_test.go @@ -24,7 +24,7 @@ func TestReadyIncludesCurrentPID(t *testing.T) { if err != nil { t.Fatal(err) } - defer listener.Close() + defer func() { _ = listener.Close() }() t.Setenv("NOTIFY_SOCKET", path) if err := Ready(); err != nil { diff --git a/store/cleanup.go b/store/cleanup.go new file mode 100644 index 0000000..f2e6901 --- /dev/null +++ b/store/cleanup.go @@ -0,0 +1,34 @@ +package store + +import ( + "database/sql" + "errors" + "fmt" +) + +// closeError runs a cleanup function and labels its failure, so a close error +// that reaches a caller says which resource failed to close. +func closeError(context string, closeFn func() error) error { + if err := closeFn(); err != nil { + return fmt.Errorf("%s: %w", context, err) + } + return nil +} + +// closeWithError joins a cleanup failure onto a function's named error return +// without discarding the primary error. Use it as: +// +// defer closeWithError(&err, "close fingerprint rows", rows.Close) +func closeWithError(errp *error, context string, closeFn func() error) { + *errp = errors.Join(*errp, closeError(context, closeFn)) +} + +// rollbackTransaction joins a rollback failure onto a named error return. +// A transaction that already committed reports sql.ErrTxDone, which is the +// expected outcome of the usual `defer rollbackTransaction(tx, &err)` guard +// and is not a failure. +func rollbackTransaction(tx *sql.Tx, errp *error) { + if err := tx.Rollback(); err != nil && !errors.Is(err, sql.ErrTxDone) { + *errp = errors.Join(*errp, fmt.Errorf("rollback transaction: %w", err)) + } +} diff --git a/store/legacy.go b/store/legacy.go index 556e943..ffdf855 100644 --- a/store/legacy.go +++ b/store/legacy.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "encoding/json" + "errors" "fmt" "strings" ) @@ -42,7 +43,7 @@ type LegacyColumn struct { // cannot overwrite metadata that gates have refreshed since. const metaLegacyMigrated = "gatekit_legacy_migrated" -func (s *Store) migrateLegacy(cols []LegacyColumn) error { +func (s *Store) migrateLegacy(cols []LegacyColumn) (err error) { if len(cols) == 0 { return nil } @@ -94,14 +95,12 @@ func (s *Store) migrateLegacy(cols []LegacyColumn) error { dest = append(dest, &raw[i]) } if err := rows.Scan(dest...); err != nil { - rows.Close() - return err + return errors.Join(err, closeError("close legacy rows", rows.Close)) } meta := map[string]any{} if strings.TrimSpace(metaJSON) != "" { if err := json.Unmarshal([]byte(metaJSON), &meta); err != nil { - rows.Close() - return fmt.Errorf("fingerprint %s: decode existing meta: %w", fp, err) + return errors.Join(fmt.Errorf("fingerprint %s: decode existing meta: %w", fp, err), closeError("close legacy rows", rows.Close)) } } for i, c := range present { @@ -114,8 +113,7 @@ func (s *Store) migrateLegacy(cols []LegacyColumn) error { } value, err := decodeLegacy(raw[i].String, c.Kind) if err != nil { - rows.Close() - return fmt.Errorf("fingerprint %s: column %s: %w", fp, c.Column, err) + return errors.Join(fmt.Errorf("fingerprint %s: column %s: %w", fp, c.Column, err), closeError("close legacy rows", rows.Close)) } if value == nil { continue @@ -124,27 +122,29 @@ func (s *Store) migrateLegacy(cols []LegacyColumn) error { } encoded, err := encodeMeta(meta) if err != nil { - rows.Close() - return err + return errors.Join(err, closeError("close legacy rows", rows.Close)) } updates = append(updates, update{fp: fp, meta: encoded}) } if err := rows.Err(); err != nil { - rows.Close() + return errors.Join(err, closeError("close legacy rows", rows.Close)) + } + // Closed before BeginTx, not deferred: reads and writes share the single + // writer connection, so holding these rows open would block the write. + if err := closeError("close legacy rows", rows.Close); err != nil { return err } - rows.Close() tx, err := s.db.BeginTx(ctx, nil) if err != nil { return err } - defer tx.Rollback() + defer rollbackTransaction(tx, &err) stmt, err := tx.PrepareContext(ctx, `UPDATE fingerprints SET meta = ? WHERE fp = ?`) if err != nil { return err } - defer stmt.Close() + defer closeWithError(&err, "close legacy update statement", stmt.Close) for _, u := range updates { if _, err := stmt.ExecContext(ctx, u.meta, u.fp); err != nil { return err diff --git a/store/legacy_test.go b/store/legacy_test.go index a4e7035..c4f4ceb 100644 --- a/store/legacy_test.go +++ b/store/legacy_test.go @@ -126,7 +126,11 @@ func TestMigrateSSHgateDatabase(t *testing.T) { if err != nil { t.Fatalf("Open legacy sshgate db: %v", err) } - defer s.Close() + defer func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }() entry, err := s.Get("fp1") if err != nil { @@ -191,7 +195,11 @@ func TestMigrateTLSgateDatabase(t *testing.T) { if err != nil { t.Fatalf("Open legacy tlsgate db: %v", err) } - defer s.Close() + defer func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }() entry, err := s.Get("fp1") if err != nil { @@ -242,13 +250,19 @@ func TestMigrateIsIdempotent(t *testing.T) { if _, err := s.Observe(Observation{Fingerprint: "fp1", Meta: map[string]any{"client_id": "new-client"}}, false); err != nil { t.Fatalf("Observe: %v", err) } - s.Close() + if err := s.Close(); err != nil { + t.Fatalf("Close before reopen: %v", err) + } s2, err := Open(Options{Path: path, Legacy: sshLegacyColumns}) if err != nil { t.Fatalf("second open: %v", err) } - defer s2.Close() + defer func() { + if err := s2.Close(); err != nil { + t.Errorf("Close second store: %v", err) + } + }() entry, err := s2.Get("fp1") if err != nil { t.Fatalf("Get: %v", err) @@ -268,7 +282,11 @@ func TestMigrateFreshDatabase(t *testing.T) { if err != nil { t.Fatalf("Open: %v", err) } - defer s.Close() + defer func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }() if _, err := s.Observe(Observation{Fingerprint: "fp1"}, false); err != nil { t.Fatalf("Observe: %v", err) } @@ -299,7 +317,11 @@ func TestObserveNewFingerprintOnMigratedDatabase(t *testing.T) { if err != nil { t.Fatalf("Open: %v", err) } - defer s.Close() + defer func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }() entry, err := s.Observe(Observation{ Fingerprint: "brandnew", IP: "192.0.2.99", @@ -330,13 +352,19 @@ func TestMigrateLeavesLegacyColumnsIntact(t *testing.T) { if err != nil { t.Fatalf("Open: %v", err) } - s.Close() + if err := s.Close(); err != nil { + t.Fatalf("Close before reopen: %v", err) + } db, err := sql.Open("sqlite", path) if err != nil { t.Fatalf("reopen raw: %v", err) } - defer db.Close() + defer func() { + if err := db.Close(); err != nil { + t.Errorf("Close raw db: %v", err) + } + }() var kex string if err := db.QueryRow(`SELECT kex FROM fingerprints WHERE fp = 'fp1'`).Scan(&kex); err != nil { t.Fatalf("legacy column gone: %v", err) diff --git a/store/limits.go b/store/limits.go index 3e2d01e..93b1ada 100644 --- a/store/limits.go +++ b/store/limits.go @@ -30,12 +30,12 @@ func trimHistory(ctx context.Context, tx *sql.Tx, fp string) error { // boundExistingHistory also handles databases created by older releases. Verdicts // and labels survive; oversized metadata is discarded rather than truncated into // invalid JSON. SQLite may retain freed pages for reuse. -func (s *Store) boundExistingHistory() error { +func (s *Store) boundExistingHistory() (err error) { tx, err := s.db.BeginTx(context.Background(), nil) if err != nil { return err } - defer tx.Rollback() + defer rollbackTransaction(tx, &err) if _, err := tx.Exec(`UPDATE fingerprints SET meta = '{}' WHERE length(CAST(meta AS BLOB)) > ?`, MaxMetadataBytes); err != nil { return err } @@ -62,7 +62,7 @@ func (s *Store) LastFingerprint() (string, error) { // ListPage reads a bounded, sorted page and only that page's observation history. // It is an eventually consistent view; concurrent observations are picked up by // the next synchronization cycle. -func (s *Store) ListPage(after, through string, limit int) ([]Entry, error) { +func (s *Store) ListPage(after, through string, limit int) (_ []Entry, err error) { if limit < 1 || limit > 128 { return nil, fmt.Errorf("page size must be between 1 and 128") } @@ -70,7 +70,7 @@ func (s *Store) ListPage(after, through string, limit int) ([]Entry, error) { if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close fingerprint rows", rows.Close) var entries []Entry for rows.Next() { e, err := scanEntry(rows) diff --git a/store/store.go b/store/store.go index 545f2ba..32258d6 100644 --- a/store/store.go +++ b/store/store.go @@ -137,23 +137,19 @@ func Open(opts Options) (*Store, error) { reader, err := sql.Open("sqlite", dsn(opts.Path, "busy_timeout=5000", "foreign_keys=ON", "journal_mode=WAL")) if err != nil { - db.Close() - return nil, err + return nil, errors.Join(err, closeError("close writer", db.Close)) } reader.SetMaxOpenConns(maxReaders) s := &Store{path: opts.Path, maxFingerprints: opts.MaxFingerprints, db: db, reader: reader} if err := s.init(); err != nil { - s.Close() - return nil, err + return nil, errors.Join(err, closeError("close store", s.Close)) } if err := s.migrateLegacy(opts.Legacy); err != nil { - s.Close() - return nil, fmt.Errorf("migrate legacy columns: %w", err) + return nil, errors.Join(fmt.Errorf("migrate legacy columns: %w", err), closeError("close store", s.Close)) } if err := s.boundExistingHistory(); err != nil { - s.Close() - return nil, fmt.Errorf("bound observation history: %w", err) + return nil, errors.Join(fmt.Errorf("bound observation history: %w", err), closeError("close store", s.Close)) } return s, nil } @@ -245,12 +241,12 @@ func (s *Store) addColumnIfMissing(ctx context.Context, table, column, def strin return err } -func (s *Store) hasColumn(ctx context.Context, table, column string) (bool, error) { +func (s *Store) hasColumn(ctx context.Context, table, column string) (_ bool, err error) { rows, err := s.db.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", table)) if err != nil { return false, err } - defer rows.Close() + defer closeWithError(&err, "close table info rows", rows.Close) for rows.Next() { var ( cid int @@ -317,7 +313,7 @@ func encodeMeta(meta map[string]any) (string, error) { // metadata bag are refreshed — status, label, and first_seen are left intact, // which preserves a prior verdict (and any pre-approved placeholder row // created by UpsertStatus) while still recording the latest handshake. -func (s *Store) Observe(obs Observation, blockUnknown bool) (Entry, error) { +func (s *Store) Observe(obs Observation, blockUnknown bool) (_ Entry, err error) { if obs.Fingerprint == "" { return Entry{}, errors.New("empty fingerprint") } @@ -334,7 +330,7 @@ func (s *Store) Observe(obs Observation, blockUnknown bool) (Entry, error) { if err != nil { return Entry{}, err } - defer tx.Rollback() + defer rollbackTransaction(tx, &err) now := encodeTime(time.Now()) status := StatusPending @@ -425,13 +421,13 @@ func (s *Store) Get(fp string) (Entry, error) { } // List returns every entry keyed by fingerprint, with IPs and ports attached. -func (s *Store) List() (map[string]Entry, error) { +func (s *Store) List() (_ map[string]Entry, err error) { ctx := context.Background() rows, err := s.reader.QueryContext(ctx, `SELECT `+entryColumns+` FROM fingerprints`) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close fingerprint rows", rows.Close) out := make(map[string]Entry) for rows.Next() { @@ -527,7 +523,7 @@ func (s *Store) Delete(fp string) error { // count is back at or below max, or until only approved entries remain. // Approved fingerprints are authoritative and never evicted. max <= 0 disables // pruning. Returns the number of entries deleted (ips/ports cascade). -func (s *Store) PruneToLimit(max int) (int, error) { +func (s *Store) PruneToLimit(max int) (_ int, err error) { if max <= 0 { return 0, nil } @@ -536,7 +532,7 @@ func (s *Store) PruneToLimit(max int) (int, error) { if err != nil { return 0, err } - defer tx.Rollback() + defer rollbackTransaction(tx, &err) deleted, err := pruneToLimit(ctx, tx, max) if err != nil { @@ -578,13 +574,13 @@ func pruneToLimit(ctx context.Context, tx *sql.Tx, max int) (int64, error) { // ResolveFingerprint maps a user-supplied query to exactly one stored // fingerprint, accepting either the full value or an unambiguous prefix so // CLIs don't force operators to paste full hashes. -func (s *Store) ResolveFingerprint(query string) (string, error) { +func (s *Store) ResolveFingerprint(query string) (_ string, err error) { query = strings.TrimSpace(query) if query == "" { return "", errors.New("empty fingerprint") } var exact string - err := s.reader.QueryRow(`SELECT fp FROM fingerprints WHERE fp = ?`, query).Scan(&exact) + err = s.reader.QueryRow(`SELECT fp FROM fingerprints WHERE fp = ?`, query).Scan(&exact) if err == nil { return exact, nil } @@ -598,7 +594,7 @@ func (s *Store) ResolveFingerprint(query string) (string, error) { if err != nil { return "", err } - defer rows.Close() + defer closeWithError(&err, "close fingerprint rows", rows.Close) var matches []string for rows.Next() { var fp string @@ -736,12 +732,12 @@ type querier interface { QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) } -func listStringsFrom(ctx context.Context, q querier, query, fp string) ([]string, error) { +func listStringsFrom(ctx context.Context, q querier, query, fp string) (_ []string, err error) { rows, err := q.QueryContext(ctx, query, fp) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close string rows", rows.Close) var out []string for rows.Next() { var v string @@ -753,12 +749,12 @@ func listStringsFrom(ctx context.Context, q querier, query, fp string) ([]string return out, rows.Err() } -func listIntsFrom(ctx context.Context, q querier, query, fp string) ([]int, error) { +func listIntsFrom(ctx context.Context, q querier, query, fp string) (_ []int, err error) { rows, err := q.QueryContext(ctx, query, fp) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close int rows", rows.Close) var out []int for rows.Next() { var v int @@ -770,12 +766,12 @@ func listIntsFrom(ctx context.Context, q querier, query, fp string) ([]int, erro return out, rows.Err() } -func allStrings(ctx context.Context, q querier, query string) (map[string][]string, error) { +func allStrings(ctx context.Context, q querier, query string) (_ map[string][]string, err error) { rows, err := q.QueryContext(ctx, query) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close string rows", rows.Close) out := make(map[string][]string) for rows.Next() { var fp, v string @@ -787,12 +783,12 @@ func allStrings(ctx context.Context, q querier, query string) (map[string][]stri return out, rows.Err() } -func allInts(ctx context.Context, q querier, query string) (map[string][]int, error) { +func allInts(ctx context.Context, q querier, query string) (_ map[string][]int, err error) { rows, err := q.QueryContext(ctx, query) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close int rows", rows.Close) out := make(map[string][]int) for rows.Next() { var fp string @@ -808,12 +804,12 @@ func allInts(ctx context.Context, q querier, query string) (map[string][]int, er return out, rows.Err() } -func listSightingsFrom(ctx context.Context, q querier, query string, args ...any) ([]Sighting, error) { +func listSightingsFrom(ctx context.Context, q querier, query string, args ...any) (_ []Sighting, err error) { rows, err := q.QueryContext(ctx, query, args...) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close sighting rows", rows.Close) var out []Sighting for rows.Next() { var sighting Sighting @@ -831,12 +827,12 @@ func listSightingsFrom(ctx context.Context, q querier, query string, args ...any return out, rows.Err() } -func allSightings(ctx context.Context, q querier, query string) (map[string][]Sighting, error) { +func allSightings(ctx context.Context, q querier, query string) (_ map[string][]Sighting, err error) { rows, err := q.QueryContext(ctx, query) if err != nil { return nil, err } - defer rows.Close() + defer closeWithError(&err, "close sighting rows", rows.Close) out := make(map[string][]Sighting) for rows.Next() { var fp, lastSeen string diff --git a/store/store_test.go b/store/store_test.go index e672398..6ec7db8 100644 --- a/store/store_test.go +++ b/store/store_test.go @@ -11,7 +11,11 @@ func openTest(t *testing.T) *Store { if err != nil { t.Fatalf("Open: %v", err) } - t.Cleanup(func() { s.Close() }) + t.Cleanup(func() { + if err := s.Close(); err != nil { + t.Errorf("Close: %v", err) + } + }) return s } @@ -354,7 +358,9 @@ func TestOpenCreatesParentDirectory(t *testing.T) { if err != nil { t.Fatalf("Open: %v", err) } - s.Close() + if err := s.Close(); err != nil { + t.Fatalf("Close: %v", err) + } } func TestOpenEmptyPath(t *testing.T) {