diff --git a/access_token_test.go b/access_token_test.go new file mode 100644 index 0000000..84df7a8 --- /dev/null +++ b/access_token_test.go @@ -0,0 +1,123 @@ +package flashduty + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" +) + +type seenRequest struct { + auth string + hasKey bool + key string +} + +func newCredentialClient(t *testing.T, accessToken bool) (*Client, *seenRequest) { + t.Helper() + seen := &seenRequest{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + seen.auth = r.Header.Get("Authorization") + seen.hasKey = r.URL.Query().Has("app_key") + seen.key = r.URL.Query().Get("app_key") + _, _ = io.WriteString(w, `{"request_id":"R","data":{}}`) + })) + t.Cleanup(srv.Close) + var ( + c *Client + err error + ) + if accessToken { + c, err = NewClientWithAccessToken("tok", WithBaseURL(srv.URL), WithLogger(noopLogger{})) + } else { + c, err = NewClient("KEY", WithBaseURL(srv.URL), WithLogger(noopLogger{})) + } + if err != nil { + t.Fatal(err) + } + return c, seen +} + +func TestNewClientWithAccessTokenRequiresToken(t *testing.T) { + if _, err := NewClientWithAccessToken(""); err == nil { + t.Fatal("expected error for empty access token") + } +} + +func TestAccessTokenClientSendsBearerOnEveryRequestKind(t *testing.T) { + ctx := context.Background() + kinds := map[string]func(c *Client) error{ + "json": func(c *Client) error { + _, err := c.do(ctx, "/incident/list", map[string]any{"p": 1}, nil) + return err + }, + "get": func(c *Client) error { + _, err := c.doGet(ctx, "/incident/info", nil, nil) + return err + }, + "upload": func(c *Client) error { + _, err := c.uploadFile(ctx, "/x/upload", url.Values{"a": {"b"}}, map[string]string{"f": "v"}, "a.txt", strings.NewReader("data"), nil) + return err + }, + } + for name, call := range kinds { + t.Run(name, func(t *testing.T) { + c, seen := newCredentialClient(t, true) + if err := call(c); err != nil { + t.Fatal(err) + } + if seen.auth != "Bearer tok" { + t.Errorf("Authorization = %q", seen.auth) + } + if seen.hasKey { + t.Errorf("app_key must not be sent, got %q", seen.key) + } + }) + } +} + +func TestNewClientStillSendsAppKey(t *testing.T) { + ctx := context.Background() + kinds := map[string]func(c *Client) error{ + "json": func(c *Client) error { + _, err := c.do(ctx, "/incident/list", nil, nil) + return err + }, + "upload": func(c *Client) error { + _, err := c.uploadFile(ctx, "/x/upload", nil, nil, "a.txt", strings.NewReader("data"), nil) + return err + }, + } + for name, call := range kinds { + t.Run(name, func(t *testing.T) { + c, seen := newCredentialClient(t, false) + if err := call(c); err != nil { + t.Fatal(err) + } + if seen.key != "KEY" || seen.auth != "" { + t.Errorf("app_key = %q, Authorization = %q", seen.key, seen.auth) + } + }) + } +} + +// The per-request trigger token is the only credential on the without-app-key +// path, for both client kinds. +func TestWithoutAppKeyPathUsesOnlyPerRequestBearer(t *testing.T) { + for _, accessToken := range []bool{false, true} { + c, seen := newCredentialClient(t, accessToken) + _, err := c.doMethodWithoutAppKey(context.Background(), http.MethodPost, "/safari/automation/triggers/t/fire", nil, nil, func(r *http.Request) { + r.Header.Set("Authorization", "Bearer trigger") + }) + if err != nil { + t.Fatal(err) + } + if seen.auth != "Bearer trigger" || seen.hasKey { + t.Errorf("accessToken=%v: Authorization = %q, app_key present = %v", accessToken, seen.auth, seen.hasKey) + } + } +} diff --git a/flashduty.go b/flashduty.go index 6ca934c..146b149 100644 --- a/flashduty.go +++ b/flashduty.go @@ -26,6 +26,7 @@ type Client struct { UserAgent string appKey string + accessToken string // when set, sent as a Bearer token instead of the app_key query parameter logger Logger requestHeaders http.Header requestHook func(*http.Request) @@ -61,6 +62,36 @@ func NewClient(appKey string, opts ...Option) (*Client, error) { return c, nil } +// NewClientWithAccessToken returns a Flashduty client that authenticates every +// request with an OAuth access token, sent as "Authorization: Bearer " +// instead of the app_key query parameter. +func NewClientWithAccessToken(accessToken string, opts ...Option) (*Client, error) { + if accessToken == "" { + return nil, fmt.Errorf("flashduty: access token is required") + } + c, err := NewClient(accessToken, opts...) + if err != nil { + return nil, err + } + c.accessToken = accessToken + c.appKey = "" + return c, nil +} + +// authQuery adds the app_key query parameter unless the client uses an access token. +func (c *Client) authQuery(q url.Values) { + if c.accessToken == "" { + q.Set("app_key", c.appKey) + } +} + +// authHeader adds the Bearer header when the client uses an access token. +func (c *Client) authHeader(req *http.Request) { + if c.accessToken != "" { + req.Header.Set("Authorization", "Bearer "+c.accessToken) + } +} + // Response wraps http.Response and surfaces Flashduty envelope metadata: the // request id (for support), pagination fields when the endpoint returns them, // and best-effort rate-limit signals. @@ -108,7 +139,7 @@ func (c *Client) newRequestWithAppKey(ctx context.Context, method, path string, u := c.BaseURL.ResolveReference(rel) if withAppKey { q := u.Query() - q.Set("app_key", c.appKey) + c.authQuery(q) u.RawQuery = q.Encode() } @@ -130,6 +161,9 @@ func (c *Client) newRequestWithAppKey(ctx context.Context, method, path string, req.Header.Set("Content-Type", "application/json") } req.Header.Set("Accept", "application/json") + if withAppKey { + c.authHeader(req) + } if c.UserAgent != "" { req.Header.Set("User-Agent", c.UserAgent) } diff --git a/upload.go b/upload.go index f58f498..24167c1 100644 --- a/upload.go +++ b/upload.go @@ -57,7 +57,7 @@ func (c *Client) uploadFile(ctx context.Context, path string, query url.Values, } u := c.BaseURL.ResolveReference(rel) q := u.Query() - q.Set("app_key", c.appKey) + c.authQuery(q) for k, vs := range query { for _, v := range vs { q.Set(k, v) @@ -71,6 +71,7 @@ func (c *Client) uploadFile(ctx context.Context, path string, query url.Values, } req.Header.Set("Content-Type", mw.FormDataContentType()) req.Header.Set("Accept", "application/json") + c.authHeader(req) if c.UserAgent != "" { req.Header.Set("User-Agent", c.UserAgent) }