package provider import ( "context" "errors" "testing" ) type recordedCall struct{ provider, op, result string } type fakeRecorder struct{ calls []recordedCall } func (f *fakeRecorder) ProviderRequest(provider, op, result string) { f.calls = append(f.calls, recordedCall{provider, op, result}) } // staticProvider returns canned values; only the classification of its // errors matters here. type staticProvider struct { createErr, deleteErr, getErr, listErr error } func (s *staticProvider) Create(context.Context, CreateRequest) (string, error) { return "id-1", s.createErr } func (s *staticProvider) Get(context.Context, string) (*Instance, error) { return &Instance{ID: "id-1"}, s.getErr } func (s *staticProvider) Delete(context.Context, string) error { return s.deleteErr } func (s *staticProvider) ListByTag(context.Context) ([]Instance, error) { return nil, s.listErr } func TestWithMetrics_recordsClassifiedResults(t *testing.T) { t.Parallel() tests := []struct { name string inner *staticProvider call func(p Provider) error wantOp string wantResult string }{ { name: "successful create is ok", inner: &staticProvider{}, call: func(p Provider) error { _, err := p.Create(context.Background(), CreateRequest{}); return err }, wantOp: "create", wantResult: "ok", }, { name: "get NotFound", inner: &staticProvider{getErr: Wrap(ErrNotFound, "get", "x", "id-1", nil)}, call: func(p Provider) error { _, err := p.Get(context.Background(), "id-1"); return err }, wantOp: "get", wantResult: "not_found", }, { name: "create quota", inner: &staticProvider{createErr: Wrap(ErrQuotaExceeded, "create", "x", "", nil)}, call: func(p Provider) error { _, err := p.Create(context.Background(), CreateRequest{}); return err }, wantOp: "create", wantResult: "quota_exceeded", }, { name: "delete permanent", inner: &staticProvider{deleteErr: Wrap(ErrPermanent, "delete", "x", "id-1", nil)}, call: func(p Provider) error { return p.Delete(context.Background(), "id-1") }, wantOp: "delete", wantResult: "permanent", }, { name: "unclassified list error is transient", inner: &staticProvider{listErr: errors.New("connection reset")}, call: func(p Provider) error { _, err := p.ListByTag(context.Background()); return err }, wantOp: "list", wantResult: "transient", }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() rec := &fakeRecorder{} p := WithMetrics("gcp-eu", tc.inner, rec) callErr := tc.call(p) if len(rec.calls) != 1 { t.Fatalf("recorded %d calls, want 1", len(rec.calls)) } want := recordedCall{provider: "gcp-eu", op: tc.wantOp, result: tc.wantResult} if rec.calls[0] != want { t.Errorf("recorded %+v, want %+v", rec.calls[0], want) } // The decorator must be transparent: errors pass through. if (tc.wantResult == "ok") != (callErr == nil) { t.Errorf("error passthrough broken: result %s but err %v", tc.wantResult, callErr) } }) } }