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
123 changes: 123 additions & 0 deletions access_token_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
}
36 changes: 35 additions & 1 deletion flashduty.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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 <token>"
// 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.
Expand Down Expand Up @@ -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()
}

Expand All @@ -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)
}
Expand Down
3 changes: 2 additions & 1 deletion upload.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
}
Expand Down
Loading