diff --git a/cmd/nylas/main.go b/cmd/nylas/main.go index de16fd9..3e5b4dd 100644 --- a/cmd/nylas/main.go +++ b/cmd/nylas/main.go @@ -20,6 +20,7 @@ import ( "github.com/nylas/cli/internal/cli/email" "github.com/nylas/cli/internal/cli/mcp" "github.com/nylas/cli/internal/cli/notetaker" + oauthcmd "github.com/nylas/cli/internal/cli/oauth" "github.com/nylas/cli/internal/cli/otp" "github.com/nylas/cli/internal/cli/rpc" "github.com/nylas/cli/internal/cli/scheduler" @@ -48,6 +49,7 @@ func main() { rootCmd.AddCommand(calendar.NewCalendarCmd()) rootCmd.AddCommand(contacts.NewContactsCmd()) rootCmd.AddCommand(dashboard.NewDashboardCmd()) + rootCmd.AddCommand(oauthcmd.NewOAuthCmd()) rootCmd.AddCommand(setup.NewSetupCmd()) rootCmd.AddCommand(scheduler.NewSchedulerCmd()) rootCmd.AddCommand(admin.NewAdminCmd()) diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 84e4310..76e850f 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -23,6 +23,7 @@ internal/ mcp/ # MCP proxy server utilities/ # Timezone, scheduling, contacts services oauth/ # OAuth callback server + filelock/ # Cross-process file lock (OAuth refresh) browser/ # Browser automation tunnel/ # Cloudflare tunnel webhookserver/ # Webhook server @@ -186,6 +187,7 @@ url := qb.BuildURL(baseURL) | `mcp/` | MCP proxy server for AI assistants | | `config/` | Configuration validation | | `oauth/` | OAuth callback server | + | `filelock/` | Cross-process advisory file lock (flock / LockFileEx) serialising OAuth token refresh | | `utilities/` | Services (contacts, email, scheduling, timezone, webhook) | | `browser/` | Browser automation | | `tunnel/` | Cloudflare tunnel | diff --git a/docs/COMMANDS.md b/docs/COMMANDS.md index 4874536..85acaec 100644 --- a/docs/COMMANDS.md +++ b/docs/COMMANDS.md @@ -113,6 +113,112 @@ nylas auth migrate # Migrate from v2 to v3 --- +## OAuth (Authorization Server) + +Log in to the Nylas OAuth 2.1 / OIDC authorization server. This authenticates +**you**, the person running the CLI, and is distinct from `nylas auth` (which +connects an end user's mailbox as a provider grant). + +It also signs in the `nylas dashboard` commands, so `nylas dashboard login` is +not needed after it. That dashboard session is for the organization you chose +on the consent screen, lasts as long as the access token, and is renewed +automatically from the OAuth session. `nylas oauth logout` ends it too. +To change organization, run `nylas dashboard orgs switch`: it opens the +browser to sign in again, and you choose the organization there. + +Without an account you can sign up on the page that opens. `--region` (or the +configured region) decides where the new organization is created; without +either it is created in the US. + +```bash +nylas oauth login # Log in via the browser (authorization code + PKCE) +nylas oauth login --scope openid,email +nylas oauth login --for mcp # Scopes + resource for `nylas mcp serve --auth oauth` +nylas oauth login --region eu # Sign up with a new organization in the EU +nylas oauth status # Show the stored session and decoded token claims +nylas oauth status --verify # Also confirm the token against /oauth/userinfo +nylas oauth token # Print a valid access token, refreshing if needed +nylas oauth logout # Revoke the session and clear stored tokens +``` + +The CLI is a static public client (client id +`b3a94d82-fc7d-4a22-803e-e603ae0f735c`, no client secret — PKCE protects the +exchange). The browser redirects to `http://127.0.0.1:/callback`, the +address the callback server binds. Tokens are stored in the system keyring; +a value too large for one keychain item (Windows allows 2560 bytes) is split +across several. The ID token is not stored: the CLI neither verifies nor uses +it. + +Every CLI process on the machine shares the one stored session. Refreshing is +serialised by a lock file, `oauth-session.lock`: in `~/.config/nylas` of your +account when the session is in the system keyring (whatever `XDG_CONFIG_HOME` +says, since the keyring does not follow it either), and beside the encrypted +secrets file when the file store is used: +the server rotates the refresh token on every use and revokes the whole family +if a consumed one is replayed, so two `nylas mcp serve` processes refreshing at +once would otherwise sign you out. A process that waited on the lock uses the +tokens the other one stored instead of refreshing again. A refresh that has +started runs to completion even if the command is interrupted, because the +server has already rotated the token it was sent. + +`nylas oauth status` decodes the access token and shows its audience, grants, +scopes and expiry. The claims are **decoded, not verified** — the CLI does not +check the signature; only the resource server's answer is authoritative. + +Default scopes are `openid`, `email`, `offline_access` and `dashboard.session`. +`offline_access` is what makes the server issue a refresh token; without it the +session ends when the access token expires (15 minutes by default). +`dashboard.session` is what lets the CLI sign the `nylas dashboard` commands in: +the consent screen shows it as *Use the Nylas Dashboard as you, with your full +role in this organization*. The server accepts it only from the CLI's built-in +client and does not list it in its discovery document, so the CLI always sends +it rather than dropping it as not offered. A login made with `--scope` and +without `dashboard.session` works, but leaves the dashboard commands signed out. + +After login, the CLI also signs the `nylas dashboard` commands in by exchanging +the access token for a dashboard session. The exchange requires the token to +carry `dashboard.session`; for a login that was not granted it, the CLI says so +and asks you to run `nylas oauth login` again. `nylas dashboard orgs switch` +adds `dashboard.session` when it signs in again. If a session from +`nylas dashboard login` is already stored for the configured server, it is kept +(its organization and app selection are unchanged); run `nylas dashboard logout` +first to use the OAuth login for the dashboard commands instead. + +Use the access token with any OAuth-protected endpoint: + +```bash +curl -H "Authorization: Bearer $(nylas oauth token)" https://example/resource +``` + +### Pointing at a local authorization server + +The authorization server is hosted by `dashboard-account`, so it uses the same +base URL as the `nylas dashboard` commands: + +```bash +NYLAS_DASHBOARD_ACCOUNT_URL=http://localhost:3001 nylas oauth login +``` + +If that server registers the CLI under a different client id, override it with +`NYLAS_OAUTH_CLIENT_ID` (it must still allow the `http://127.0.0.1/callback` +redirect URI). + +The dashboard session exchange only accepts tokens from dashboard-account's +built-in first-party clients (the Nylas CLI and Nylas Mail), and only that +client may request `dashboard.session`. With any other client id the CLI leaves +`dashboard.session` out of the request (the server would refuse the whole +request otherwise), so `nylas oauth login` still succeeds, but the `nylas +dashboard` commands stay signed out; use `nylas dashboard login` for those. + +The CLI resolves every endpoint from the server's +`/.well-known/oauth-authorization-server` document, and that document is built +from the server's `OAUTH_ISSUER`. If `OAUTH_ISSUER` names a host the CLI cannot +reach (for example a Cloudflare tunnel that is no longer running), login fails +even though the local port responds — set `OAUTH_ISSUER` to the address you +actually browse to. + +--- + ## Dashboard Manage your Nylas Dashboard account, applications, domains, and API keys directly from the CLI. @@ -134,6 +240,14 @@ nylas dashboard status # Show current auth status nylas dashboard refresh # Refresh session tokens ``` +A dashboard session is tied to the servers it was issued for: the account URL +and both gateway URLs (`NYLAS_DASHBOARD_ACCOUNT_URL`, +`NYLAS_DASHBOARD_GATEWAY_URL`, `NYLAS_DASHBOARD_GATEWAY_US_URL`, +`NYLAS_DASHBOARD_GATEWAY_EU_URL`, or the config file). If any of them changes, +the stored session is refused rather than sent to a server that did not issue +it; log in again, or restore the settings. `nylas dashboard logout` then clears +it locally without contacting the new server. + ### SSO (Direct) ```bash @@ -584,6 +698,7 @@ nylas mcp install --all # Install for all detected assistants nylas mcp status # Check installation status nylas mcp uninstall --assistant cursor # Remove configuration nylas mcp serve # Start MCP server (used by assistants) +nylas mcp serve --auth oauth # ...authenticating with `nylas oauth login --for mcp` ``` **Supported assistants:** diff --git a/docs/DEVELOPMENT.md b/docs/DEVELOPMENT.md index 63a0c88..422115d 100644 --- a/docs/DEVELOPMENT.md +++ b/docs/DEVELOPMENT.md @@ -68,6 +68,35 @@ make test-integration **CRITICAL:** Integration tests create real resources. Always use `make ci-full` for automatic cleanup. +### OAuth authorization server tests + +`internal/cli/integration/oauth_test.go` drives a real dashboard-account +authorization server instead of the Nylas API, so it needs its own variable and +skips without it: + +```bash +NYLAS_OAUTH_AS_URL=http://localhost:3001 \ + go test -tags integration -run TestOAuthAS ./internal/cli/integration/ +``` + +Requirements on the server side: + +- dashboard-account running (in a Tilt stack it is on port 3001) +- `/dev` routes enabled — `ENABLE_DEV_ROUTES=true` or `IS_E2E=true`. The tests + seed their own user, consent grant and authorization code through them, which + is what lets the token exchange run without a browser. +- the CLI's static public client registered, with the redirect URI + `http://127.0.0.1/callback`. Set `NYLAS_OAUTH_CLIENT_ID` if the local server + registers it under another id. + +The tests front the server with a small proxy that rewrites the issuer origin in +the discovery document. dashboard-account builds every advertised endpoint from +`OAUTH_ISSUER`, and in a local stack that is frequently a tunnel hostname that is +stale or unreachable; the client under test is spec-correct and follows whatever +the document says. If you would rather fix it at the source, set +`OAUTH_ISSUER=http://localhost:3001` in `infra/.env.local` and restart the +service — the proxy then rewrites nothing. + --- ## Project Structure diff --git a/docs/commands/mcp.md b/docs/commands/mcp.md index 1708da9..61fc6c3 100644 --- a/docs/commands/mcp.md +++ b/docs/commands/mcp.md @@ -43,8 +43,19 @@ nylas mcp install --assistant claude-code # Specific assistant nylas mcp install --assistant cursor # Cursor IDE nylas mcp install --all # All detected assistants nylas mcp install --binary /path/to/nylas # Custom binary path +nylas mcp install --assistant claude-code --auth oauth # Proxy authenticates with OAuth ``` +Every assistant is configured to launch `nylas mcp serve` over STDIO, and no +credential is written into any assistant config. `--auth oauth` adds +`--auth oauth` to the launcher (run `nylas oauth login --for mcp` first). + +Pointing an assistant directly at the hosted server +(`https://mcp.{us,eu}.nylas.com`) and letting it run OAuth itself is not +configured by `install` yet: which supported assistants handle remote MCP with +OAuth, and in which config format, has not been verified. The local proxy is the +compatibility path for all of them. + ### Status Check installation status: @@ -67,9 +78,50 @@ nylas mcp uninstall --all Start the MCP server (called by AI assistants, not directly): ```bash -nylas mcp serve +nylas mcp serve # authenticate with the API key (default) +nylas mcp serve --auth oauth # authenticate with an OAuth session +``` + +#### OAuth (`--auth oauth`) + +Log in once for the MCP server, then point the assistant at +`nylas mcp serve --auth oauth`: + +```bash +nylas oauth login --for mcp ``` +`--for mcp` requests the data scopes the MCP tools use (`email.read`, +`email.send`, `calendar.read`, `calendar.write`, `contacts.read`, +`notetaker.read`, `grants.read`) plus `offline_access` and `dashboard.session` +(which signs the `nylas dashboard` commands in too; the MCP server ignores it), +and sends the MCP server +of your configured region as the RFC 8707 `resource`, so the token is issued +for that server only. Scopes the authorization server does not offer are left +out and listed. + +With `--auth oauth` the proxy: + +- asks for a valid token before **every** request and refreshes it as it nears + expiry (access tokens last 15 minutes). Refreshing is serialised across every + `nylas mcp serve` on the machine, so several assistants can share one login. +- sends requests to the MCP server named in the token's audience (`aud`), not + the configured region, and refuses a token whose audience names neither. +- offers the default grant (`X-Nylas-Grant-Id` and the injected `grant_id`) + only when the token's `grants` claim lists it, and does not answer + `get_grant` from the local grant store. +- on `401` with a `WWW-Authenticate` challenge, refreshes once and retries; if + that fails it tells you to run `nylas oauth login --for mcp`. +- on `403 insufficient_scope`, names the missing scope and the login command. + +#### Protocol + +The hosted server is stateless. The proxy sends no `Mcp-Session-Id` and ignores +one if offered. It sends `Mcp-Method` on every request, `Mcp-Name` for +`tools/call` and `prompts/get` (the server refuses a name that disagrees with +the body), and, once `initialize` has answered, `Mcp-Protocol-Version` with the +version the server negotiated. + --- ## Supported Assistants @@ -158,6 +210,8 @@ region: eu # or "us" (default) ``` The MCP proxy reads this setting and routes requests to the appropriate regional endpoint. +With `--auth oauth` the region is used once, at `nylas oauth login --for mcp`, +to choose the token's resource; requests then follow the token's audience. --- diff --git a/docs/security/overview.md b/docs/security/overview.md index 14f4939..3c6f47c 100644 --- a/docs/security/overview.md +++ b/docs/security/overview.md @@ -56,6 +56,34 @@ Non-sensitive settings stored in `~/.config/nylas/config.yaml`: - Callback port - Local default grant mirror +### OAuth Login Callback + +`nylas oauth login` and `nylas auth login` receive the redirect on a loopback +callback server (`127.0.0.1`, plus `::1` for `localhost`): + +- The expected `state` is set before the browser opens, and is compared in + constant time. +- Only a request carrying that state can end the login, with a code or an + `error`. Any other request to the port is refused with 400 and the login + keeps waiting, so a stray or hostile request cannot abort it. +- The `error` value is shown only if it matches `^[a-z_]{1,64}$`, so control + characters from the URL never reach the terminal. + +### Session Locks + +Session writes are serialised across CLI processes by advisory file locks, +always taken in this order: + +| Lock | Guards | +|------|--------| +| `dashboard-session.lock` | Dashboard session keys (renewal from OAuth, clear, reset) | +| `oauth-session.lock` | OAuth session keys (login, refresh, logout, reset) | +| `.secrets.lock` | Each read and write of the encrypted file store | + +With the encrypted file store the first two sit in the config directory; with +the system keyring they sit in `~/.config/nylas/` under the account's home +directory, whatever `XDG_CONFIG_HOME` is. + --- ## Testing diff --git a/go.mod b/go.mod index 959233e..205446c 100644 --- a/go.mod +++ b/go.mod @@ -21,6 +21,7 @@ require ( github.com/zalando/go-keyring v0.2.6 golang.org/x/crypto v0.46.0 golang.org/x/mod v0.30.0 + golang.org/x/sys v0.43.0 golang.org/x/term v0.38.0 golang.org/x/text v0.32.0 golang.org/x/time v0.14.0 @@ -60,6 +61,5 @@ require ( github.com/tetratelabs/wazero v1.11.0 // indirect github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect golang.org/x/sync v0.19.0 // indirect - golang.org/x/sys v0.43.0 // indirect lukechampine.com/adiantum v1.1.1 // indirect ) diff --git a/internal/adapters/dashboard/account_client.go b/internal/adapters/dashboard/account_client.go index 6001895..5c1251b 100644 --- a/internal/adapters/dashboard/account_client.go +++ b/internal/adapters/dashboard/account_client.go @@ -162,6 +162,19 @@ func (c *AccountClient) SSOStart(ctx context.Context, loginType, mode string, pr return &result, nil } +// ExchangeOAuthToken trades an OAuth access token for a DPoP-bound dashboard +// session. The token goes in the body: this server reads Authorization as a +// dashboard session token. +func (c *AccountClient) ExchangeOAuthToken(ctx context.Context, accessToken string) (*domain.DashboardOAuthExchangeResponse, error) { + body := map[string]any{"accessToken": accessToken} + + var result domain.DashboardOAuthExchangeResponse + if err := c.doPost(ctx, "/auth/cli/oauth/exchange", body, nil, "", &result); err != nil { + return nil, fmt.Errorf("failed to exchange the OAuth session for a dashboard session: %w", err) + } + return &result, nil +} + // SSOPoll polls the SSO device flow for completion. func (c *AccountClient) SSOPoll(ctx context.Context, flowID, orgPublicID string) (*domain.DashboardSSOPollResponse, error) { body := map[string]any{ diff --git a/internal/adapters/dashboard/account_client_test.go b/internal/adapters/dashboard/account_client_test.go index c1ecc16..fe62831 100644 --- a/internal/adapters/dashboard/account_client_test.go +++ b/internal/adapters/dashboard/account_client_test.go @@ -513,6 +513,41 @@ func TestAccountClientSSOPollVariants(t *testing.T) { }) } +func TestAccountClientExchangeOAuthTokenSendsTheTokenInTheBody(t *testing.T) { + t.Parallel() + + server := newAccountClientTestServer(t, func(t *testing.T, w http.ResponseWriter, r *http.Request, _ []byte, body map[string]any) { + assert.Equal(t, http.MethodPost, r.Method) + assert.Equal(t, "/auth/cli/oauth/exchange", r.URL.Path) + assert.Equal(t, "eyJ.access.token", body["accessToken"]) + assert.Empty(t, r.Header.Get("Authorization"), "the server reads Authorization as a dashboard token") + assert.Equal(t, "test-proof", r.Header.Get("DPoP")) + + writeDashboardEnvelope(t, w, map[string]any{ + "userToken": "user-token", + "orgToken": "org-token", + "user": map[string]any{"publicId": "usr_1"}, + "organizations": []any{}, + "orgPublicId": "org_1", + "expiresAt": "2026-09-24T12:15:00.000Z", + }) + }) + defer server.Close() + + client := &AccountClient{ + baseURL: server.URL, + httpClient: server.Client(), + dpop: &mockDPoP{proof: "test-proof"}, + } + + resp, err := client.ExchangeOAuthToken(context.Background(), "eyJ.access.token") + require.NoError(t, err) + assert.Equal(t, "user-token", resp.UserToken) + assert.Equal(t, "usr_1", resp.User.PublicID) + assert.Equal(t, "org_1", resp.OrgPublicID) + assert.Equal(t, 2026, resp.ExpiresAt.Year()) +} + func TestAccountClientRefreshPropagatesUnderlyingError(t *testing.T) { t.Parallel() diff --git a/internal/adapters/dashboard/gateway_client.go b/internal/adapters/dashboard/gateway_client.go index 31c114d..4bbb35b 100644 --- a/internal/adapters/dashboard/gateway_client.go +++ b/internal/adapters/dashboard/gateway_client.go @@ -277,6 +277,13 @@ func (c *GatewayClient) doGraphQL(ctx context.Context, url, query string, variab // NYLAS_DASHBOARD_GATEWAY_US_URL → overrides US only // NYLAS_DASHBOARD_GATEWAY_EU_URL → overrides EU only // NYLAS_DASHBOARD_GATEWAY_URL → overrides both (single local gateway) +// +// GatewayURL returns the gateway GraphQL URL for region ("eu", else US), +// honouring the NYLAS_DASHBOARD_GATEWAY_* overrides. +func GatewayURL(region string) string { + return gatewayURL(region) +} + func gatewayURL(region string) string { if region == "eu" { if envURL := os.Getenv("NYLAS_DASHBOARD_GATEWAY_EU_URL"); envURL != "" { diff --git a/internal/adapters/dashboard/mock.go b/internal/adapters/dashboard/mock.go index be5fb07..c7e269d 100644 --- a/internal/adapters/dashboard/mock.go +++ b/internal/adapters/dashboard/mock.go @@ -17,6 +17,7 @@ type MockAccountClient struct { LogoutFn func(ctx context.Context, userToken, orgToken string) error SSOStartFn func(ctx context.Context, loginType, mode string, privacyPolicyAccepted bool, email string) (*domain.DashboardSSOStartResponse, error) SSOPollFn func(ctx context.Context, flowID, orgPublicID string) (*domain.DashboardSSOPollResponse, error) + ExchangeOAuthTokenFn func(ctx context.Context, accessToken string) (*domain.DashboardOAuthExchangeResponse, error) GetCurrentSessionFn func(ctx context.Context, userToken, orgToken string) (*domain.DashboardSessionResponse, error) SwitchOrgFn func(ctx context.Context, orgPublicID, userToken, orgToken string) (*domain.DashboardSwitchOrgResponse, error) ListDomainsFn func(ctx context.Context, limit int, pageToken, userToken, orgToken string) (domain.DashboardInboxDomainPage, error) @@ -56,6 +57,9 @@ func (m *MockAccountClient) SSOStart(ctx context.Context, loginType, mode string func (m *MockAccountClient) SSOPoll(ctx context.Context, flowID, orgPublicID string) (*domain.DashboardSSOPollResponse, error) { return m.SSOPollFn(ctx, flowID, orgPublicID) } +func (m *MockAccountClient) ExchangeOAuthToken(ctx context.Context, accessToken string) (*domain.DashboardOAuthExchangeResponse, error) { + return m.ExchangeOAuthTokenFn(ctx, accessToken) +} func (m *MockAccountClient) GetCurrentSession(ctx context.Context, userToken, orgToken string) (*domain.DashboardSessionResponse, error) { if m.GetCurrentSessionFn != nil { return m.GetCurrentSessionFn(ctx, userToken, orgToken) diff --git a/internal/adapters/dpop/dpop.go b/internal/adapters/dpop/dpop.go index 0e73353..0fb899e 100644 --- a/internal/adapters/dpop/dpop.go +++ b/internal/adapters/dpop/dpop.go @@ -8,6 +8,7 @@ import ( "crypto/sha256" "encoding/base64" "encoding/json" + "errors" "fmt" "net/url" "strings" @@ -30,9 +31,15 @@ type Service struct { func New(secrets ports.SecretStore) (*Service, error) { s := &Service{} - // Try to load existing key + // Try to load existing key. Only a key that is genuinely absent is + // replaced: a failed read (a locked or unreachable keychain) must not + // overwrite the stored key, because the server binds a dashboard session, + // and each CLI access token it exchanges, to the key that proved them. seedB64, err := secrets.Get(ports.KeyDashboardDPoPKey) - if err == nil && seedB64 != "" { + if err != nil && !errors.Is(err, domain.ErrSecretNotFound) { + return nil, fmt.Errorf("%w: failed to read the stored key: %w", domain.ErrDashboardDPoP, err) + } + if seedB64 != "" { seed, decErr := base64.StdEncoding.DecodeString(seedB64) if decErr == nil && len(seed) == ed25519.SeedSize { s.privateKey = ed25519.NewKeyFromSeed(seed) diff --git a/internal/adapters/dpop/dpop_test.go b/internal/adapters/dpop/dpop_test.go index 03b14fc..d5cd682 100644 --- a/internal/adapters/dpop/dpop_test.go +++ b/internal/adapters/dpop/dpop_test.go @@ -5,9 +5,11 @@ import ( "crypto/sha256" "encoding/base64" "encoding/json" + "errors" "strings" "testing" + "github.com/nylas/cli/internal/domain" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -306,3 +308,42 @@ func extractClaim(t *testing.T, proof, key string) string { } return "" } + +// readFailingStore fails every Get, like a locked or unreachable keychain. +type readFailingStore struct{ *mockSecretStore } + +func (readFailingStore) Get(string) (string, error) { return "", errors.New("keychain locked") } + +func TestNew_ReadFailureDoesNotReplaceTheStoredKey(t *testing.T) { + t.Parallel() + store := newMockSecretStore() + _, err := New(store) + require.NoError(t, err) + stored := store.data["dashboard_dpop_key"] + + _, err = New(readFailingStore{store}) + + require.ErrorIs(t, err, domain.ErrDashboardDPoP) + assert.Equal(t, stored, store.data["dashboard_dpop_key"], + "the server binds sessions to this key; a failed read must not overwrite it") +} + +func TestNew_MissingKeyIsGenerated(t *testing.T) { + t.Parallel() + store := notFoundStore{newMockSecretStore()} + + _, err := New(store) + + require.NoError(t, err) + assert.NotEmpty(t, store.data["dashboard_dpop_key"]) +} + +// notFoundStore reports a missing key the way the real keyring does. +type notFoundStore struct{ *mockSecretStore } + +func (s notFoundStore) Get(key string) (string, error) { + if v, ok := s.data[key]; ok { + return v, nil + } + return "", domain.ErrSecretNotFound +} diff --git a/internal/adapters/filelock/filelock.go b/internal/adapters/filelock/filelock.go new file mode 100644 index 0000000..4aa524f --- /dev/null +++ b/internal/adapters/filelock/filelock.go @@ -0,0 +1,77 @@ +// Package filelock implements ports.CrossProcessLock with an advisory lock on +// a file: flock(2) on Unix, LockFileEx on Windows. +// +// Both are released by the kernel when the holding process exits, including +// when it crashes, so a killed `nylas mcp serve` cannot wedge the others. +package filelock + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "time" +) + +// defaultPollInterval is how often a waiter retries. Acquisition is +// non-blocking plus polling, rather than a blocking syscall, because a +// blocking flock cannot be cancelled by a context. +const defaultPollInterval = 20 * time.Millisecond + +// Lock is a cross-process lock backed by the file at path. +type Lock struct { + path string + poll time.Duration +} + +// New returns a lock on path. The file and its directory are created on +// first use; the file's content is never read or written. +func New(path string) *Lock { + return &Lock{path: path, poll: defaultPollInterval} +} + +// Lock acquires the lock, waiting until ctx is done. +func (l *Lock) Lock(ctx context.Context) (func() error, error) { + if err := os.MkdirAll(filepath.Dir(l.path), 0o700); err != nil { + return nil, fmt.Errorf("failed to create lock directory: %w", err) + } + // #nosec G304 -- the path is built by the caller from the CLI's own config dir. + file, err := os.OpenFile(l.path, os.O_RDWR|os.O_CREATE, 0o600) + if err != nil { + return nil, fmt.Errorf("failed to open lock file: %w", err) + } + + for { + acquired, err := tryLock(file) + if err != nil { + _ = file.Close() + return nil, fmt.Errorf("failed to lock %s: %w", l.path, err) + } + if acquired { + return unlocker(file), nil + } + + timer := time.NewTimer(l.poll) + select { + case <-ctx.Done(): + timer.Stop() + _ = file.Close() + return nil, fmt.Errorf("waiting for lock %s: %w", l.path, ctx.Err()) + case <-timer.C: + } + } +} + +func unlocker(file *os.File) func() error { + released := false + return func() error { + if released { + return errors.New("lock already released") + } + released = true + unlockErr := unlockFile(file) + closeErr := file.Close() + return errors.Join(unlockErr, closeErr) + } +} diff --git a/internal/adapters/filelock/filelock_other.go b/internal/adapters/filelock/filelock_other.go new file mode 100644 index 0000000..d6e7f61 --- /dev/null +++ b/internal/adapters/filelock/filelock_other.go @@ -0,0 +1,16 @@ +//go:build !unix && !windows + +package filelock + +import ( + "errors" + "os" +) + +// errUnsupported fails closed: without a cross-process lock two refreshers +// could replay a rotated token and sign the user out. +var errUnsupported = errors.New("cross-process file locking is not supported on this platform") + +func tryLock(*os.File) (bool, error) { return false, errUnsupported } + +func unlockFile(*os.File) error { return errUnsupported } diff --git a/internal/adapters/filelock/filelock_test.go b/internal/adapters/filelock/filelock_test.go new file mode 100644 index 0000000..0144ccf --- /dev/null +++ b/internal/adapters/filelock/filelock_test.go @@ -0,0 +1,173 @@ +//go:build !integration + +package filelock + +import ( + "bufio" + "context" + "errors" + "os" + "os/exec" + "path/filepath" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const helperEnv = "NYLAS_FILELOCK_HELPER_PATH" + +// TestMain lets the test binary double as a separate process that holds the +// lock, because the property that matters is exclusion between processes. +func TestMain(m *testing.M) { + if path := os.Getenv(helperEnv); path != "" { + holdLockUntilStdinCloses(path) + return + } + os.Exit(m.Run()) +} + +func holdLockUntilStdinCloses(path string) { + unlock, err := New(path).Lock(context.Background()) + if err != nil { + _, _ = os.Stdout.WriteString("error: " + err.Error() + "\n") + os.Exit(2) + } + _, _ = os.Stdout.WriteString("locked\n") + _, _ = bufio.NewReader(os.Stdin).ReadString('\n') + _ = unlock() + os.Exit(0) +} + +// startHolder runs a child process that holds the lock until its stdin closes. +func startHolder(t *testing.T, path string) (*exec.Cmd, func()) { + t.Helper() + + // #nosec G204 -- re-executes this test binary. + cmd := exec.Command(os.Args[0], "-test.run=^$") + cmd.Env = append(os.Environ(), helperEnv+"="+path) + stdin, err := cmd.StdinPipe() + require.NoError(t, err) + stdout, err := cmd.StdoutPipe() + require.NoError(t, err) + require.NoError(t, cmd.Start()) + + line, err := bufio.NewReader(stdout).ReadString('\n') + require.NoError(t, err) + require.Equal(t, "locked\n", line) + + release := func() { + _ = stdin.Close() + _ = cmd.Wait() + } + t.Cleanup(func() { + _ = stdin.Close() + if cmd.ProcessState == nil { + _ = cmd.Process.Kill() + _ = cmd.Wait() + } + }) + return cmd, release +} + +func tryLockFor(path string, d time.Duration) (func() error, error) { + ctx, cancel := context.WithTimeout(context.Background(), d) + defer cancel() + return New(path).Lock(ctx) +} + +func TestLock_ExcludesAnotherProcess(t *testing.T) { + path := filepath.Join(t.TempDir(), "refresh.lock") + _, release := startHolder(t, path) + + _, err := tryLockFor(path, 150*time.Millisecond) + require.ErrorIs(t, err, context.DeadlineExceeded, "the lock is held by the other process") + + release() + + unlock, err := tryLockFor(path, 2*time.Second) + require.NoError(t, err, "released by the other process") + require.NoError(t, unlock()) +} + +func TestLock_IsReleasedWhenTheHolderIsKilled(t *testing.T) { + // A crashed `nylas mcp serve` must not sign every other process out by + // wedging the refresh lock forever. + path := filepath.Join(t.TempDir(), "refresh.lock") + cmd, _ := startHolder(t, path) + + require.NoError(t, cmd.Process.Kill()) + _ = cmd.Wait() + + unlock, err := tryLockFor(path, 2*time.Second) + require.NoError(t, err) + require.NoError(t, unlock()) +} + +func TestLock_ExcludesWithinOneProcess(t *testing.T) { + path := filepath.Join(t.TempDir(), "refresh.lock") + + var inside, maxInside atomic.Int32 + var wg sync.WaitGroup + for range 8 { + wg.Add(1) + go func() { + defer wg.Done() + unlock, err := New(path).Lock(context.Background()) + if !assert.NoError(t, err) { + return + } + now := inside.Add(1) + for { + prev := maxInside.Load() + if now <= prev || maxInside.CompareAndSwap(prev, now) { + break + } + } + time.Sleep(5 * time.Millisecond) + inside.Add(-1) + assert.NoError(t, unlock()) + }() + } + wg.Wait() + + assert.Equal(t, int32(1), maxInside.Load(), "never more than one holder at a time") +} + +func TestLock_CreatesDirectoryAndPrivateFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "nested", "dir", "refresh.lock") + + unlock, err := New(path).Lock(context.Background()) + require.NoError(t, err) + t.Cleanup(func() { _ = unlock() }) + + info, err := os.Stat(path) + require.NoError(t, err) + if filepath.Separator == '/' { + assert.Equal(t, os.FileMode(0o600), info.Mode().Perm()) + } +} + +func TestLock_UnlockTwiceIsAnError(t *testing.T) { + unlock, err := New(filepath.Join(t.TempDir(), "refresh.lock")).Lock(context.Background()) + require.NoError(t, err) + + require.NoError(t, unlock()) + assert.Error(t, unlock()) +} + +func TestLock_HonoursAlreadyCancelledContext(t *testing.T) { + path := filepath.Join(t.TempDir(), "refresh.lock") + holder, err := New(path).Lock(context.Background()) + require.NoError(t, err) + t.Cleanup(func() { _ = holder() }) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err = New(path).Lock(ctx) + + require.True(t, errors.Is(err, context.Canceled)) +} diff --git a/internal/adapters/filelock/filelock_unix.go b/internal/adapters/filelock/filelock_unix.go new file mode 100644 index 0000000..9e2bdb0 --- /dev/null +++ b/internal/adapters/filelock/filelock_unix.go @@ -0,0 +1,33 @@ +//go:build unix + +package filelock + +import ( + "errors" + "os" + + "golang.org/x/sys/unix" +) + +// tryLock takes an exclusive flock without blocking. flock locks belong to +// the open file description, so two opens in one process exclude each other +// just as two processes do. +func tryLock(file *os.File) (bool, error) { + for { + err := unix.Flock(int(file.Fd()), unix.LOCK_EX|unix.LOCK_NB) + switch { + case err == nil: + return true, nil + case errors.Is(err, unix.EINTR): + continue + case errors.Is(err, unix.EWOULDBLOCK): + return false, nil + default: + return false, err + } + } +} + +func unlockFile(file *os.File) error { + return unix.Flock(int(file.Fd()), unix.LOCK_UN) +} diff --git a/internal/adapters/filelock/filelock_windows.go b/internal/adapters/filelock/filelock_windows.go new file mode 100644 index 0000000..097773c --- /dev/null +++ b/internal/adapters/filelock/filelock_windows.go @@ -0,0 +1,35 @@ +//go:build windows + +package filelock + +import ( + "errors" + "os" + + "golang.org/x/sys/windows" +) + +// tryLock takes an exclusive LockFileEx lock on the first byte without +// blocking. Windows locks belong to the handle, so two opens in one process +// exclude each other just as two processes do. +func tryLock(file *os.File) (bool, error) { + overlapped := new(windows.Overlapped) + err := windows.LockFileEx( + windows.Handle(file.Fd()), + windows.LOCKFILE_EXCLUSIVE_LOCK|windows.LOCKFILE_FAIL_IMMEDIATELY, + 0, 1, 0, overlapped, + ) + switch { + case err == nil: + return true, nil + case errors.Is(err, windows.ERROR_LOCK_VIOLATION): + return false, nil + default: + return false, err + } +} + +func unlockFile(file *os.File) error { + overlapped := new(windows.Overlapped) + return windows.UnlockFileEx(windows.Handle(file.Fd()), 0, 1, 0, overlapped) +} diff --git a/internal/adapters/keyring/file.go b/internal/adapters/keyring/file.go index 8943ac8..35a46a7 100644 --- a/internal/adapters/keyring/file.go +++ b/internal/adapters/keyring/file.go @@ -1,6 +1,7 @@ package keyring import ( + "context" "crypto/rand" "encoding/base64" "encoding/json" @@ -10,7 +11,9 @@ import ( "path/filepath" "strings" "sync" + "time" + "github.com/nylas/cli/internal/adapters/filelock" "github.com/nylas/cli/internal/domain" "golang.org/x/crypto/argon2" ) @@ -18,6 +21,9 @@ import ( const ( fileStorePassphraseEnv = "NYLAS_FILE_STORE_PASSPHRASE" fileStoreSaltSize = 16 + // fileStoreLockTimeout bounds how long one operation waits for another + // process's. Each holds the lock for a single read-modify-write. + fileStoreLockTimeout = 10 * time.Second ) // EncryptedFileStore implements SecretStore using an encrypted file. @@ -39,7 +45,13 @@ type EncryptedFileStore struct { passphrase []byte migrationKey []byte legacyKey []byte - mu sync.RWMutex + // mu serialises goroutines; fileLock serialises processes. Every call is + // a read-modify-write of the whole file, so without fileLock two CLI + // processes interleave and the last rename silently discards the other's + // write — for an OAuth session, a just-rotated refresh token, whose loss + // makes the next refresh a replay that revokes the whole token family. + mu sync.RWMutex + fileLock *filelock.Lock } // NewEncryptedFileStore creates a new EncryptedFileStore rooted in configDir. @@ -89,13 +101,33 @@ func NewEncryptedFileStore(configDir string) (*EncryptedFileStore, error) { passphrase: passphrase, migrationKey: migrationKey, legacyKey: legacyKey, + fileLock: filelock.New(filepath.Join(configDir, ".secrets.lock")), + }, nil +} + +// lock takes both locks. The returned function releases them. +func (f *EncryptedFileStore) lock() (func(), error) { + f.mu.Lock() + ctx, cancel := context.WithTimeout(context.Background(), fileStoreLockTimeout) + defer cancel() + unlock, err := f.fileLock.Lock(ctx) + if err != nil { + f.mu.Unlock() + return nil, fmt.Errorf("%w: %v", domain.ErrSecretStoreFailed, err) + } + return func() { + _ = unlock() + f.mu.Unlock() }, nil } // Set stores a secret value for the given key. func (f *EncryptedFileStore) Set(key, value string) error { - f.mu.Lock() - defer f.mu.Unlock() + release, err := f.lock() + if err != nil { + return err + } + defer release() secrets, err := f.loadSecrets() if err != nil && !os.IsNotExist(err) { @@ -121,8 +153,11 @@ func (f *EncryptedFileStore) Set(key, value string) error { // CLI workloads aren't read-heavy, so serializing reads is the right // trade for guaranteed migration correctness. func (f *EncryptedFileStore) Get(key string) (string, error) { - f.mu.Lock() - defer f.mu.Unlock() + release, err := f.lock() + if err != nil { + return "", err + } + defer release() secrets, err := f.loadSecrets() if err != nil { @@ -141,8 +176,11 @@ func (f *EncryptedFileStore) Get(key string) (string, error) { // Delete removes a secret for the given key. func (f *EncryptedFileStore) Delete(key string) error { - f.mu.Lock() - defer f.mu.Unlock() + release, err := f.lock() + if err != nil { + return err + } + defer release() secrets, err := f.loadSecrets() if err != nil { diff --git a/internal/adapters/keyring/file_crossprocess_test.go b/internal/adapters/keyring/file_crossprocess_test.go new file mode 100644 index 0000000..ac624f7 --- /dev/null +++ b/internal/adapters/keyring/file_crossprocess_test.go @@ -0,0 +1,54 @@ +package keyring + +import ( + "fmt" + "sync" + "testing" +) + +// TestEncryptedFileStore_WritesFromSeparateStoresAllSurvive stands in for two +// CLI processes: two stores on one directory share no mutex, so only the +// file lock stops a read-modify-write in one from discarding the other's — +// the lost write that, for an OAuth session, puts back a rotated refresh +// token and revokes the whole family on the next refresh. +func TestEncryptedFileStore_WritesFromSeparateStoresAllSurvive(t *testing.T) { + dir := t.TempDir() + setFileStorePassphrase(t) + + stores := make([]*EncryptedFileStore, 2) + for i := range stores { + store, err := NewEncryptedFileStore(dir) + if err != nil { + t.Fatalf("NewEncryptedFileStore: %v", err) + } + stores[i] = store + } + + const perStore = 8 + var wg sync.WaitGroup + for i, store := range stores { + wg.Add(1) + go func(i int, store *EncryptedFileStore) { + defer wg.Done() + for n := range perStore { + if err := store.Set(fmt.Sprintf("key-%d-%d", i, n), "value"); err != nil { + t.Errorf("Set: %v", err) + } + } + }(i, store) + } + wg.Wait() + + reader, err := NewEncryptedFileStore(dir) + if err != nil { + t.Fatalf("NewEncryptedFileStore: %v", err) + } + for i := range stores { + for n := range perStore { + key := fmt.Sprintf("key-%d-%d", i, n) + if _, err := reader.Get(key); err != nil { + t.Errorf("%s was lost: %v", key, err) + } + } + } +} diff --git a/internal/adapters/keyring/keyring.go b/internal/adapters/keyring/keyring.go index a3541b5..3f3c5e6 100644 --- a/internal/adapters/keyring/keyring.go +++ b/internal/adapters/keyring/keyring.go @@ -2,9 +2,13 @@ package keyring import ( + "crypto/rand" + "encoding/hex" "errors" "fmt" "os" + "strconv" + "strings" "github.com/nylas/cli/internal/domain" "github.com/nylas/cli/internal/ports" @@ -21,9 +25,52 @@ func NewSystemKeyring() *SystemKeyring { return &SystemKeyring{} } +// Keychains cap the size of one item: Windows Credential Manager at 2560 +// bytes, and the macOS `security` command at 4096 for the whole command line. +// A token can outgrow that, so a value longer than keychainChunkSize is split +// across several items and the key holds a header naming them. +const ( + keychainChunkSize = 2000 + keychainChunkPrefix = "nylas:chunked:v1:" + keychainMaxChunks = 64 +) + // Set stores a secret value for the given key. +// +// A long value's chunks are written under a fresh generation id before the +// header that points at them, so a reader never pairs a header with chunks +// from another write. The previous generation is removed afterwards. func (k *SystemKeyring) Set(key, value string) error { - return keyring.Set(serviceName, key, value) + previous, _ := keyring.Get(serviceName, key) + + if len(value) <= keychainChunkSize && !strings.HasPrefix(value, keychainChunkPrefix) { + if err := keyring.Set(serviceName, key, value); err != nil { + return err + } + deleteChunks(key, previous) + return nil + } + + parts := splitChunks(value) + if len(parts) > keychainMaxChunks { + return fmt.Errorf("secret %s is too large for the system keyring (%d bytes)", key, len(value)) + } + gen, err := newChunkGeneration() + if err != nil { + return err + } + for i, part := range parts { + if err := keyring.Set(serviceName, chunkKey(key, gen, i), part); err != nil { + deleteChunks(key, chunkHeader(gen, i)) + return err + } + } + if err := keyring.Set(serviceName, key, chunkHeader(gen, len(parts))); err != nil { + deleteChunks(key, chunkHeader(gen, len(parts))) + return err + } + deleteChunks(key, previous) + return nil } // Get retrieves a secret value for the given key. @@ -32,16 +79,91 @@ func (k *SystemKeyring) Get(key string) (string, error) { if errors.Is(err, keyring.ErrNotFound) { return "", domain.ErrSecretNotFound } - return value, err + if err != nil { + return "", err + } + gen, n, ok := parseChunkHeader(value) + if !ok { + return value, nil + } + var b strings.Builder + for i := range n { + part, err := keyring.Get(serviceName, chunkKey(key, gen, i)) + if err != nil { + // Another process replaced the value between the two reads, or a + // chunk was removed by hand. Either way this copy is unusable. + return "", fmt.Errorf("secret %s is incomplete in the system keyring: %w", key, err) + } + b.WriteString(part) + } + return b.String(), nil } // Delete removes a secret for the given key. func (k *SystemKeyring) Delete(key string) error { + previous, _ := keyring.Get(serviceName, key) err := keyring.Delete(serviceName, key) - if err == keyring.ErrNotFound { - return nil // Already deleted + if err != nil && !errors.Is(err, keyring.ErrNotFound) { + return err + } + deleteChunks(key, previous) + return nil +} + +func splitChunks(value string) []string { + var parts []string + for len(value) > keychainChunkSize { + parts = append(parts, value[:keychainChunkSize]) + value = value[keychainChunkSize:] + } + return append(parts, value) +} + +func newChunkGeneration() (string, error) { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("failed to generate a keyring chunk id: %w", err) + } + return hex.EncodeToString(b), nil +} + +func chunkKey(key, gen string, i int) string { + return fmt.Sprintf("%s.chunk.%s.%d", key, gen, i) +} + +func chunkHeader(gen string, n int) string { + return fmt.Sprintf("%s%s:%d", keychainChunkPrefix, gen, n) +} + +func parseChunkHeader(value string) (gen string, n int, ok bool) { + rest, found := strings.CutPrefix(value, keychainChunkPrefix) + if !found { + return "", 0, false + } + gen, count, found := strings.Cut(rest, ":") + if !found || len(gen) != 16 { + return "", 0, false + } + if _, err := hex.DecodeString(gen); err != nil { + return "", 0, false + } + n, err := strconv.Atoi(count) + if err != nil || n < 1 || n > keychainMaxChunks { + return "", 0, false + } + return gen, n, true +} + +// deleteChunks removes the chunks a header points at, best effort: a leftover +// chunk is unreachable once no header names it. +func deleteChunks(key, header string) { + gen, n, ok := parseChunkHeader(header) + if !ok { + return + } + for i := range n { + _ = keyring.Delete(serviceName, chunkKey(key, gen, i)) } - return err } // IsAvailable checks if the system keychain is available. diff --git a/internal/adapters/keyring/keyring_chunk_test.go b/internal/adapters/keyring/keyring_chunk_test.go new file mode 100644 index 0000000..a480d4f --- /dev/null +++ b/internal/adapters/keyring/keyring_chunk_test.go @@ -0,0 +1,118 @@ +package keyring + +import ( + "strings" + "testing" + + "github.com/nylas/cli/internal/domain" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + gokeyring "github.com/zalando/go-keyring" +) + +// useSizeLimitedKeychain swaps in go-keyring's in-memory provider. It has no +// size limit of its own, so the tests check each stored item against the +// smallest real one (Windows Credential Manager, 2560 bytes) directly. +func useSizeLimitedKeychain(t *testing.T) { + t.Helper() + gokeyring.MockInit() + t.Cleanup(gokeyring.MockInit) +} + +func TestSystemKeyring_StoresAValueLargerThanOneKeychainItem(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + value := strings.Repeat("eyJhbGciOiJFZERTQSJ9.", 300) // ~6.6 KB, like a large JWT + + require.NoError(t, kr.Set("oauth_access_token", value)) + + got, err := kr.Get("oauth_access_token") + require.NoError(t, err) + assert.Equal(t, value, got) + + header, err := gokeyring.Get(serviceName, "oauth_access_token") + require.NoError(t, err) + gen, n, ok := parseChunkHeader(header) + require.True(t, ok, "the key holds a header, not the value") + for i := range n { + part, err := gokeyring.Get(serviceName, chunkKey("oauth_access_token", gen, i)) + require.NoError(t, err) + assert.LessOrEqual(t, len(part), 2560, "every item fits the smallest keychain limit") + } +} + +func TestSystemKeyring_ReplacingALargeValueRemovesItsChunks(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + require.NoError(t, kr.Set("k", strings.Repeat("a", 5000))) + header, _ := gokeyring.Get(serviceName, "k") + gen, _, _ := parseChunkHeader(header) + + require.NoError(t, kr.Set("k", "short")) + + got, err := kr.Get("k") + require.NoError(t, err) + assert.Equal(t, "short", got) + _, err = gokeyring.Get(serviceName, chunkKey("k", gen, 0)) + assert.ErrorIs(t, err, gokeyring.ErrNotFound) +} + +func TestSystemKeyring_ReplacingALargeValueWithAnotherUsesNewChunks(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + require.NoError(t, kr.Set("k", strings.Repeat("a", 5000))) + first, _ := gokeyring.Get(serviceName, "k") + + second := strings.Repeat("b", 4100) + require.NoError(t, kr.Set("k", second)) + + got, err := kr.Get("k") + require.NoError(t, err) + assert.Equal(t, second, got) + oldGen, _, _ := parseChunkHeader(first) + _, err = gokeyring.Get(serviceName, chunkKey("k", oldGen, 0)) + assert.ErrorIs(t, err, gokeyring.ErrNotFound) +} + +func TestSystemKeyring_DeleteRemovesTheChunks(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + require.NoError(t, kr.Set("k", strings.Repeat("a", 5000))) + header, _ := gokeyring.Get(serviceName, "k") + gen, n, _ := parseChunkHeader(header) + + require.NoError(t, kr.Delete("k")) + + _, err := kr.Get("k") + assert.ErrorIs(t, err, domain.ErrSecretNotFound) + for i := range n { + _, err := gokeyring.Get(serviceName, chunkKey("k", gen, i)) + assert.ErrorIs(t, err, gokeyring.ErrNotFound) + } +} + +func TestSystemKeyring_AMissingChunkIsAnErrorNotATruncatedValue(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + require.NoError(t, kr.Set("k", strings.Repeat("a", 5000))) + header, _ := gokeyring.Get(serviceName, "k") + gen, _, _ := parseChunkHeader(header) + require.NoError(t, gokeyring.Delete(serviceName, chunkKey("k", gen, 1))) + + _, err := kr.Get("k") + + require.Error(t, err) + assert.NotErrorIs(t, err, domain.ErrSecretNotFound) +} + +func TestSystemKeyring_AShortValueThatLooksLikeAHeaderRoundTrips(t *testing.T) { + useSizeLimitedKeychain(t) + kr := NewSystemKeyring() + value := keychainChunkPrefix + "0123456789abcdef:1" + + require.NoError(t, kr.Set("k", value)) + + got, err := kr.Get("k") + require.NoError(t, err) + assert.Equal(t, value, got) +} diff --git a/internal/adapters/mcp/proxy.go b/internal/adapters/mcp/proxy.go index 16a5ecf..04d8eb0 100644 --- a/internal/adapters/mcp/proxy.go +++ b/internal/adapters/mcp/proxy.go @@ -6,23 +6,27 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" + "log" + "maps" "math" "net/http" "os" "strings" "sync" + "github.com/nylas/cli/internal/domain" "github.com/nylas/cli/internal/httputil" "github.com/nylas/cli/internal/ports" ) const ( // NylasMCPEndpointUS is the US regional MCP endpoint. - NylasMCPEndpointUS = "https://mcp.us.nylas.com" + NylasMCPEndpointUS = domain.MCPResourceUS // NylasMCPEndpointEU is the EU regional MCP endpoint. - NylasMCPEndpointEU = "https://mcp.eu.nylas.com" + NylasMCPEndpointEU = domain.MCPResourceEU ) // GetMCPEndpoint returns the appropriate MCP endpoint for the given region. @@ -47,17 +51,27 @@ type rpcRequest struct { } `json:"params"` } +// mcpHTTPClient sends every request with its credential (an API key or an +// OAuth access token), so it follows no redirect. +var mcpHTTPClient = httputil.NewNoRedirectClient(httputil.DefaultClientTimeout) + // Proxy forwards MCP requests from STDIO to the Nylas MCP server. type Proxy struct { + // endpoint is the regional default, used when a credential names no + // server of its own. An OAuth credential always names one. endpoint string apiKey string - authHeader string // Cached "Bearer " value + creds ports.MCPCredentialSource + oauth bool defaultGrant string grantStore ports.GrantStore httpClient *http.Client - sessionID string - grantTools map[string]bool // Dynamically discovered tools that accept grant_id - mu sync.RWMutex + // protocolVersion is what the server answered initialize with, sent back + // as Mcp-Protocol-Version on every later request. There is no session: + // the hosted server is stateless and issues no Mcp-Session-Id. + protocolVersion string + grantTools map[string]bool // Dynamically discovered tools that accept grant_id + mu sync.RWMutex } // NewProxy creates a new MCP proxy with the given API key and region. @@ -65,8 +79,21 @@ func NewProxy(apiKey, region string) *Proxy { return &Proxy{ endpoint: GetMCPEndpoint(region), apiKey: apiKey, - authHeader: "Bearer " + apiKey, // Cache auth header - httpClient: httputil.DefaultClient, + creds: apiKeyCredentials{apiKey: apiKey}, + httpClient: mcpHTTPClient, + } +} + +// NewOAuthProxy creates an MCP proxy that authenticates with an OAuth +// session. The source is asked for a credential before every request, so a +// fifteen-minute access token is refreshed as it ages rather than failing +// the first request after it expires, and the MCP host is whatever the +// credential names — the audience the token was issued for. +func NewOAuthProxy(creds ports.MCPCredentialSource) *Proxy { + return &Proxy{ + creds: creds, + oauth: true, + httpClient: mcpHTTPClient, } } @@ -89,8 +116,13 @@ func (p *Proxy) SetGrantStore(store ports.GrantStore) { // Run starts the proxy, reading from stdin and writing to stdout. func (p *Proxy) Run(ctx context.Context) error { - reader := bufio.NewReader(os.Stdin) - writer := bufio.NewWriter(os.Stdout) + return p.serve(ctx, os.Stdin, os.Stdout) +} + +// serve runs the proxy loop over the given streams. +func (p *Proxy) serve(ctx context.Context, in io.Reader, out io.Writer) error { + reader := bufio.NewReader(in) + writer := bufio.NewWriter(out) for { select { @@ -146,6 +178,12 @@ func (p *Proxy) Run(ctx context.Context) error { // Forward to Nylas MCP server response, err := p.forward(ctx, line, &req) if err != nil { + // A notification must never be answered, not even with an error + // (JSON-RPC 2.0 §4.1), so its failure goes to the log instead. + if isNotification(line) { + log.Printf("mcp: forwarding %s notification failed: %v", req.Method, err) + continue + } // Write error response errorResp := p.createErrorResponse(&req, err) if _, writeErr := writer.Write(append(errorResp, '\n')); writeErr != nil { @@ -165,6 +203,18 @@ func (p *Proxy) Run(ctx context.Context) error { } } +// isNotification reports whether a JSON-RPC message is a notification: one +// with no "id" member at all. An explicit "id": null is still a request, and +// rpcRequest.ID cannot tell the two apart. +func isNotification(message []byte) bool { + var members map[string]json.RawMessage + if err := json.Unmarshal(message, &members); err != nil { + return false + } + _, hasID := members["id"] + return !hasID +} + // forward sends a request to the Nylas MCP server and returns the response. // The parsed rpcRequest is optional - if nil, request is forwarded as-is. func (p *Proxy) forward(ctx context.Context, request []byte, parsed *rpcRequest) ([]byte, error) { @@ -172,48 +222,109 @@ func (p *Proxy) forward(ctx context.Context, request []byte, parsed *rpcRequest) isToolsList := parsed != nil && parsed.Method == "tools/list" isInitialize := parsed != nil && parsed.Method == "initialize" + cred, err := p.credential(ctx) + if err != nil { + return nil, err + } + + renewed := false + for { + resp, err := p.send(ctx, request, parsed, cred) + if err != nil { + return nil, err + } + + // A 401 carrying a Bearer challenge means the token was refused, not + // that the request was bad: renew once and resend. A second refusal + // is reported rather than retried, so a revoked session cannot loop. + if resp.StatusCode == http.StatusUnauthorized && !renewed && + domain.ParseBearerChallenge(resp.Header.Get("WWW-Authenticate")) != nil { + next, renewErr := p.renew(ctx, cred) + if renewErr == nil { + drainAndClose(resp) + cred = next + renewed = true + continue + } + if !errors.Is(renewErr, domain.ErrMCPCredentialNotRenewable) { + drainAndClose(resp) + return nil, renewErr + } + } + + body, err := p.readResponse(resp, cred) + if err != nil { + return nil, err + } + // Modify responses as needed + if isToolsList { + body = p.modifyToolsListResponse(body) + } + if isInitialize { + p.rememberProtocolVersion(body) + body = p.modifyInitializeResponse(body) + } + return body, nil + } +} + +// send makes one HTTP request with cred. The caller closes the response. +func (p *Proxy) send(ctx context.Context, request []byte, parsed *rpcRequest, cred *domain.MCPCredential) (*http.Response, error) { + endpoint := cred.Endpoint + if endpoint == "" { + endpoint = p.endpoint + } + if endpoint == "" { + return nil, errors.New("no MCP server to send the request to") + } + + p.mu.RLock() + defaultGrant := p.defaultGrant + protocolVersion := p.protocolVersion + p.mu.RUnlock() + + // The default grant is only a hint, and an OAuth token only acts on the + // grants it names: offering any other would be refused at best. + grantHint := "" + if cred.AllowsGrantHint(defaultGrant) { + grantHint = defaultGrant + } + + // Every attempt starts from the request as the assistant sent it: the + // grant injection below writes into the parsed arguments, and a retry + // after renewal must not inherit the previous credential's grant. + parsed = cloneRPCRequest(parsed) + // Inject default grant into tool calls if not specified - request = p.injectDefaultGrant(request, parsed) + request = p.injectGrant(request, parsed, grantHint) // Normalize tool arguments (type coercion, timestamp rounding) request = p.normalizeToolArguments(request, parsed) - req, err := http.NewRequestWithContext(ctx, "POST", p.endpoint, bytes.NewReader(request)) + req, err := http.NewRequestWithContext(ctx, "POST", endpoint, bytes.NewReader(request)) if err != nil { return nil, fmt.Errorf("creating request: %w", err) } - // Set required headers (use cached auth header) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json, text/event-stream") - req.Header.Set("Authorization", p.authHeader) - - // Include session ID and default grant if we have them (read lock) - p.mu.RLock() - if p.sessionID != "" { - req.Header.Set("Mcp-Session-Id", p.sessionID) + req.Header.Set("Authorization", "Bearer "+cred.Token) + setProtocolHeaders(req.Header, parsed, protocolVersion) + if grantHint != "" { + req.Header.Set("X-Nylas-Grant-Id", grantHint) } - if p.defaultGrant != "" { - req.Header.Set("X-Nylas-Grant-Id", p.defaultGrant) - } - p.mu.RUnlock() - // Send request resp, err := p.httpClient.Do(req) if err != nil { return nil, fmt.Errorf("sending request: %w", err) } - defer func() { _ = resp.Body.Close() }() - - // Store session ID if provided - if sessionID := resp.Header.Get("Mcp-Session-Id"); sessionID != "" { - p.mu.Lock() - p.sessionID = sessionID - p.mu.Unlock() - } + return resp, nil +} - // Handle response based on content type - contentType := resp.Header.Get("Content-Type") +// readResponse turns an HTTP response into the JSON-RPC payload to write +// back, closing it. +func (p *Proxy) readResponse(resp *http.Response, cred *domain.MCPCredential) ([]byte, error) { + defer func() { _ = resp.Body.Close() }() // Handle 202 Accepted (no body) if resp.StatusCode == http.StatusAccepted { @@ -222,24 +333,13 @@ func (p *Proxy) forward(ctx context.Context, request []byte, parsed *rpcRequest) // Handle errors if resp.StatusCode != http.StatusOK { - body, _ := io.ReadAll(resp.Body) - return nil, fmt.Errorf("server returned %d: %s", resp.StatusCode, string(body)) + body, _ := io.ReadAll(io.LimitReader(resp.Body, maxErrorBody)) + return nil, p.statusError(resp, body, cred) } // Handle SSE stream - if strings.HasPrefix(contentType, "text/event-stream") { - body, err := p.readSSE(resp.Body) - if err != nil { - return nil, err - } - // Modify responses as needed - if isToolsList { - body = p.modifyToolsListResponse(body) - } - if isInitialize { - body = p.modifyInitializeResponse(body) - } - return body, nil + if strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") { + return p.readSSE(resp.Body) } // Handle JSON response @@ -247,15 +347,6 @@ func (p *Proxy) forward(ctx context.Context, request []byte, parsed *rpcRequest) if err != nil { return nil, fmt.Errorf("reading response: %w", err) } - - // Modify responses as needed - if isToolsList { - body = p.modifyToolsListResponse(body) - } - if isInitialize { - body = p.modifyInitializeResponse(body) - } - return body, nil } @@ -363,7 +454,12 @@ func (p *Proxy) injectDefaultGrant(request []byte, parsed *rpcRequest) []byte { p.mu.RLock() defaultGrant := p.defaultGrant p.mu.RUnlock() + return p.injectGrant(request, parsed, defaultGrant) +} +// injectGrant injects defaultGrant as grant_id into a tool call that accepts +// one and names none. An empty defaultGrant injects nothing. +func (p *Proxy) injectGrant(request []byte, parsed *rpcRequest, defaultGrant string) []byte { if defaultGrant == "" { return request } @@ -531,3 +627,14 @@ func toInt64(v any) (int64, bool) { return 0, false } } + +// cloneRPCRequest copies req deeply enough for injectGrant and +// normalizeToolArguments to change the copy's arguments freely. +func cloneRPCRequest(req *rpcRequest) *rpcRequest { + if req == nil { + return nil + } + clone := *req + clone.Params.Arguments = maps.Clone(req.Params.Arguments) + return &clone +} diff --git a/internal/adapters/mcp/proxy_auth.go b/internal/adapters/mcp/proxy_auth.go new file mode 100644 index 0000000..9497461 --- /dev/null +++ b/internal/adapters/mcp/proxy_auth.go @@ -0,0 +1,99 @@ +package mcp + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + + "github.com/nylas/cli/internal/domain" +) + +// OAuthLoginCommand is what a user runs to get a session the hosted MCP +// server accepts. Every OAuth failure the proxy reports names it. +const OAuthLoginCommand = "nylas oauth login --for mcp" + +// maxErrorBody bounds how much of an error response is read and echoed. +const maxErrorBody = 4 << 10 + +// apiKeyCredentials is the credential source behind NewProxy: the same API +// key for every request, and nothing to renew. +type apiKeyCredentials struct { + apiKey string +} + +func (a apiKeyCredentials) Credential(context.Context) (*domain.MCPCredential, error) { + return &domain.MCPCredential{Token: a.apiKey}, nil +} + +func (apiKeyCredentials) Renew(context.Context, *domain.MCPCredential) (*domain.MCPCredential, error) { + return nil, domain.ErrMCPCredentialNotRenewable +} + +// credential asks the source for this request's credential. +func (p *Proxy) credential(ctx context.Context) (*domain.MCPCredential, error) { + if p.creds == nil { + return nil, errors.New("MCP proxy has no credential source") + } + cred, err := p.creds.Credential(ctx) + if err != nil { + return nil, p.loginError("no usable OAuth session", err) + } + if cred == nil || cred.Token == "" { + return nil, p.loginError("no usable OAuth session", errors.New("credential source returned no token")) + } + return cred, nil +} + +// renew asks the source for a replacement after a 401. +func (p *Proxy) renew(ctx context.Context, rejected *domain.MCPCredential) (*domain.MCPCredential, error) { + next, err := p.creds.Renew(ctx, rejected) + if err != nil { + if errors.Is(err, domain.ErrMCPCredentialNotRenewable) { + return nil, err + } + return nil, p.loginError("the Nylas MCP server rejected the OAuth token and it could not be refreshed", err) + } + if next == nil || next.Token == "" { + return nil, p.loginError("the Nylas MCP server rejected the OAuth token", errors.New("refresh returned no token")) + } + return next, nil +} + +// loginError says what went wrong and what to run. An API key proxy has no +// login command to offer, so its errors pass through unchanged. +func (p *Proxy) loginError(what string, cause error) error { + if !p.oauth { + return cause + } + return fmt.Errorf("%s: %w. Run `%s` to sign in again", what, cause, OAuthLoginCommand) +} + +// statusError explains a non-200 answer. The two OAuth refusals get a +// sentence the user can act on; everything else keeps the status and body. +func (p *Proxy) statusError(resp *http.Response, body []byte, cred *domain.MCPCredential) error { + challenge := domain.ParseBearerChallenge(resp.Header.Get("WWW-Authenticate")) + + if resp.StatusCode == http.StatusForbidden && challenge != nil && challenge.Error == "insufficient_scope" { + missing := "a scope" + if scopes := challenge.Scopes(); len(scopes) > 0 { + missing = "scope " + strings.Join(scopes, ", ") + } + return fmt.Errorf("the Nylas MCP server requires %s, which this login was not granted. Run `%s` to sign in with the scopes the MCP tools use", + missing, OAuthLoginCommand) + } + + if resp.StatusCode == http.StatusUnauthorized && p.oauth && cred != nil { + return fmt.Errorf("the Nylas MCP server rejected the OAuth token (HTTP 401). Run `%s` to sign in again", + OAuthLoginCommand) + } + + return fmt.Errorf("server returned %d: %s", resp.StatusCode, string(body)) +} + +func drainAndClose(resp *http.Response) { + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, maxErrorBody)) + _ = resp.Body.Close() +} diff --git a/internal/adapters/mcp/proxy_e2e_test.go b/internal/adapters/mcp/proxy_e2e_test.go index e33777b..f75fdc8 100644 --- a/internal/adapters/mcp/proxy_e2e_test.go +++ b/internal/adapters/mcp/proxy_e2e_test.go @@ -14,11 +14,12 @@ import ( // mockMCPServer simulates the upstream Nylas MCP server for E2E proxy tests. // It returns realistic tools/list, initialize, and tools/call responses. type mockMCPServer struct { - t *testing.T - lastToolCall string - lastArgs map[string]any - receivedGrant string - receivedMethod string + t *testing.T + lastToolCall string + lastArgs map[string]any + receivedGrant string + receivedMethod string + receivedSessionID string } func (m *mockMCPServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { @@ -32,6 +33,7 @@ func (m *mockMCPServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { } m.receivedMethod = req.Method + m.receivedSessionID = r.Header.Get("Mcp-Session-Id") w.Header().Set("Content-Type", "application/json") w.Header().Set("Mcp-Session-Id", "e2e-session-001") @@ -301,10 +303,12 @@ func TestE2E_ProxyLifecycle(t *testing.T) { } }) - // === Step 9: session ID stored from server response === - t.Run("session_id_stored", func(t *testing.T) { - if proxy.sessionID != "e2e-session-001" { - t.Errorf("expected session ID 'e2e-session-001', got %q", proxy.sessionID) + // === Step 9: no session — the server is stateless === + // The mock hands out an Mcp-Session-Id on every answer; the proxy must + // not carry it into the requests that follow. + t.Run("session_id_not_echoed", func(t *testing.T) { + if mock.receivedSessionID != "" { + t.Errorf("expected no Mcp-Session-Id on requests, got %q", mock.receivedSessionID) } }) diff --git a/internal/adapters/mcp/proxy_forward_test.go b/internal/adapters/mcp/proxy_forward_test.go index 58c71e4..0f13bfa 100644 --- a/internal/adapters/mcp/proxy_forward_test.go +++ b/internal/adapters/mcp/proxy_forward_test.go @@ -6,11 +6,13 @@ import ( "net/http" "net/http/httptest" "strings" + "sync/atomic" "testing" ) func TestProxy_forward(t *testing.T) { t.Parallel() + var sawSessionID atomic.Bool // Create a mock server server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -21,6 +23,9 @@ func TestProxy_forward(t *testing.T) { if r.Header.Get("Content-Type") != "application/json" { t.Errorf("expected Content-Type 'application/json', got '%s'", r.Header.Get("Content-Type")) } + if r.Header.Get("Mcp-Session-Id") != "" { + sawSessionID.Store(true) + } // Return a response w.Header().Set("Content-Type", "application/json") @@ -49,9 +54,13 @@ func TestProxy_forward(t *testing.T) { t.Errorf("expected jsonrpc '2.0', got '%v'", resp["jsonrpc"]) } - // Verify session ID was stored - if proxy.sessionID != "test-session-123" { - t.Errorf("expected sessionID 'test-session-123', got '%s'", proxy.sessionID) + // The hosted server is stateless: a session id it sends is not kept, + // so there is nothing to echo on the next request. + if _, err := proxy.forward(t.Context(), request, nil); err != nil { + t.Fatalf("second forward failed: %v", err) + } + if sawSessionID.Load() { + t.Error("the proxy must not send Mcp-Session-Id to the stateless server") } } diff --git a/internal/adapters/mcp/proxy_notification_test.go b/internal/adapters/mcp/proxy_notification_test.go new file mode 100644 index 0000000..dba03e4 --- /dev/null +++ b/internal/adapters/mcp/proxy_notification_test.go @@ -0,0 +1,63 @@ +package mcp + +import ( + "bufio" + "bytes" + "encoding/json" + "net/http" + "strings" + "testing" + + "github.com/nylas/cli/internal/domain" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestProxy_FailedNotificationIsNeverAnswered(t *testing.T) { + // Every upstream call is refused, so each message fails to forward. Only + // the two requests may get a reply; the notification between them must not. + refuse := respondStatus(http.StatusUnauthorized, "") + _, server := newOAuthTestProxy(t, refuse, refuse, refuse) + proxy := NewOAuthProxy(&fakeCredentials{credentials: []*domain.MCPCredential{ + oauthCred(server, "a"), oauthCred(server, "b"), oauthCred(server, "c"), + }}) + + in := strings.Join([]string{ + `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`, + `{"jsonrpc":"2.0","method":"notifications/initialized"}`, + `{"jsonrpc":"2.0","id":null,"method":"ping"}`, + }, "\n") + "\n" + var out bytes.Buffer + + require.NoError(t, proxy.serve(t.Context(), strings.NewReader(in), &out)) + + var ids []any + scanner := bufio.NewScanner(&out) + for scanner.Scan() { + var resp map[string]any + require.NoError(t, json.Unmarshal(scanner.Bytes(), &resp)) + require.Contains(t, resp, "error") + ids = append(ids, resp["id"]) + } + assert.Equal(t, []any{float64(1), nil}, ids) +} + +func TestIsNotification(t *testing.T) { + tests := []struct { + name string + message string + want bool + }{ + {"no id member", `{"jsonrpc":"2.0","method":"notifications/initialized"}`, true}, + {"numeric id", `{"jsonrpc":"2.0","id":1,"method":"ping"}`, false}, + {"string id", `{"jsonrpc":"2.0","id":"a","method":"ping"}`, false}, + {"explicit null id is a request", `{"jsonrpc":"2.0","id":null,"method":"ping"}`, false}, + {"not an object", `[{"jsonrpc":"2.0","method":"x"}]`, false}, + {"invalid JSON", `{`, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.want, isNotification([]byte(tt.message))) + }) + } +} diff --git a/internal/adapters/mcp/proxy_oauth_test.go b/internal/adapters/mcp/proxy_oauth_test.go new file mode 100644 index 0000000..9888c90 --- /dev/null +++ b/internal/adapters/mcp/proxy_oauth_test.go @@ -0,0 +1,381 @@ +package mcp + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/nylas/cli/internal/domain" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// fakeMCPServer records what reached it and answers each request with the +// next scripted response. +type fakeMCPServer struct { + mu sync.Mutex + requests []recordedRequest + script []func(w http.ResponseWriter) +} + +type recordedRequest struct { + authorization string + grantHeader string + body map[string]any +} + +func (f *fakeMCPServer) ServeHTTP(w http.ResponseWriter, r *http.Request) { + raw, _ := io.ReadAll(r.Body) + var body map[string]any + _ = json.Unmarshal(raw, &body) + + f.mu.Lock() + f.requests = append(f.requests, recordedRequest{ + authorization: r.Header.Get("Authorization"), + grantHeader: r.Header.Get("X-Nylas-Grant-Id"), + body: body, + }) + var respond func(w http.ResponseWriter) + if len(f.script) > 0 { + respond, f.script = f.script[0], f.script[1:] + } + f.mu.Unlock() + + if respond == nil { + respond = respondOK + } + respond(w) +} + +func (f *fakeMCPServer) recorded() []recordedRequest { + f.mu.Lock() + defer f.mu.Unlock() + return append([]recordedRequest(nil), f.requests...) +} + +func respondOK(w http.ResponseWriter) { + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"ok":true}}`)) +} + +func respondStatus(status int, challenge string) func(w http.ResponseWriter) { + return func(w http.ResponseWriter) { + if challenge != "" { + w.Header().Set("WWW-Authenticate", challenge) + } + w.WriteHeader(status) + _, _ = w.Write([]byte(`{"error":"refused"}`)) + } +} + +// fakeCredentials is a scripted ports.MCPCredentialSource. +type fakeCredentials struct { + mu sync.Mutex + credentials []*domain.MCPCredential // returned in turn by Credential + renewed *domain.MCPCredential + credErr error + renewErr error + calls int + renewCalls []*domain.MCPCredential +} + +func (f *fakeCredentials) Credential(context.Context) (*domain.MCPCredential, error) { + f.mu.Lock() + defer f.mu.Unlock() + if f.credErr != nil { + return nil, f.credErr + } + cred := f.credentials[min(f.calls, len(f.credentials)-1)] + f.calls++ + return cred, nil +} + +func (f *fakeCredentials) Renew(_ context.Context, rejected *domain.MCPCredential) (*domain.MCPCredential, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.renewCalls = append(f.renewCalls, rejected) + if f.renewErr != nil { + return nil, f.renewErr + } + return f.renewed, nil +} + +func newOAuthTestProxy(t *testing.T, script ...func(w http.ResponseWriter)) (*fakeMCPServer, *httptest.Server) { + t.Helper() + fake := &fakeMCPServer{script: script} + server := httptest.NewServer(fake) + t.Cleanup(server.Close) + return fake, server +} + +func oauthCred(server *httptest.Server, token string, grantIDs ...string) *domain.MCPCredential { + cred := &domain.MCPCredential{Token: token, Endpoint: server.URL, GrantScoped: true} + for _, id := range grantIDs { + cred.Grants = append(cred.Grants, domain.OAuthTokenGrant{ID: id, ApplicationID: "app-1"}) + } + return cred +} + +const toolCall = `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"list_messages","arguments":{}}}` + +func TestOAuthProxy_AsksForACredentialBeforeEveryRequest(t *testing.T) { + // A fifteen-minute token must not be cached for the life of the process. + fake, server := newOAuthTestProxy(t) + creds := &fakeCredentials{credentials: []*domain.MCPCredential{ + oauthCred(server, "token-1"), + oauthCred(server, "token-2"), + }} + proxy := NewOAuthProxy(creds) + + for range 2 { + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + require.NoError(t, err) + } + + got := fake.recorded() + require.Len(t, got, 2) + assert.Equal(t, "Bearer token-1", got[0].authorization) + assert.Equal(t, "Bearer token-2", got[1].authorization) +} + +func TestOAuthProxy_SendsToTheCredentialsEndpoint(t *testing.T) { + fake, server := newOAuthTestProxy(t) + proxy := NewOAuthProxy(&fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "t")}}) + proxy.endpoint = "http://127.0.0.1:1" // a regional default the credential must override + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + require.NoError(t, err) + assert.Len(t, fake.recorded(), 1) +} + +func TestOAuthProxy_GrantHintOnlyForGrantsTheTokenNames(t *testing.T) { + tests := []struct { + name string + tokenGrant []string + wantHint string + }{ + {"listed", []string{"grant-1", "grant-2"}, "grant-1"}, + {"not listed", []string{"grant-2"}, ""}, + {"no grants", nil, ""}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + fake, server := newOAuthTestProxy(t) + proxy := NewOAuthProxy(&fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "t", tt.tokenGrant...)}}) + proxy.SetDefaultGrant("grant-1") + + req := parseRPC(toolCall) + _, err := proxy.forward(t.Context(), req.raw, req.parsed) + require.NoError(t, err) + + got := fake.recorded() + require.Len(t, got, 1) + assert.Equal(t, tt.wantHint, got[0].grantHeader) + + args := got[0].body["params"].(map[string]any)["arguments"].(map[string]any) + if tt.wantHint == "" { + assert.NotContains(t, args, "grant_id", "an unlisted grant must not be injected either") + } else { + assert.Equal(t, tt.wantHint, args["grant_id"]) + } + }) + } +} + +func TestOAuthProxy_RenewsOnceOn401ChallengeAndRetries(t *testing.T) { + fake, server := newOAuthTestProxy(t, + respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`), + respondOK, + ) + stale := oauthCred(server, "stale") + creds := &fakeCredentials{credentials: []*domain.MCPCredential{stale}, renewed: oauthCred(server, "fresh")} + proxy := NewOAuthProxy(creds) + + resp, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + require.NoError(t, err) + assert.Contains(t, string(resp), `"ok":true`) + + got := fake.recorded() + require.Len(t, got, 2) + assert.Equal(t, "Bearer stale", got[0].authorization) + assert.Equal(t, "Bearer fresh", got[1].authorization) + require.Len(t, creds.renewCalls, 1) + assert.Same(t, stale, creds.renewCalls[0], "renew is told which token was refused") +} + +func TestOAuthProxy_SecondRefusalTellsTheUserToLogIn(t *testing.T) { + challenge := `Bearer error="invalid_token"` + fake, server := newOAuthTestProxy(t, + respondStatus(http.StatusUnauthorized, challenge), + respondStatus(http.StatusUnauthorized, challenge), + respondOK, + ) + creds := &fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "a")}, renewed: oauthCred(server, "b")} + proxy := NewOAuthProxy(creds) + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.Error(t, err) + assert.Contains(t, err.Error(), OAuthLoginCommand) + assert.Len(t, creds.renewCalls, 1, "renew once, never loop") + assert.Len(t, fake.recorded(), 2) +} + +func TestOAuthProxy_401WithoutChallengeIsNotRenewed(t *testing.T) { + fake, server := newOAuthTestProxy(t, respondStatus(http.StatusUnauthorized, "")) + creds := &fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "a")}, renewed: oauthCred(server, "b")} + proxy := NewOAuthProxy(creds) + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.Error(t, err) + assert.Contains(t, err.Error(), OAuthLoginCommand) + assert.Empty(t, creds.renewCalls) + assert.Len(t, fake.recorded(), 1) +} + +func TestOAuthProxy_FailedRenewalTellsTheUserToLogIn(t *testing.T) { + _, server := newOAuthTestProxy(t, respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`)) + creds := &fakeCredentials{ + credentials: []*domain.MCPCredential{oauthCred(server, "a")}, + renewErr: &domain.OAuthError{Code: "invalid_grant", StatusCode: http.StatusBadRequest}, + } + proxy := NewOAuthProxy(creds) + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.Error(t, err) + assert.Contains(t, err.Error(), OAuthLoginCommand) + var oauthErr *domain.OAuthError + assert.True(t, errors.As(err, &oauthErr), "the cause is kept for anyone who inspects it") +} + +func TestOAuthProxy_InsufficientScopeNamesTheScope(t *testing.T) { + fake, server := newOAuthTestProxy(t, + respondStatus(http.StatusForbidden, `Bearer error="insufficient_scope", scope="email.send"`), + ) + creds := &fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "a")}, renewed: oauthCred(server, "b")} + proxy := NewOAuthProxy(creds) + + req := parseRPC(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"send_message","arguments":{}}}`) + _, err := proxy.forward(t.Context(), req.raw, req.parsed) + + require.Error(t, err) + assert.Contains(t, err.Error(), "email.send") + assert.Contains(t, err.Error(), OAuthLoginCommand) + assert.Empty(t, creds.renewCalls, "a refresh cannot add a scope that was never granted") + assert.Len(t, fake.recorded(), 1) +} + +func TestOAuthProxy_NoSessionFailsBeforeAnyRequest(t *testing.T) { + fake, _ := newOAuthTestProxy(t) + proxy := NewOAuthProxy(&fakeCredentials{credErr: domain.ErrOAuthNotLoggedIn}) + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.ErrorIs(t, err, domain.ErrOAuthNotLoggedIn) + assert.Contains(t, err.Error(), OAuthLoginCommand) + assert.Empty(t, fake.recorded(), "nothing may be sent without a credential") +} + +func TestOAuthProxy_ErrorsNeverCarryTheToken(t *testing.T) { + _, server := newOAuthTestProxy(t, + respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`), + respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`), + ) + secret := "eyJ-very-secret-access-token" + proxy := NewOAuthProxy(&fakeCredentials{ + credentials: []*domain.MCPCredential{oauthCred(server, secret)}, + renewed: oauthCred(server, secret+"-2"), + }) + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.Error(t, err) + assert.NotContains(t, err.Error(), secret) +} + +func TestAPIKeyProxy_401ChallengeIsNotRenewed(t *testing.T) { + // An API key has nothing to refresh; the answer is reported as it was. + fake, server := newOAuthTestProxy(t, respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`)) + proxy := NewProxy("api-key", "us") + proxy.endpoint = server.URL + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + require.Error(t, err) + assert.Contains(t, err.Error(), "401") + assert.NotContains(t, err.Error(), OAuthLoginCommand) + assert.Len(t, fake.recorded(), 1) + assert.Equal(t, "Bearer api-key", fake.recorded()[0].authorization) +} + +func TestAPIKeyProxy_GrantHintIsUnrestricted(t *testing.T) { + fake, server := newOAuthTestProxy(t) + proxy := NewProxy("api-key", "us") + proxy.endpoint = server.URL + proxy.SetDefaultGrant("grant-9") + + _, err := proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + require.NoError(t, err) + + assert.Equal(t, "grant-9", fake.recorded()[0].grantHeader) +} + +func TestOAuthProxy_RetryAfterRenewalRebuildsTheGrantFromTheOriginalRequest(t *testing.T) { + // The first attempt writes grant_id into the parsed arguments. The retry + // must start from what the assistant sent, not from that edited copy. + for _, tool := range []string{"list_messages", "list_events"} { + for _, tt := range []struct { + name string + freshGrant []string + want any + }{ + {"renewed token still lists the grant", []string{"grant-1"}, "grant-1"}, + {"renewed token no longer lists it", []string{"grant-2"}, nil}, + } { + t.Run(tool+"/"+tt.name, func(t *testing.T) { + fake, server := newOAuthTestProxy(t, + respondStatus(http.StatusUnauthorized, `Bearer error="invalid_token"`), + respondOK, + ) + creds := &fakeCredentials{ + credentials: []*domain.MCPCredential{oauthCred(server, "stale", "grant-1")}, + renewed: oauthCred(server, "fresh", tt.freshGrant...), + } + proxy := NewOAuthProxy(creds) + proxy.SetDefaultGrant("grant-1") + + req := parseRPC(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + tool + `","arguments":{}}}`) + _, err := proxy.forward(t.Context(), req.raw, req.parsed) + require.NoError(t, err) + + got := fake.recorded() + require.Len(t, got, 2) + args := got[1].body["params"].(map[string]any)["arguments"].(map[string]any) + assert.Equal(t, tt.want, args["grant_id"]) + assert.Equal(t, "Bearer fresh", got[1].authorization) + }) + } + } +} + +func TestProxy_DoesNotFollowRedirectsWithTheCredential(t *testing.T) { + stolen := 0 + elsewhere := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { stolen++ })) + t.Cleanup(elsewhere.Close) + _, server := newOAuthTestProxy(t, func(w http.ResponseWriter) { + w.Header().Set("Location", elsewhere.URL) + w.WriteHeader(http.StatusTemporaryRedirect) + }) + proxy := NewOAuthProxy(&fakeCredentials{credentials: []*domain.MCPCredential{oauthCred(server, "t")}}) + + _, _ = proxy.forward(t.Context(), []byte(`{"jsonrpc":"2.0","id":1,"method":"ping"}`), nil) + + assert.Zero(t, stolen, "the bearer token must not follow a redirect") +} diff --git a/internal/adapters/mcp/proxy_protocol.go b/internal/adapters/mcp/proxy_protocol.go new file mode 100644 index 0000000..66dccba --- /dev/null +++ b/internal/adapters/mcp/proxy_protocol.go @@ -0,0 +1,69 @@ +package mcp + +import ( + "encoding/json" + "net/http" + "regexp" +) + +// Streamable HTTP request headers the hosted Nylas MCP server reads +// (api-v3 src/mcp/server.go). Mcp-Method and Mcp-Name let it route and +// meter a request without parsing the body, and it refuses a request whose +// Mcp-Name disagrees with the body. There is deliberately no Mcp-Session-Id: +// the server runs stateless — every request stands alone — so the proxy +// neither stores nor sends one. +const ( + headerMCPMethod = "Mcp-Method" + headerMCPName = "Mcp-Name" + headerMCPProtocolVersion = "Mcp-Protocol-Version" +) + +// headerSafeName allow-lists what may be copied from a JSON-RPC body into a +// header: method and tool names are short identifiers. Anything else is left +// to the body, which the server still parses for clients that send no +// headers at all. +var headerSafeName = regexp.MustCompile(`^[A-Za-z0-9_./\-]{1,128}$`) + +// protocolVersionPattern is the MCP date-stamped version shape. +var protocolVersionPattern = regexp.MustCompile(`^\d{4}-\d{2}-\d{2}$`) + +// setProtocolHeaders adds the MCP routing headers for a parsed request. +// Mcp-Name is sent for the methods whose params carry a name; if that name +// cannot go in a header, neither header is sent, because a method header +// without its name reads as a mismatch rather than as a legacy request. +func setProtocolHeaders(h http.Header, parsed *rpcRequest, protocolVersion string) { + if protocolVersion != "" { + h.Set(headerMCPProtocolVersion, protocolVersion) + } + if parsed == nil || !headerSafeName.MatchString(parsed.Method) { + return + } + switch parsed.Method { + case "tools/call", "prompts/get": + if !headerSafeName.MatchString(parsed.Params.Name) { + return + } + h.Set(headerMCPName, parsed.Params.Name) + } + h.Set(headerMCPMethod, parsed.Method) +} + +// rememberProtocolVersion records the version the server negotiated in its +// initialize answer. Later requests carry it as Mcp-Protocol-Version, as the +// streamable HTTP transport requires once a version has been agreed. +func (p *Proxy) rememberProtocolVersion(initializeResponse []byte) { + var resp struct { + Result struct { + ProtocolVersion string `json:"protocolVersion"` + } `json:"result"` + } + if err := json.Unmarshal(initializeResponse, &resp); err != nil { + return + } + if !protocolVersionPattern.MatchString(resp.Result.ProtocolVersion) { + return + } + p.mu.Lock() + p.protocolVersion = resp.Result.ProtocolVersion + p.mu.Unlock() +} diff --git a/internal/adapters/mcp/proxy_protocol_test.go b/internal/adapters/mcp/proxy_protocol_test.go new file mode 100644 index 0000000..847232b --- /dev/null +++ b/internal/adapters/mcp/proxy_protocol_test.go @@ -0,0 +1,110 @@ +package mcp + +import ( + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// headerRecorder answers like the stateless hosted server: it hands out a +// session id the proxy must ignore, and negotiates a protocol version. +type headerRecorder struct { + mu sync.Mutex + headers []http.Header +} + +func (h *headerRecorder) ServeHTTP(w http.ResponseWriter, r *http.Request) { + h.mu.Lock() + h.headers = append(h.headers, r.Header.Clone()) + h.mu.Unlock() + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Mcp-Session-Id", "should-be-ignored") + _, _ = w.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2025-06-18","instructions":"x"}}`)) +} + +func (h *headerRecorder) all() []http.Header { + h.mu.Lock() + defer h.mu.Unlock() + return append([]http.Header(nil), h.headers...) +} + +func newProtocolProxy(t *testing.T) (*Proxy, *headerRecorder) { + t.Helper() + rec := &headerRecorder{} + server := httptest.NewServer(rec) + t.Cleanup(server.Close) + proxy := NewProxy("k", "us") + proxy.endpoint = server.URL + return proxy, rec +} + +func TestProxy_SendsDocumentedProtocolHeaders(t *testing.T) { + proxy, rec := newProtocolProxy(t) + + for _, raw := range []string{ + `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18"}}`, + `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`, + `{"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"name":"list_messages","arguments":{}}}`, + `{"jsonrpc":"2.0","id":4,"method":"prompts/get","params":{"name":"summarise"}}`, + } { + req := parseRPC(raw) + _, err := proxy.forward(t.Context(), req.raw, req.parsed) + require.NoError(t, err) + } + + got := rec.all() + require.Len(t, got, 4) + + assert.Equal(t, "initialize", got[0].Get("Mcp-Method")) + assert.Empty(t, got[0].Get("Mcp-Protocol-Version"), "nothing is negotiated before initialize answers") + + assert.Equal(t, "tools/list", got[1].Get("Mcp-Method")) + assert.Empty(t, got[1].Get("Mcp-Name")) + assert.Equal(t, "2025-06-18", got[1].Get("Mcp-Protocol-Version"), "the negotiated version follows") + + assert.Equal(t, "tools/call", got[2].Get("Mcp-Method")) + assert.Equal(t, "list_messages", got[2].Get("Mcp-Name"), "the server refuses a Mcp-Name that disagrees with the body") + + assert.Equal(t, "prompts/get", got[3].Get("Mcp-Method")) + assert.Equal(t, "summarise", got[3].Get("Mcp-Name")) + + for i, h := range got { + assert.Empty(t, h.Get("Mcp-Session-Id"), "request %d: the server is stateless", i) + } +} + +func TestProxy_UnsafeNamesStayInTheBody(t *testing.T) { + proxy, rec := newProtocolProxy(t) + + req := parseRPC(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"bad name\r\nX-Injected: 1","arguments":{}}}`) + _, err := proxy.forward(t.Context(), req.raw, req.parsed) + require.NoError(t, err) + + got := rec.all() + require.Len(t, got, 1) + assert.Empty(t, got[0].Get("Mcp-Method"), "no method header without its name") + assert.Empty(t, got[0].Get("Mcp-Name")) + assert.Empty(t, got[0].Get("X-Injected")) +} + +func TestProxy_UnparsedRequestGetsNoRoutingHeaders(t *testing.T) { + proxy, rec := newProtocolProxy(t) + + _, err := proxy.forward(t.Context(), []byte(`[{"jsonrpc":"2.0","id":1,"method":"ping"}]`), nil) + require.NoError(t, err) + + assert.Empty(t, rec.all()[0].Get("Mcp-Method")) +} + +func TestRememberProtocolVersion_IgnoresMalformedValues(t *testing.T) { + proxy := NewProxy("k", "us") + proxy.rememberProtocolVersion([]byte(`{"result":{"protocolVersion":"2025-06-18\r\nX: y"}}`)) + assert.Empty(t, proxy.protocolVersion) + + proxy.rememberProtocolVersion([]byte(`{"result":{"protocolVersion":"2026-07-28"}}`)) + assert.Equal(t, "2026-07-28", proxy.protocolVersion) +} diff --git a/internal/adapters/oauth/mock.go b/internal/adapters/oauth/mock.go index 711213b..774599f 100644 --- a/internal/adapters/oauth/mock.go +++ b/internal/adapters/oauth/mock.go @@ -12,10 +12,14 @@ type MockServer struct { Port int AuthCode string ExpectedState string + SetState string StartCalled bool StopCalled bool WaitForCallbackCalled bool TimeoutAfter time.Duration + + // RedirectURI overrides the advertised redirect URI when set. + RedirectURI string } // NewMockServer creates a new MockServer. @@ -38,6 +42,12 @@ func (m *MockServer) Stop() error { return nil } +// SetExpectedState records the state the service set before the browser +// opened. WaitForCallback records its own copy in ExpectedState. +func (m *MockServer) SetExpectedState(state string) { + m.SetState = state +} + // WaitForCallback waits for the OAuth callback. func (m *MockServer) WaitForCallback(ctx context.Context, expectedState string) (string, error) { m.WaitForCallbackCalled = true @@ -60,5 +70,8 @@ func (m *MockServer) WaitForCallback(ctx context.Context, expectedState string) // GetRedirectURI returns the redirect URI. func (m *MockServer) GetRedirectURI() string { + if m.RedirectURI != "" { + return m.RedirectURI + } return "http://localhost:8080/callback" } diff --git a/internal/adapters/oauth/server.go b/internal/adapters/oauth/server.go index 23a6887..f69f0e8 100644 --- a/internal/adapters/oauth/server.go +++ b/internal/adapters/oauth/server.go @@ -8,6 +8,7 @@ import ( "fmt" "net" "net/http" + "regexp" "strings" "sync" "time" @@ -26,6 +27,13 @@ type CallbackServer struct { once sync.Once mu sync.RWMutex state string + + // ipLiteral makes the server advertise http://127.0.0.1:/callback + // and bind IPv4 loopback only, so what it listens on is exactly what it + // tells the authorization server. The default advertises "localhost", + // which Nylas hosted auth has registered and which may resolve to either + // loopback family. + ipLiteral bool } // NewCallbackServer creates a new callback server. @@ -37,6 +45,16 @@ func NewCallbackServer(port int) *CallbackServer { } } +// NewLoopbackIPCallbackServer creates a callback server that binds and +// advertises the IPv4 loopback literal 127.0.0.1 (RFC 8252 section 7.3 +// recommends the literal over "localhost", whose resolution the client does +// not control). +func NewLoopbackIPCallbackServer(port int) *CallbackServer { + server := NewCallbackServer(port) + server.ipLiteral = true + return server +} + // Start starts the callback server. func (s *CallbackServer) Start() error { s.once = sync.Once{} @@ -59,7 +77,7 @@ func (s *CallbackServer) Start() error { // redirect URI, which can resolve to either IPv4 or IPv6 loopback // depending on host configuration. Listen on both loopback families when // available without accepting LAN traffic. - listeners, port, err := listenLoopback(s.port) + listeners, port, err := listenLoopback(s.port, !s.ipLiteral) if err != nil { return fmt.Errorf("failed to start callback server: %w", err) } @@ -74,7 +92,7 @@ func (s *CallbackServer) Start() error { return nil } -func listenLoopback(port int) ([]net.Listener, int, error) { +func listenLoopback(port int, includeIPv6 bool) ([]net.Listener, int, error) { ipv4, err := net.Listen("tcp4", fmt.Sprintf("127.0.0.1:%d", port)) if err != nil { return nil, 0, err @@ -91,6 +109,9 @@ func listenLoopback(port int) ([]net.Listener, int, error) { } listeners := []net.Listener{ipv4} + if !includeIPv6 { + return listeners, actualPort, nil + } ipv6, err := net.Listen("tcp6", fmt.Sprintf("[::1]:%d", actualPort)) if err != nil { if !isIPv6LoopbackUnavailable(err) { @@ -130,6 +151,13 @@ func (s *CallbackServer) Stop() error { return nil } +// SetExpectedState sets the state a callback must carry to be accepted. Call +// it before the browser opens the authorization URL, so a fast redirect is +// never checked against an unset state. +func (s *CallbackServer) SetExpectedState(state string) { + s.setExpectedState(state) +} + // WaitForCallback waits for the OAuth callback and returns the auth code. func (s *CallbackServer) WaitForCallback(ctx context.Context, expectedState string) (string, error) { s.setExpectedState(expectedState) @@ -146,15 +174,28 @@ func (s *CallbackServer) WaitForCallback(ctx context.Context, expectedState stri // GetRedirectURI returns the redirect URI for OAuth. func (s *CallbackServer) GetRedirectURI() string { + if s.ipLiteral { + return fmt.Sprintf("http://127.0.0.1:%d/callback", s.port) + } return fmt.Sprintf("http://localhost:%d/callback", s.port) } func (s *CallbackServer) handleCallback(w http.ResponseWriter, r *http.Request) { - code := r.URL.Query().Get("code") + query := r.URL.Query() + + // State first: only the redirect answering this login's request may end + // the wait. Anything else — a stale tab, another page probing the fixed + // port — is refused without touching the login in progress. + if !s.matchesExpectedState(query.Get("state")) { + http.Error(w, "Authentication failed: invalid OAuth state", http.StatusBadRequest) + return + } + + code := query.Get("code") if code == "" { - errMsg := r.URL.Query().Get("error") - if errMsg == "" { - errMsg = "no authorization code received" + errMsg := "no authorization code received" + if raw := query.Get("error"); raw != "" { + errMsg = sanitizeOAuthErrorCode(raw) } s.once.Do(func() { s.errChan <- fmt.Errorf("%w: %s", domain.ErrAuthFailed, errMsg) @@ -163,14 +204,6 @@ func (s *CallbackServer) handleCallback(w http.ResponseWriter, r *http.Request) return } - if !s.matchesExpectedState(r.URL.Query().Get("state")) { - s.once.Do(func() { - s.errChan <- fmt.Errorf("%w: invalid OAuth state", domain.ErrAuthFailed) - }) - http.Error(w, "Authentication failed: invalid OAuth state", http.StatusBadRequest) - return - } - s.once.Do(func() { s.codeChan <- code }) @@ -184,7 +217,7 @@ func (s *CallbackServer) handleCallback(w http.ResponseWriter, r *http.Request)