Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .golangci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ version: "2"
linters:
default: none
enable:
- errcheck
- govet
- ineffassign
- staticcheck
Expand Down
14 changes: 14 additions & 0 deletions controlplane/cleanup.go
Original file line number Diff line number Diff line change
@@ -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))
}
}
8 changes: 4 additions & 4 deletions controlplane/controlplane.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
}
Expand All @@ -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
Expand All @@ -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)
}
Expand Down
20 changes: 15 additions & 5 deletions controlplane/controlplane_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down Expand Up @@ -163,15 +167,17 @@ 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"},
{Fingerprint: "unseen", Status: store.StatusApproved, Label: "preapproved"},
{Fingerprint: "", Status: store.StatusApproved},
{Fingerprint: "junk", Status: store.Status("nonsense")},
},
})
}); err != nil {
t.Errorf("encode policy response: %v", err)
}
}))
defer srv.Close()

Expand Down Expand Up @@ -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()

Expand All @@ -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()

Expand Down
8 changes: 7 additions & 1 deletion proxy/proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}()
}
Expand Down
6 changes: 3 additions & 3 deletions proxy/proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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()
Expand Down
14 changes: 14 additions & 0 deletions sdnotify/cleanup.go
Original file line number Diff line number Diff line change
@@ -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))
}
}
4 changes: 2 additions & 2 deletions sdnotify/sdnotify.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
}
Expand Down
2 changes: 1 addition & 1 deletion sdnotify/sdnotify_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
34 changes: 34 additions & 0 deletions store/cleanup.go
Original file line number Diff line number Diff line change
@@ -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))
}
}
26 changes: 13 additions & 13 deletions store/legacy.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"strings"
)
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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 {
Expand All @@ -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
Expand All @@ -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
Expand Down
Loading