diff --git a/README.md b/README.md index 2239a3a..09d43d1 100644 --- a/README.md +++ b/README.md @@ -99,7 +99,7 @@ Before using most `nbctl` commands, you need to configure your Nudgebee API cred This command interactively guides you through setting up a new configuration profile or updating an existing one. If `profile-name` is not provided, it defaults to `default`. You will be prompted for: * **Nudgebee API Endpoint**: The URL of the Nudgebee API (e.g., `https://api.nudgebee.com`). -* **Nudgebee API Key**: Your personal API key for authentication. +* **Nudgebee API Key**: Your personal API token (`sk-nb-…`), created under **Settings → API Tokens**. `nbctl` sends it directly as a Bearer token on every request. Tokens created before direct token auth was supported are rejected with a 401 and must be recreated. * **Nudgebee Username**: Your Nudgebee account username (e.g., your email). * **Default Account ID**: The ID of the Nudgebee account you wish to interact with by default. diff --git a/TESTING.md b/TESTING.md index 0739ac3..6d39850 100644 --- a/TESTING.md +++ b/TESTING.md @@ -8,8 +8,8 @@ Helpers The package `pkg/testutil` exposes the following useful helpers: - `RunWithSimpleGraphQL(mockData any, cmd *cobra.Command, args []string) (string, error)` - - Convenience for mocking a GraphQL response. Automatically mocks `/api/auth/token` and - returns `{ "data": mockData }` at `/api/graphql`. Useful for simple, static tests. + - Convenience for mocking a GraphQL response. Returns `{ "data": mockData }` at + `/api/graphql`. Useful for simple, static tests. - `RunWithMockServer(handler http.HandlerFunc, viperOverrides map[string]any, cmd *cobra.Command, args []string) (string, error)` - More flexible: provide a handler to simulate complex behavior, and a small map of viper overrides (e.g. `api-key`, `username`). The helper sets the `endpoint` viper key to the test server URL and restores previous viper values after the test. diff --git a/cmd/configure.go b/cmd/configure.go index 1dbd2dd..2674a71 100644 --- a/cmd/configure.go +++ b/cmd/configure.go @@ -102,7 +102,7 @@ var configureAddCmd = &cobra.Command{ // Validate the configuration by making a simple API call fmt.Println("Validating configuration...") - gqlClient := client.NewClient(client.WithApiKey(apiKey), client.WithEndpoint(endpoint), client.WithUsername(username)) + gqlClient := client.NewClient(client.WithApiKey(apiKey), client.WithEndpoint(endpoint)) req := client.NewRequest(` query { cloud_accounts: accounts_list(where: {}, limit: 1, offset: 0) { diff --git a/cmd/nubi_test.go b/cmd/nubi_test.go index b1827e9..00822ef 100644 --- a/cmd/nubi_test.go +++ b/cmd/nubi_test.go @@ -76,8 +76,6 @@ func TestNubiCmd_AsyncQuery_TriggerError_JSON(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": w.WriteHeader(http.StatusOK) _ = json.NewEncoder(w).Encode(map[string]any{ @@ -116,8 +114,6 @@ func TestNubiCmd_Query_AccessDenied_Text(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": w.WriteHeader(http.StatusOK) _ = json.NewEncoder(w).Encode(map[string]any{ @@ -162,8 +158,6 @@ func TestNubiCmd_SyncQuery(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": resp := map[string]interface{}{ "data": map[string]interface{}{ @@ -225,8 +219,6 @@ func TestNubiCmd_SyncQuery_JSON(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": resp := map[string]interface{}{ "data": map[string]interface{}{ @@ -294,8 +286,6 @@ func TestNubiCmd_Query_Timeout(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": resp := map[string]interface{}{ "data": map[string]interface{}{ @@ -343,8 +333,6 @@ func TestNubiCmd_Query_Timeout_JSON(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": resp := map[string]interface{}{ "data": map[string]interface{}{ @@ -396,8 +384,6 @@ func TestNubiCmd_Query_TransientRetry(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": pollCount++ if pollCount == 2 { @@ -815,8 +801,6 @@ func TestNubiCmd_Get_WithAccountId(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": var payload struct { Query string `json:"query"` @@ -882,8 +866,6 @@ func TestNubiCmd_Get_SessionId_WithAccountId(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": var payload struct { Query string `json:"query"` @@ -962,8 +944,6 @@ func TestNubiCmd_SyncQuery_SessionIdDiffersFromConversationId(t *testing.T) { handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") switch r.URL.Path { - case "/api/auth/token": - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) case "/api/graphql": var payload struct { Query string `json:"query"` diff --git a/cmd/workflow_apply_test.go b/cmd/workflow_apply_test.go index ca037be..5617aab 100644 --- a/cmd/workflow_apply_test.go +++ b/cmd/workflow_apply_test.go @@ -34,11 +34,6 @@ definition: require.NoError(t, tmpFile.Close()) handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` @@ -108,11 +103,6 @@ definition: require.NoError(t, tmpFile.Close()) handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` diff --git a/cmd/workflow_cancel_execution_test.go b/cmd/workflow_cancel_execution_test.go index 80d7023..3e46194 100644 --- a/cmd/workflow_cancel_execution_test.go +++ b/cmd/workflow_cancel_execution_test.go @@ -17,11 +17,6 @@ func TestWorkflowCancelExecutionCmd(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` @@ -62,11 +57,6 @@ func TestWorkflowCancelExecutionCmd_JSON(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]any{ diff --git a/cmd/workflow_delete_test.go b/cmd/workflow_delete_test.go index b31cc41..899e639 100644 --- a/cmd/workflow_delete_test.go +++ b/cmd/workflow_delete_test.go @@ -17,11 +17,6 @@ func TestWorkflowDeleteCmd(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` @@ -62,11 +57,6 @@ func TestWorkflowDeleteCmd_JSON(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]any{ diff --git a/cmd/workflow_pause_test.go b/cmd/workflow_pause_test.go index 201d028..9fb0cf8 100644 --- a/cmd/workflow_pause_test.go +++ b/cmd/workflow_pause_test.go @@ -17,11 +17,6 @@ func TestWorkflowPauseCmd(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` @@ -62,11 +57,6 @@ func TestWorkflowPauseCmd_JSON(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]any{ diff --git a/cmd/workflow_resume_test.go b/cmd/workflow_resume_test.go index 92dce20..0c96c81 100644 --- a/cmd/workflow_resume_test.go +++ b/cmd/workflow_resume_test.go @@ -17,11 +17,6 @@ func TestWorkflowResumeCmd(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { var reqBody struct { Query string `json:"query"` @@ -62,11 +57,6 @@ func TestWorkflowResumeCmd_JSON(t *testing.T) { defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]any{ diff --git a/cmd/workflow_validate_test.go b/cmd/workflow_validate_test.go index 27a4519..9b5f604 100644 --- a/cmd/workflow_validate_test.go +++ b/cmd/workflow_validate_test.go @@ -68,11 +68,6 @@ definition: } handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") // Return top-level errors @@ -186,11 +181,6 @@ definition: // We need to use RunWithMockServer to simulate top-level errors (not wrapped in data) handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } if r.URL.Path == "/api/graphql" { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) diff --git a/pkg/client/client.go b/pkg/client/client.go index e8413a1..e63e1be 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -48,6 +48,7 @@ func (r *Request) Header(key, value string) { // Client is a GraphQL client. type Client struct { endpoint string + apiKey string httpClient *http.Client } @@ -103,29 +104,12 @@ func (t *loggingTransport) RoundTrip(req *http.Request) (*http.Response, error) type clientOptions struct { endpoint string apiKey string - username string } type ClientOption interface { apply(opts *clientOptions) } -type clientUsernameOption struct { - username string -} - -func (o clientUsernameOption) apply(opts *clientOptions) { - if o.username != "" { - opts.username = o.username - } -} - -func WithUsername(username string) ClientOption { - return clientUsernameOption{ - username: username, - } -} - type clientApiKeyOption struct { apiKey string } @@ -158,287 +142,118 @@ func WithEndpoint(endpoint string) ClientOption { } } -// NewClient creates a new GraphQL client. -func NewClient(opts ...ClientOption) *Client { +// resolveOptions applies opts and falls back to viper config and defaults. +// The returned endpoint has no trailing slash. +func resolveOptions(opts []ClientOption) clientOptions { config := clientOptions{} for _, o := range opts { o.apply(&config) } - endpoint := config.endpoint - if endpoint == "" { - endpoint = viper.GetString("endpoint") - } - if endpoint == "" { - endpoint = "https://app.nudgebee.com" + if config.endpoint == "" { + config.endpoint = viper.GetString("endpoint") } - if endpoint[len(endpoint)-1] == '/' { - endpoint = endpoint[:len(endpoint)-1] + if config.endpoint == "" { + config.endpoint = "https://app.nudgebee.com" } - graphqlEndpoint := endpoint + "/api/graphql" + config.endpoint = strings.TrimRight(config.endpoint, "/") - apiKey := config.apiKey - if apiKey == "" { - apiKey = viper.GetString("api-key") + if config.apiKey == "" { + config.apiKey = viper.GetString("api-key") } + return config +} - username := config.username - if username == "" { - username = viper.GetString("username") - } - tokenEndpoint := endpoint + "/api/auth/token" - - // create a new http client with the auth header - // transport that injects bearer tokens obtained from token endpoint - transport := &authTransport{ - apiKey: apiKey, - username: username, - tokenEndpoint: tokenEndpoint, - wrapped: http.DefaultTransport, - httpClient: &http.Client{Timeout: 30 * time.Second}, - } +var ( + verboseLogger *log.Logger + verboseLoggerOnce sync.Once +) - var finalTransport http.RoundTripper = transport - verbose := viper.GetBool("verbose") - if verbose { +// getVerboseLogger opens nbctl_graphql.log once per process, so long-running +// callers (e.g. the MCP server) that build many clients don't leak descriptors. +func getVerboseLogger() *log.Logger { + verboseLoggerOnce.Do(func() { logFile, err := os.OpenFile("nbctl_graphql.log", os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o644) if err != nil { log.Printf("Error opening log file: %v\n", err) - } else { - logger := log.New(logFile, "", log.LstdFlags) - finalTransport = &loggingTransport{ - wrapped: transport, - logger: logger, - } + return } - } - - httpClient := &http.Client{ - Transport: finalTransport, - Timeout: 30 * time.Second, - } - - return &Client{ - endpoint: graphqlEndpoint, - httpClient: httpClient, - } + verboseLogger = log.New(logFile, "", log.LstdFlags) + }) + return verboseLogger } -// NewHTTPClient creates a new authenticated http.Client. -func NewHTTPClient(opts ...ClientOption) *http.Client { - config := clientOptions{} - for _, o := range opts { - o.apply(&config) - } - - endpoint := config.endpoint - if endpoint == "" { - endpoint = viper.GetString("endpoint") - } - if endpoint == "" { - endpoint = "https://app.nudgebee.com" - } - if endpoint[len(endpoint)-1] == '/' { - endpoint = endpoint[:len(endpoint)-1] - } - - apiKey := config.apiKey - if apiKey == "" { - apiKey = viper.GetString("api-key") - } - - username := config.username - if username == "" { - username = viper.GetString("username") - } - tokenEndpoint := endpoint + "/api/auth/token" - - transport := &authTransport{ - apiKey: apiKey, - username: username, - tokenEndpoint: tokenEndpoint, - wrapped: http.DefaultTransport, - httpClient: &http.Client{Timeout: 30 * time.Second}, +// newTransport returns the authenticating transport, wrapped in a request +// logger when --verbose is set. +func newTransport(apiKey string) http.RoundTripper { + var transport http.RoundTripper = &authTransport{ + apiKey: apiKey, + wrapped: http.DefaultTransport, } - var finalTransport http.RoundTripper = transport - verbose := viper.GetBool("verbose") - if verbose { - logFile, err := os.OpenFile("nbctl_graphql.log", os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o644) - if err != nil { - log.Printf("Error opening log file: %v\n", err) - } else { - logger := log.New(logFile, "", log.LstdFlags) - finalTransport = &loggingTransport{ + if viper.GetBool("verbose") { + if logger := getVerboseLogger(); logger != nil { + transport = &loggingTransport{ wrapped: transport, logger: logger, } } } + return transport +} +func newHTTPClient(config clientOptions) *http.Client { return &http.Client{ - Transport: finalTransport, + Transport: newTransport(config.apiKey), Timeout: 30 * time.Second, } } -type authTransport struct { - apiKey string - username string - tokenEndpoint string - wrapped http.RoundTripper - - // http client used to fetch tokens - httpClient *http.Client - - // cached token state - mu sync.Mutex - accessToken string - expiry time.Time -} - -func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { - // ensure we have a valid access token - token, err := t.getAccessToken(req.Context()) - if err != nil { - return nil, err - } - - // avoid mutating original request - req2 := cloneRequest(req) - req2.Header.Set("Authorization", "Bearer "+token) - - resp, err := t.wrapped.RoundTrip(req2) - if err != nil { - return resp, err - } - - // if unauthorized, try refreshing token and retry once - if resp.StatusCode == http.StatusUnauthorized { - // discard body from first response; handle potential copy error - if _, err := io.Copy(io.Discard, resp.Body); err != nil { - _ = err // intentionally ignore copy error - } - if err := resp.Body.Close(); err != nil { - _ = err // intentionally ignore close error - } - - // force refresh - if err := t.forceRefresh(req.Context()); err != nil { - return nil, err - } - - // retry request with new token - token, err = t.getAccessToken(req.Context()) - if err != nil { - return nil, err - } - - req3 := cloneRequest(req) - req3.Header.Set("Authorization", "Bearer "+token) - return t.wrapped.RoundTrip(req3) +// NewClient creates a new GraphQL client. +func NewClient(opts ...ClientOption) *Client { + config := resolveOptions(opts) + return &Client{ + endpoint: config.endpoint + "/api/graphql", + apiKey: config.apiKey, + httpClient: newHTTPClient(config), } - - return resp, nil -} - -// cloneRequest creates a deep copy of the request, including the Header -func cloneRequest(r *http.Request) *http.Request { - // Clone returns a deep copy of r with its context changed to ctx. - // The Request.Header map is also deep copied. - return r.Clone(r.Context()) } -// tokenResponse models the expected JSON response from the token endpoint. -// We accept multiple common field names. -type tokenResponse struct { - Token string `json:"token"` - Expiry int64 `json:"expiry"` +// NewHTTPClient creates a new authenticated http.Client. +func NewHTTPClient(opts ...ClientOption) *http.Client { + return newHTTPClient(resolveOptions(opts)) } -// getAccessToken returns a valid access token, fetching a new one if needed. -func (t *authTransport) getAccessToken(ctx context.Context) (string, error) { - t.mu.Lock() - token := t.accessToken - exp := t.expiry - t.mu.Unlock() - - if token != "" && time.Now().Before(exp) { - return token, nil - } +// ApiTokenPrefix is the prefix on every Nudgebee API token (`sk-nb-...`). +const ApiTokenPrefix = "sk-nb-" - // fetch new token - if err := t.fetchToken(ctx); err != nil { - return "", err - } - - t.mu.Lock() - defer t.mu.Unlock() - return t.accessToken, nil -} - -// forceRefresh forces fetching a new token regardless of cached expiry -func (t *authTransport) forceRefresh(ctx context.Context) error { - return t.fetchToken(ctx) +// authTransport sends the configured API token as a Bearer on every request. +// The gateway authenticates the raw token directly; there is no exchange step. +type authTransport struct { + apiKey string + wrapped http.RoundTripper } -// fetchToken calls the token endpoint with the API key and stores the access token and expiry -func (t *authTransport) fetchToken(ctx context.Context) error { - // prepare request. We POST an empty JSON body and include the api-key in header 'X-Api-Key'. - // NOTE: This is an assumption; if your token endpoint expects a different shape (e.g. JSON body), - // adjust accordingly or set a custom token-endpoint implementation. - bodyBytes, err := json.Marshal(map[string]string{ - "email": t.username, - "secret": t.apiKey, - }) - if err != nil { - return err - } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, t.tokenEndpoint, bytes.NewReader(bodyBytes)) - if err != nil { - return err - } - req.Header.Set("Content-Type", "application/json") - - resp, err := t.httpClient.Do(req) - if err != nil { - return err - } - defer func() { - _ = resp.Body.Close() - }() - - if resp.StatusCode < 200 || resp.StatusCode >= 300 { - return errors.New("token endpoint returned non-2xx status: " + resp.Status) - } - - var tr tokenResponse - dec := json.NewDecoder(resp.Body) - if err := dec.Decode(&tr); err != nil { - return err +func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if t.apiKey == "" { + return nil, errors.New("no API key configured: run 'nbctl configure add' or set api-key") } + // avoid mutating original request + req2 := req.Clone(req.Context()) + req2.Header.Set("Authorization", "Bearer "+t.apiKey) + return t.wrapped.RoundTrip(req2) +} - token := tr.Token - if token == "" { - return errors.New("token endpoint did not return access_token") +// unauthorizedError explains a 401 from the gateway. The most likely causes are +// a deleted/expired token, or a token created before direct-token auth existed +// (those must be recreated). +func unauthorizedError(apiKey string) error { + msg := "authentication failed (401 Unauthorized): the API key was rejected" + if !strings.HasPrefix(apiKey, ApiTokenPrefix) { + msg += fmt.Sprintf("; API keys must start with %q", ApiTokenPrefix) } - - // calculate expiry - var expiry time.Time - if tr.Expiry > 0 { - // subtract small buffer - expiry = time.Now().Add(time.Duration(tr.Expiry)*time.Second - 10*time.Second) - } else { - // default to 55 minutes - expiry = time.Now().Add(55 * time.Minute) - } - - t.mu.Lock() - t.accessToken = token - t.expiry = expiry - t.mu.Unlock() - - return nil + return errors.New(msg + ". The token may be deleted, expired, or created before direct token auth was supported. " + + "Create a new token under Settings → API Tokens and run 'nbctl configure add'") } type GraphQLError struct { @@ -520,6 +335,10 @@ func (c *Client) Run(ctx context.Context, req *Request, resp any) error { _ = httpResp.Body.Close() }() + if httpResp.StatusCode == http.StatusUnauthorized { + return unauthorizedError(c.apiKey) + } + // 4. Decode Response // We want to handle errors specifically, so we decode into a raw map first or a struct with Errors. // We'll use a struct that captures Data as RawMessage to allow unmarshalling later. diff --git a/pkg/client/client_test.go b/pkg/client/client_test.go index 33de2cb..a96bc4b 100644 --- a/pkg/client/client_test.go +++ b/pkg/client/client_test.go @@ -2,133 +2,86 @@ package client import ( "context" - "encoding/json" "net/http" "net/http/httptest" + "strings" "testing" - "time" "github.com/spf13/viper" ) -func TestTokenFetchAndAuthHeader(t *testing.T) { - // token server returns a token with short expiry - tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "token": "test-token-1", - "expiry": 2, // seconds - }) - })) - defer tokenSrv.Close() - - // graphQL server checks Authorization header - var seenToken string - callCount := 0 +func TestApiKeySentAsBearer(t *testing.T) { + var seenAuth []string apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - callCount++ - seenToken = r.Header.Get("Authorization") - // respond OK - w.WriteHeader(200) + if r.URL.Path != "/api/graphql" { + t.Errorf("unexpected request to %s; no token exchange should happen", r.URL.Path) + } + seenAuth = append(seenAuth, r.Header.Get("Authorization")) w.Header().Set("Content-Type", "application/json") _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) })) defer apiSrv.Close() - // configure viper - viper.Set("endpoint", tokenSrv.URL) // NewClient will append /api/graphql and /api/token, but tests will override below - viper.Set("api-key", "dummy-key") - - // Build a client but override transport endpoints to point to test servers - transport := &authTransport{ - apiKey: "dummy-key", - tokenEndpoint: tokenSrv.URL, - wrapped: http.DefaultTransport, - httpClient: &http.Client{Timeout: 5 * time.Second}, - } - httpClient := &http.Client{Transport: transport} - client := &Client{ - endpoint: apiSrv.URL, - httpClient: httpClient, - } - - // make a request - this should trigger token fetch - req := NewRequest(`query { ok }`) - var resp map[string]any - if err := client.Run(context.Background(), req, &resp); err != nil { - t.Fatalf("run failed: %v", err) - } - - if seenToken != "Bearer test-token-1" { - t.Fatalf("expected Authorization header to contain fetched token, got %q", seenToken) + client := NewClient(WithEndpoint(apiSrv.URL+"/"), WithApiKey("sk-nb-test")) + for i := 0; i < 2; i++ { + var resp map[string]any + if err := client.Run(context.Background(), NewRequest(`query { ok }`), &resp); err != nil { + t.Fatalf("run %d failed: %v", i, err) + } } - // wait until token expires - time.Sleep(3 * time.Second) - - // next request should fetch a fresh token from token server (which always returns test-token-1) - req2 := NewRequest(`query { ok }`) - if err := client.Run(context.Background(), req2, &resp); err != nil { - t.Fatalf("second run failed: %v", err) + if len(seenAuth) != 2 { + t.Fatalf("expected 2 requests, got %d", len(seenAuth)) } - - if callCount < 2 { - t.Fatalf("expected at least two GraphQL calls, got %d", callCount) + for _, auth := range seenAuth { + if auth != "Bearer sk-nb-test" { + t.Fatalf("expected raw API key as Bearer, got %q", auth) + } } } -func TestRetryOn401(t *testing.T) { - // token server: first returns token1, then token2 - tokens := []string{"token1", "token2"} - tokenSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - tkn := tokens[0] - // rotate - if len(tokens) > 1 { - tokens = tokens[1:] - } - _ = json.NewEncoder(w).Encode(map[string]any{"token": tkn, "expiry": 60}) +func TestUnauthorizedReturnsHint(t *testing.T) { + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":"Unauthorized"}`)) })) - defer tokenSrv.Close() + defer apiSrv.Close() - // GraphQL server: first request with token1 responds 401, second with token2 responds 200 - call := 0 - apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - call++ - auth := r.Header.Get("Authorization") - if call == 1 { - // first call: token1 -> return 401 - if auth != "Bearer token1" { - t.Fatalf("expected first call to use token1, got %q", auth) - } - w.WriteHeader(401) - return + t.Run("sk-nb key", func(t *testing.T) { + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("sk-nb-revoked")) + err := client.Run(context.Background(), NewRequest(`query { ok }`), nil) + if err == nil || !strings.Contains(err.Error(), "401") || !strings.Contains(err.Error(), "Create a new token") { + t.Fatalf("expected 401 hint, got %v", err) } - // subsequent calls should carry token2 - if auth != "Bearer token2" { - t.Fatalf("expected retry to use token2, got %q", auth) + if strings.Contains(err.Error(), "must start with") { + t.Fatalf("did not expect prefix hint for sk-nb key, got %v", err) } - w.WriteHeader(200) - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) + }) + + t.Run("legacy key", func(t *testing.T) { + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("old-secret")) + err := client.Run(context.Background(), NewRequest(`query { ok }`), nil) + if err == nil || !strings.Contains(err.Error(), `must start with "sk-nb-"`) { + t.Fatalf("expected prefix hint, got %v", err) + } + }) +} + +func TestMissingApiKey(t *testing.T) { + viper.Set("api-key", "") + defer viper.Set("api-key", nil) + called := false + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true })) defer apiSrv.Close() - transport := &authTransport{ - apiKey: "dummy-key", - tokenEndpoint: tokenSrv.URL, - wrapped: http.DefaultTransport, - httpClient: &http.Client{Timeout: 5 * time.Second}, - } - httpClient := &http.Client{Transport: transport} - client := &Client{ - endpoint: apiSrv.URL, - httpClient: httpClient, + err := NewClient(WithEndpoint(apiSrv.URL)).Run(context.Background(), NewRequest(`query { ok }`), nil) + if err == nil || !strings.Contains(err.Error(), "no API key configured") { + t.Fatalf("expected missing key error, got %v", err) } - - req := NewRequest(`query { ok }`) - var resp map[string]any - if err := client.Run(context.Background(), req, &resp); err != nil { - t.Fatalf("run failed: %v", err) + if called { + t.Fatal("request should not be sent without an API key") } } @@ -137,7 +90,6 @@ func TestNewClient(t *testing.T) { client := NewClient( WithEndpoint("http://test.com"), WithApiKey("test-key"), - WithUsername("test-user"), ) if client == nil { t.Fatal("expected client to be non-nil") @@ -150,7 +102,6 @@ func TestNewClient(t *testing.T) { t.Run("with viper", func(t *testing.T) { viper.Set("endpoint", "http://viper.com") viper.Set("api-key", "viper-key") - viper.Set("username", "viper-user") // no need to reset viper as it's global and tests run in sequence or we just let it be // better to use cleanup if we care about other tests, but here we are ok. diff --git a/pkg/nubi/nubi_test.go b/pkg/nubi/nubi_test.go index 2374635..b9b5cbf 100644 --- a/pkg/nubi/nubi_test.go +++ b/pkg/nubi/nubi_test.go @@ -14,19 +14,8 @@ import ( ) func newTestNubiClient(handler http.HandlerFunc) (*NubiClient, func()) { - // Wrapper handler to handle token endpoint - wrappedHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/api/auth/token" { - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return - } - // Delegate to the test-specific handler for other requests (presumably graphql) - handler(w, r) - }) - - srv := httptest.NewServer(wrappedHandler) - c := client.NewClient(client.WithEndpoint(srv.URL)) + srv := httptest.NewServer(handler) + c := client.NewClient(client.WithEndpoint(srv.URL), client.WithApiKey("sk-nb-test")) nubiClient := New(c, "test-account", "test-user", "test-session", srv.URL) return nubiClient, srv.Close } diff --git a/pkg/testutil/helpers.go b/pkg/testutil/helpers.go index 8acb3fb..144a8fd 100644 --- a/pkg/testutil/helpers.go +++ b/pkg/testutil/helpers.go @@ -110,7 +110,7 @@ func RunWithMockServer(handler http.HandlerFunc, viperOverrides map[string]any, } // RunWithSimpleGraphQL is a convenience helper that starts a mock server which -// automatically handles the token endpoint and serves the provided mockData as +// serves the provided mockData as // the GraphQL response payload under the "data" field. It's suitable for the // common case where tests only need to mock returned data. func RunWithSimpleGraphQL(mockData any, cmd *cobra.Command, args []string) (string, error) { @@ -118,10 +118,6 @@ func RunWithSimpleGraphQL(mockData any, cmd *cobra.Command, args []string) (stri defer func() { _ = os.Unsetenv("NBCTL_TESTING") }() handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { - case "/api/auth/token": - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{"token": "fake-token", "expiry": 3600}) - return case "/api/graphql": w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]any{"data": mockData})