diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index aefe144..51921ed 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -29,7 +29,7 @@ behavior to **legal@nudgebee.com**. ``` cmd/ Cobra command definitions (one file per command) -pkg/client/ GraphQL client, auth/token handling +pkg/client/ GraphQL client, API-key auth pkg/config/ Profile + viper config loading pkg/format/ Tabular / JSON output rendering pkg/log/ Logger helpers diff --git a/README.md b/README.md index 2239a3a..c58adaa 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 key for authentication, created under Settings → API Tokens. nbctl sends it as the bearer token on every request. Keys start with `sk-nb-`; current Nudgebee servers no longer accept a key without that prefix, so create a new one. * **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/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..0c67e3b 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 } @@ -188,8 +189,8 @@ func NewClient(opts ...ClientOption) *Client { } tokenEndpoint := endpoint + "/api/auth/token" - // create a new http client with the auth header - // transport that injects bearer tokens obtained from token endpoint + // create a new http client that sends the API key as the bearer token, + // falling back to the token endpoint on servers that predate that transport := &authTransport{ apiKey: apiKey, username: username, @@ -206,6 +207,8 @@ func NewClient(opts ...ClientOption) *Client { log.Printf("Error opening log file: %v\n", err) } else { logger := log.New(logFile, "", log.LstdFlags) + // Wraps the auth transport, so it logs each request before the + // Authorization header is added and the API key never reaches the file. finalTransport = &loggingTransport{ wrapped: transport, logger: logger, @@ -220,11 +223,15 @@ func NewClient(opts ...ClientOption) *Client { return &Client{ endpoint: graphqlEndpoint, + apiKey: apiKey, httpClient: httpClient, } } -// NewHTTPClient creates a new authenticated http.Client. +// NewHTTPClient creates a new authenticated http.Client. A request with a body +// must be replayable, i.e. have GetBody set (http.NewRequest does this for +// bytes and strings readers): against a server that predates direct API-key +// auth, the first attempt is refused and the request is sent again. func NewHTTPClient(opts ...ClientOption) *http.Client { config := clientOptions{} for _, o := range opts { @@ -282,6 +289,19 @@ func NewHTTPClient(opts ...ClientOption) *http.Client { } } +// apiKeyPrefix and gatewayKeyPrefix are the formats Nudgebee issues API keys +// in. A current server accepts a platform key (sk-nb-…) directly as the bearer +// token; an AI Gateway key (sk-nb-gw-…) only works against the AI Gateway. +const ( + apiKeyPrefix = "sk-nb-" + gatewayKeyPrefix = "sk-nb-gw-" +) + +// authTransport sends the API key itself as the bearer token. A server that +// predates direct API-key auth refuses that with a 401 but still exchanges the +// key for a session token at tokenEndpoint, so on the first refusal the +// transport tries the exchange and, if it works, uses exchanged tokens from +// then on. Drop the exchange once no supported server predates direct auth. type authTransport struct { apiKey string username string @@ -293,11 +313,47 @@ type authTransport struct { // cached token state mu sync.Mutex + legacy bool accessToken string expiry time.Time } func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { + t.mu.Lock() + legacy := t.legacy + t.mu.Unlock() + if legacy { + return t.roundTripExchanged(req) + } + + req2, err := replayableRequest(req) + if err != nil { + return nil, err + } + req2.Header.Set("Authorization", "Bearer "+t.apiKey) + resp, err := t.wrapped.RoundTrip(req2) + if err != nil || resp.StatusCode != http.StatusUnauthorized { + return resp, err + } + + // Refused: either the key is bad, or the server predates direct API-key + // auth. Only the exchange can tell the two apart. If it fails too, the + // original 401 stands. + if err := t.fetchToken(req.Context()); err != nil { + return resp, nil + } + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + + t.mu.Lock() + t.legacy = true + t.mu.Unlock() + return t.roundTripExchanged(req) +} + +// roundTripExchanged authenticates with a session token obtained from the +// token endpoint, for servers that predate direct API-key auth. +func (t *authTransport) roundTripExchanged(req *http.Request) (*http.Response, error) { // ensure we have a valid access token token, err := t.getAccessToken(req.Context()) if err != nil { @@ -305,7 +361,10 @@ func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { } // avoid mutating original request - req2 := cloneRequest(req) + req2, err := replayableRequest(req) + if err != nil { + return nil, err + } req2.Header.Set("Authorization", "Bearer "+token) resp, err := t.wrapped.RoundTrip(req2) @@ -334,7 +393,10 @@ func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { return nil, err } - req3 := cloneRequest(req) + req3, err := replayableRequest(req) + if err != nil { + return nil, err + } req3.Header.Set("Authorization", "Bearer "+token) return t.wrapped.RoundTrip(req3) } @@ -342,6 +404,22 @@ func (t *authTransport) RoundTrip(req *http.Request) (*http.Response, error) { return resp, nil } +// replayableRequest clones req with a body of its own. A request can be sent +// more than once here (after a 401), and a plain clone shares the body reader +// an earlier attempt may already have drained, which fails with +// "ContentLength=N with Body length 0". +func replayableRequest(req *http.Request) (*http.Request, error) { + r := cloneRequest(req) + if req.GetBody != nil { + body, err := req.GetBody() + if err != nil { + return nil, err + } + r.Body = body + } + return r, 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. @@ -349,6 +427,21 @@ func cloneRequest(r *http.Request) *http.Request { return r.Clone(r.Context()) } +// unauthorizedError explains a 401. The app's 401 body carries neither `data` +// nor `errors`, so without this a rejected key would look like an empty result. +func unauthorizedError(apiKey string) error { + switch { + case apiKey == "": + return errors.New("not authenticated: no API key configured; run `nbctl configure`") + case strings.HasPrefix(apiKey, gatewayKeyPrefix): + return errors.New("not authenticated: this is an AI Gateway key, which only works against the AI Gateway; create a platform API key in Settings → API Tokens and run `nbctl configure`") + case !strings.HasPrefix(apiKey, apiKeyPrefix): + return errors.New("not authenticated: this key predates the \"sk-nb-\" format, which current Nudgebee servers no longer accept; create a new key in Settings → API Tokens and run `nbctl configure`") + default: + return errors.New("not authenticated: the API key was rejected; it may have been deleted, suspended or expired") + } +} + // tokenResponse models the expected JSON response from the token endpoint. // We accept multiple common field names. type tokenResponse struct { @@ -520,6 +613,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..fc685ac 100644 --- a/pkg/client/client_test.go +++ b/pkg/client/client_test.go @@ -3,8 +3,11 @@ package client import ( "context" "encoding/json" + "io" "net/http" "net/http/httptest" + "os" + "strings" "testing" "time" @@ -42,6 +45,7 @@ func TestTokenFetchAndAuthHeader(t *testing.T) { // Build a client but override transport endpoints to point to test servers transport := &authTransport{ apiKey: "dummy-key", + legacy: true, tokenEndpoint: tokenSrv.URL, wrapped: http.DefaultTransport, httpClient: &http.Client{Timeout: 5 * time.Second}, @@ -115,6 +119,7 @@ func TestRetryOn401(t *testing.T) { transport := &authTransport{ apiKey: "dummy-key", + legacy: true, tokenEndpoint: tokenSrv.URL, wrapped: http.DefaultTransport, httpClient: &http.Client{Timeout: 5 * time.Second}, @@ -132,6 +137,184 @@ func TestRetryOn401(t *testing.T) { } } +func TestSendsApiKeyAsBearer(t *testing.T) { + var paths []string + var seenAuth string + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + paths = append(paths, r.URL.Path) + seenAuth = r.Header.Get("Authorization") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) + })) + defer apiSrv.Close() + + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("sk-nb-test-key")) + + var resp map[string]any + if err := client.Run(context.Background(), NewRequest(`query { ok }`), &resp); err != nil { + t.Fatalf("run failed: %v", err) + } + + // A current server accepts the key itself: no exchange call precedes the request. + if len(paths) != 1 || paths[0] != "/api/graphql" { + t.Fatalf("expected a single call to /api/graphql, got %v", paths) + } + if seenAuth != "Bearer sk-nb-test-key" { + t.Fatalf("expected the API key as the bearer token, got %q", seenAuth) + } +} + +func TestFallsBackToExchangeOnOlderServer(t *testing.T) { + // An older server refuses the key as a bearer but exchanges it for a session token. + directAttempts, exchanges := 0, 0 + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/auth/token": + exchanges++ + _ = json.NewEncoder(w).Encode(map[string]any{"token": "session-token", "expiry": 3600}) + case "/api/graphql": + if r.Header.Get("Authorization") != "Bearer session-token" { + directAttempts++ + w.WriteHeader(http.StatusUnauthorized) + return + } + // The retried request must still carry the query. + body, _ := io.ReadAll(r.Body) + if !strings.Contains(string(body), "query { ok }") { + t.Errorf("expected the retried request to carry the query, got %q", body) + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) + } + })) + defer apiSrv.Close() + + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("sk-nb-test-key"), WithUsername("a@b.com")) + + 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) + } + if resp["ok"] != true { + t.Fatalf("run %d: expected data, got %v", i, resp) + } + } + + // Once the server is known to need the exchange, later requests go straight to it. + if directAttempts != 1 || exchanges != 1 { + t.Fatalf("expected 1 direct attempt and 1 exchange, got %d and %d", directAttempts, exchanges) + } +} + +func TestFallbackResendsBodyUnderConcurrency(t *testing.T) { + // Several first requests hit an older server at once. Each one is refused, + // exchanged and re-sent, and every re-send must carry its full body. + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/api/auth/token": + _ = json.NewEncoder(w).Encode(map[string]any{"token": "session-token", "expiry": 3600}) + case "/api/graphql": + if r.Header.Get("Authorization") != "Bearer session-token" { + w.WriteHeader(http.StatusUnauthorized) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) + } + })) + defer apiSrv.Close() + + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("sk-nb-test-key"), WithUsername("a@b.com")) + query := `query { ok }` + strings.Repeat(" ", 20000) + + const workers = 8 + errs := make(chan error, workers) + for i := 0; i < workers; i++ { + go func() { + var resp map[string]any + errs <- client.Run(context.Background(), NewRequest(query), &resp) + }() + } + for i := 0; i < workers; i++ { + if err := <-errs; err != nil { + t.Errorf("concurrent fallback request failed: %v", err) + } + } +} + +func TestUnauthorizedIsAnError(t *testing.T) { + // The app's 401 body has neither `data` nor `errors`, so decoding it alone + // would report success with an empty result. A current server no longer has + // the token endpoint, so the fallback fails too. + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + if r.URL.Path == "/api/auth/token" { + w.WriteHeader(http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":"not_authenticated","description":"The user does not have an active session"}`)) + })) + defer apiSrv.Close() + + cases := []struct { + name string + apiKey string + hint string + }{ + {"no key", "", "no API key configured"}, + {"key from before the sk-nb- format", "0123456789abcdef", "predates"}, + {"AI Gateway key", "sk-nb-gw-abc", "AI Gateway key"}, + {"deleted or expired key", "sk-nb-abc", "rejected"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey(tc.apiKey)) + // NewClient falls back to viper for an empty key. + client.apiKey = tc.apiKey + + var resp map[string]any + err := client.Run(context.Background(), NewRequest(`query { ok }`), &resp) + if err == nil { + t.Fatal("expected a 401 to be an error") + } + if !strings.Contains(err.Error(), tc.hint) { + t.Fatalf("expected error to mention %q, got %q", tc.hint, err.Error()) + } + }) + } +} + +func TestVerboseLogOmitsApiKey(t *testing.T) { + t.Chdir(t.TempDir()) + viper.Set("verbose", true) + defer viper.Set("verbose", false) + + apiSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data": {"ok": true}}`)) + })) + defer apiSrv.Close() + + client := NewClient(WithEndpoint(apiSrv.URL), WithApiKey("sk-nb-secret-key")) + var resp map[string]any + if err := client.Run(context.Background(), NewRequest(`query { ok }`), &resp); err != nil { + t.Fatalf("run failed: %v", err) + } + + logged, err := os.ReadFile("nbctl_graphql.log") + if err != nil { + t.Fatalf("expected a verbose log file: %v", err) + } + if !strings.Contains(string(logged), "POST /api/graphql") { + t.Fatalf("expected the request to be logged, got %q", logged) + } + if strings.Contains(string(logged), "sk-nb-secret-key") { + t.Fatal("the verbose log must not contain the API key") + } +} + func TestNewClient(t *testing.T) { t.Run("with options", func(t *testing.T) { client := NewClient( diff --git a/pkg/nubi/nubi_test.go b/pkg/nubi/nubi_test.go index 2374635..65ce5d3 100644 --- a/pkg/nubi/nubi_test.go +++ b/pkg/nubi/nubi_test.go @@ -14,18 +14,7 @@ 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) + srv := httptest.NewServer(handler) c := client.NewClient(client.WithEndpoint(srv.URL)) 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..5320625 100644 --- a/pkg/testutil/helpers.go +++ b/pkg/testutil/helpers.go @@ -110,18 +110,14 @@ 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 -// the GraphQL response payload under the "data" field. It's suitable for the +// 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) { _ = os.Setenv("NBCTL_TESTING", "true") 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})