diff --git a/.env.example b/.env.example index 2129a74..b31a5b9 100644 --- a/.env.example +++ b/.env.example @@ -42,3 +42,14 @@ AUTH_COOKIE_SECURE=false # Trust the last X-Forwarded-For entry for login rate limiting, only behind a # reverse proxy you control that appends the peer address to the header. AUTH_TRUST_PROXY=false +# OpenID Connect login (see docs/AUTHENTICATION.md). Leave AUTH_OIDC_ISSUER empty to disable it. +# AUTH_PUBLIC_URL is required when it is set. +AUTH_OIDC_ISSUER= +AUTH_OIDC_CLIENT_ID= +AUTH_OIDC_CLIENT_SECRET= +# AUTH_OIDC_SCOPES=openid profile email +# AUTH_OIDC_GROUPS_CLAIM=groups +# AUTH_OIDC_USERNAME_CLAIM=preferred_username +# AUTH_OIDC_USER_PROVISIONING=true +# AUTH_OIDC_TEAM_SYNC=true +# AUTH_OIDC_BUTTON_LABEL=Single Sign-On diff --git a/README.md b/README.md index 1557927..ec59dab 100644 --- a/README.md +++ b/README.md @@ -246,7 +246,7 @@ npm run dev - [🚀 Installation Guide](./docs/INSTALLATION.md) - Complete installation instructions - [⚙️ Configuration Guide](./docs/CONFIGURATION.md) - Environment variables and settings - [🔧 Development Guide](./docs/DEVELOPMENT.md) - Set up development environment -- [🔐 Authentication](./docs/AUTHENTICATION.md) - Users, teams, permissions and API keys +- [🔐 Authentication](./docs/AUTHENTICATION.md) - Users, teams, permissions, API keys and single sign-on ### User Guides - [📖 User Guide](./docs/USER_GUIDE.md) - How to use Tracker diff --git a/cmd/serv.go b/cmd/serv.go index 95ab51d..c2ccc7b 100644 --- a/cmd/serv.go +++ b/cmd/serv.go @@ -22,6 +22,7 @@ import ( lock "github.com/bananaops/tracker/generated/proto/lock/v1alpha1" "github.com/bananaops/tracker/internal/auth" "github.com/bananaops/tracker/internal/auth/identity" + "github.com/bananaops/tracker/internal/auth/sso" store "github.com/bananaops/tracker/internal/stores" "github.com/bananaops/tracker/server" "github.com/go-openapi/runtime/middleware" @@ -149,6 +150,25 @@ var serv = &cobra.Command{ // Cookie based auth endpoints (login, logout, password change) server.NewAuthHTTP(userStore, sessions, authCfg).Register(mux) + // OpenID Connect login, only when AUTH_OIDC_ISSUER is set. Discovery is + // lazy: an unreachable identity provider must not keep Tracker from + // starting, the local admin account stays the way in. + if authCfg.OIDC.Enabled() { + codec, err := sso.NewTransactionCodec(sessionSecret) + if err != nil { + log.Fatalf("cannot create the OIDC transaction codec: %v", err) + } + provider := sso.NewOIDCProvider(authCfg.OIDC, authCfg.OIDCRedirectURL()) + go func() { + // The error is not logged: go-oidc embeds the raw response body. + if err := provider.Discover(context.Background()); err != nil { + slog.Warn("OIDC discovery failed at startup, it is retried on the next login", "issuer", authCfg.OIDC.Issuer, "reason", "provider_unavailable") + } + }() + server.NewOIDCHTTP(userStore, teamStore, sessions, provider, codec, authCfg).Register(mux) + slog.Info("OIDC login enabled", "oidc", authCfg.OIDC, "redirect_uri", authCfg.OIDCRedirectURL()) + } + // Register Homer proxy endpoint server.RegisterHomerHandler(mux, os.Getenv("HOMER_URL")) diff --git a/docs/AUTHENTICATION.md b/docs/AUTHENTICATION.md index 1931512..385b643 100644 --- a/docs/AUTHENTICATION.md +++ b/docs/AUTHENTICATION.md @@ -101,6 +101,8 @@ early. | `POST /api/v1alpha1/auth/password` | Body `{"currentPassword","newPassword"}`. Requires a session. | | `GET /api/v1alpha1/auth/me` | Identity, teams and effective permissions of the caller. Public. | | `GET /api/v1alpha1/auth/config` | Login options and anonymous permissions. Public. | +| `GET /api/v1alpha1/auth/oidc/login` | Starts an OpenID Connect login and redirects to the identity provider. Only when OIDC is enabled, `404` otherwise. See [Single Sign-On](#single-sign-on-openid-connect). | +| `GET /api/v1alpha1/auth/oidc/callback` | Redirect URI of the identity provider. Only when OIDC is enabled, `404` otherwise. | ### Browser cross-site requests @@ -119,11 +121,300 @@ any password work. API keys and `Authorization: Bearer` tokens are explicit credentials, not ambient ones, and are never dropped. Requests carrying neither header, which is every non browser client, are unaffected. +## Single Sign-On (OpenID Connect) + +Tracker can sign users in through any OpenID Connect identity provider (IdP): +Keycloak, Microsoft Entra ID, Google, GitLab, Dex and others. It uses the +Authorization Code flow with PKCE (`S256`), a random `state` and a `nonce`. +The `id_token` is verified for signature, issuer, audience, expiry and nonce. +Only the `id_token` is read: Tracker never calls the userinfo endpoint, so +every claim it needs must be in the `id_token`. + +The local `admin` account keeps working next to SSO and is the way back in +when the IdP is misconfigured or down. SSO is off unless `AUTH_OIDC_ISSUER` is +set. When it is on, `GET /api/v1alpha1/auth/config` reports `oidcEnabled` and +`oidcButtonLabel`. The Single Sign-On button of the login page ships with the +web PR #201: until it is merged, start a login by opening +`/api/v1alpha1/auth/oidc/login` directly. + +### Configuration + +| Variable | Default | Description | +|----------|---------|-------------| +| `AUTH_OIDC_ISSUER` | - | Issuer URL of the IdP. Enables SSO when set. Absolute `https` URL without query or fragment (`http` only for `localhost` and loopback IPs). It must match the `iss` of the tokens exactly, trailing slash included. | +| `AUTH_OIDC_CLIENT_ID` | - | Client ID. Required when the issuer is set. | +| `AUTH_OIDC_CLIENT_SECRET` | - | Client secret. Required when the issuer is set (confidential client). Never logged. | +| `AUTH_OIDC_SCOPES` | `openid profile email` | Requested scopes, separated by spaces or commas. `openid` is always added first. | +| `AUTH_OIDC_GROUPS_CLAIM` | `groups` | Name of the `id_token` claim carrying the groups: an array of strings or a single string. | +| `AUTH_OIDC_USERNAME_CLAIM` | `preferred_username` | Claim used as the Tracker username. When it is missing or not a valid username, `email` is tried. | +| `AUTH_OIDC_USER_PROVISIONING` | `true` | Create the Tracker account at first login. When `false`, only users already known to Tracker can sign in. | +| `AUTH_OIDC_TEAM_SYNC` | `true` | Synchronize team membership from the groups claim at each login. | +| `AUTH_OIDC_BUTTON_LABEL` | `Single Sign-On` | Label of the login button, at most 64 characters. | +| `AUTH_PUBLIC_URL` | - | Required with OIDC. `scheme://host[:port]` without path: it is the base of the redirect URI. | + +Startup fails with a message naming the variable when the configuration is +invalid: an issuer that is not `https` (outside loopback), a missing client ID +or client secret, a client ID or secret without an issuer, a missing or +path-bearing `AUTH_PUBLIC_URL`, a boolean that is not `true` or `false`, or a +button label over 64 characters. Pass the client secret through a secret +manager or a Kubernetes `Secret`, never in an image or a committed file. + +Register this redirect URI with the IdP: + +``` +/api/v1alpha1/auth/oidc/callback +``` + +The redirect URI is compared exactly by most IdPs: scheme, host, port and +path must match `AUTH_PUBLIC_URL`. The transaction cookie is encrypted with a +key derived from the session secret, so replicas must share `AUTH_SESSION_SECRET` +(or the persisted secret in MongoDB). + +Discovery (`/.well-known/openid-configuration`) is lazy. Tracker starts, and +the local login works, even when the IdP is unreachable. Discovery is tried +in the background at startup and again on the next SSO login, at most every +5 seconds after a failure. While the IdP is down, SSO logins end on +`oidc_unavailable` and users sign in with `admin` to keep managing Tracker. + +### Endpoints + +| Endpoint | Description | +|----------|-------------| +| `GET /api/v1alpha1/auth/oidc/login?redirect=/path` | Starts the login. `redirect` is where the user lands afterwards; only a local absolute path is accepted, anything else becomes `/`. | +| `GET /api/v1alpha1/auth/oidc/callback` | Receives the IdP response, creates the session and redirects. | + +Both answer `404` when `AUTH_OIDC_ISSUER` is not set. + +### Accounts + +- An OIDC account is identified by the pair `(issuer, subject)`. It is never + linked to an existing account by username or email, so a local account, or + an account of another IdP, cannot be taken over by claiming its name. +- At first login the account is created with the username taken from + `AUTH_OIDC_USERNAME_CLAIM` (then `email`), and no team. When the username + is already taken, a suffix is added: `alice`, then `alice-2`, `alice-3`... +- The username is frozen at creation. Email and display name are refreshed at + every login. The display name comes from `name`, then `given_name` and + `family_name`, then the username. +- OIDC accounts have no Tracker password: `POST /api/v1alpha1/auth/password` + answers `400` and administrators cannot set one. +- An administrator can disable an OIDC account in Tracker. Its next login is + refused with a `403` page. +- With `AUTH_OIDC_USER_PROVISIONING=false`, a user who is not known yet gets a + `403` page. Only accounts already known by their `(issuer, subject)` pair, that +is created by an earlier login while provisioning was on, can sign in. +- Changing `AUTH_OIDC_ISSUER` to another value creates new accounts: the + identity is bound to the issuer. Keep the issuer stable. + +### Teams + +With `AUTH_OIDC_TEAM_SYNC=true`, a team lists its OIDC groups in `oidcGroups` +(see [Teams](#teams)). At each login: + +- The user joins every team having at least one group in common with the + groups claim, and leaves every other team that has `oidcGroups`. +- Comparison is exact and case sensitive. +- Teams without `oidcGroups` are never touched: their members are managed by + hand in Tracker. A manual membership in a team that has `oidcGroups` is + overwritten at the next login. +- `Administrators` is synchronized only if it has `oidcGroups`. The last + enabled member of `Administrators` is never removed by a sync (the server + logs an error when it keeps them). +- A group claim that is present but empty means "no group": the user leaves + all mapped teams. +- If at least one team has `oidcGroups` and the claim is **absent** from the + `id_token`, the result depends on the user. A user who already holds a + membership in a team that has `oidcGroups` is refused with `oidc_failed` and + nothing is written (no account creation, no profile refresh, no membership + change), so a broken IdP mapper cannot silently strip anybody of their + rights. A user who holds no such membership, a first login included, is + treated as having no group: the login succeeds, no mapped team is added and + the server logs a `WARN` (`auth.oidc.sync`, reason + `groups_claim_missing_accepted`, with the claim name and the username). +- A claim that is neither a string nor an array (object, number, boolean, + `null`) is treated exactly like an absent claim, and the server logs a `WARN` + with reason `groups_claim_unexpected_type`. In an array, entries that are not + strings are ignored. + +> **Recommended: configure the IdP so the groups claim is always emitted in +> the `id_token`, as an empty array for users who have no group.** Many IdPs +> omit the claim for such users. That is harmless for a user without a mapped +> membership, but a user who already holds one is refused as soon as the claim +> goes missing, instead of being removed from the mapped teams. Check the +> claim by decoding a test `id_token`, and see the IdP recipes below. + +Without team mapping (`AUTH_OIDC_TEAM_SYNC=false`, or no team has +`oidcGroups`), no groups claim is needed and teams are managed by hand. + +### Security notes + +- The login transaction (state, nonce, PKCE verifier, redirect) lives in the + encrypted `tracker_oidc` cookie: `HttpOnly`, valid 10 minutes, `Path` set to + the callback route and `Secure` under the same rule as the session cookie. + It is `SameSite=Lax`, not `Strict`, because the callback is a cross-site + top level navigation coming from the IdP, and a `Strict` cookie would not be + sent. It is deleted on every callback. +- There is one login in progress per browser: starting a second login, for + example in another tab, replaces the first transaction and the first tab + ends on `oidc_state`. +- The Tracker session is independent of the IdP session. Disabling a user in + the IdP does not revoke their existing Tracker session before + `AUTH_SESSION_TTL` expires: disable the account in Tracker too, which + invalidates its sessions. Logging out of Tracker does not log out of the IdP. +- Values received from the IdP are never reflected: error codes below are + constants, and the IdP error text only goes to the server logs. + +### Errors + +A failed login redirects to `/login?error=`: + +| Code | Meaning | +|------|---------| +| `oidc_denied` | The IdP returned an error (user cancelled, access denied, client not allowed). | +| `oidc_state` | The login transaction is missing, expired (10 minutes), unreadable or does not match the `state`. | +| `oidc_failed` | Code exchange failed (including a token endpoint that is unreachable or answers with an error), `id_token` verification failed, the token has no usable username, the groups claim is missing for a user who holds a mapped team membership, or an internal error occurred. | +| `oidc_unavailable` | Discovery of the IdP failed: at the start of the login, or at the callback when discovery had not succeeded yet. | + +Two refusals are shown on a `403` page instead: the account is not registered +in Tracker (provisioning disabled) and the Tracker account is disabled. + +The web login page does not display the `?error=` code yet: a follow-up will +add it. Until then, read the reason in the server logs, where every failure is +an `auth.login` entry with `method=oidc` and a `reason` (`state_mismatch`, +`transaction_missing`, `id_token_verification_failed`, `not_provisioned`, +`user_disabled`, `provider_unavailable`...). Team sync problems are logged as +`auth.oidc.sync` (`groups_claim_missing`, and the warnings +`groups_claim_missing_accepted` and `groups_claim_unexpected_type`). Secrets, codes, tokens and cookie +values are never logged. + +### Identity provider recipes + +In every recipe, use the redirect URI above, keep a confidential client, and +set `AUTH_PUBLIC_URL` to the URL users type in their browser. + +#### Keycloak + +1. In the realm, create a client of type OpenID Connect. Enable Client + authentication (confidential) and the Standard flow only. +2. In Advanced settings, set the PKCE Code Challenge Method to `S256`. +3. Add the redirect URI in Valid redirect URIs. Copy the client secret from + the Credentials tab. +4. Add a mapper of type Group Membership to the client's dedicated scope, with + Token Claim Name `groups` and Add to ID token enabled. Keycloak puts the + claim in the token from the user's groups; check the result for a user who + has no group (Client scopes, Evaluate tab, generated ID token) and adapt the + mapper if the claim is missing. Full group path is your choice: when it is + on, values look like `/parent/child`, and that is what `oidcGroups` must + contain. + +```bash +AUTH_OIDC_ISSUER=https://keycloak.example.com/realms/ +AUTH_OIDC_CLIENT_ID=tracker +AUTH_OIDC_CLIENT_SECRET= +``` + +#### Microsoft Entra ID + +1. In App registrations, create an application. Add the Web platform with the + redirect URI, and create a client secret in Certificates and secrets. +2. In Token configuration, choose Add groups claim. Pick Security groups, or + Groups assigned to the application for large tenants. Group values are + Object IDs (GUIDs): put those in `oidcGroups`. +3. Groups overage: an `id_token` (JWT) carries at most 200 groups. Above that, + Entra removes the `groups` claim and sends `_claim_names` instead, which + Tracker does not follow. Users in that case are treated as having no group, + or refused with `oidc_failed` if they already hold a mapped membership. + Prevent it with Groups assigned to the application, which restricts the + claim to the groups assigned to the app (Enterprise applications, Users and + groups), or with app roles: define roles, assign them to groups, and set + `AUTH_OIDC_GROUPS_CLAIM=roles`. +4. Whether the claim appears for users with no group depends on the tenant + setup: check the `id_token` of such a user, and check the Microsoft + documentation if it is absent. +5. `preferred_username` is the user principal name (UPN). + +```bash +AUTH_OIDC_ISSUER=https://login.microsoftonline.com//v2.0 +AUTH_OIDC_CLIENT_ID= +AUTH_OIDC_CLIENT_SECRET= +``` + +#### Google + +Create an OAuth client of type Web application in the Google Cloud console and +add the redirect URI. Google issues no groups claim in its `id_token`, so team +mapping cannot work: set `AUTH_OIDC_TEAM_SYNC=false` and manage team +membership by hand. Use the email as username. To restrict access to your +organization, set the OAuth consent screen user type to Internal. Google adds +an `hd` claim for Workspace accounts, but Tracker does not check it. + +```bash +AUTH_OIDC_ISSUER=https://accounts.google.com +AUTH_OIDC_CLIENT_ID=.apps.googleusercontent.com +AUTH_OIDC_CLIENT_SECRET= +AUTH_OIDC_USERNAME_CLAIM=email +AUTH_OIDC_TEAM_SYNC=false +``` + +#### GitLab + +Create an application (instance, group or user level), confidential, with the +redirect URI and the scopes `openid`, `profile` and `email`. Per the GitLab +documentation, the `id_token` carries the `groups_direct` claim (direct group +memberships), while `groups` is only served by the userinfo endpoint, which +Tracker does not call. Use `groups_direct`, and check on your GitLab version +that it is present in a test `id_token`. Check the GitLab documentation for +the claim format. + +```bash +AUTH_OIDC_ISSUER=https://gitlab.com # or the URL of your instance +AUTH_OIDC_CLIENT_ID= +AUTH_OIDC_CLIENT_SECRET= +AUTH_OIDC_GROUPS_CLAIM=groups_direct +``` + +#### Dex + +Add a static client and request the `groups` scope, which makes Dex put the +user's groups in the `id_token`. Which groups Dex knows depends on the +connector (LDAP, GitHub, GitLab, OIDC...): check the connector documentation, +including what is emitted for a user without groups. + +```yaml +staticClients: + - id: tracker + name: Tracker + secret: + redirectURIs: + - https://tracker.example.com/api/v1alpha1/auth/oidc/callback +``` + +```bash +AUTH_OIDC_ISSUER=https://dex.example.com +AUTH_OIDC_CLIENT_ID=tracker +AUTH_OIDC_CLIENT_SECRET= +AUTH_OIDC_SCOPES=openid profile email groups +``` + +### Troubleshooting + +| Symptom | Cause and fix | +|---------|---------------| +| `oidc_failed`, log `id_token_verification_failed`, issuer mismatch | `AUTH_OIDC_ISSUER` differs from the `iss` claim, often by a trailing slash. Copy the `issuer` from the IdP discovery document. | +| The IdP shows a `redirect_uri` error | The registered URI is not exactly `/api/v1alpha1/auth/oidc/callback`. Check scheme, host, port. | +| `oidc_failed`, log `id_token_verification_failed`, token expired | Clock skew between Tracker and the IdP. Fix NTP on the hosts. | +| `oidc_failed`, log `groups_claim_missing` | A user who already holds a mapped team membership signed in with an `id_token` that has no claim named `AUTH_OIDC_GROUPS_CLAIM`. Nothing was changed. Decode a test `id_token`, check the claim name and that the mapper adds it to the ID token (not only the access token). Emitting an empty array for users without groups is recommended. | +| Users land in the wrong teams | `oidcGroups` values must equal the claim values exactly (Object IDs on Entra, `/parent/child` with Keycloak full paths). | +| Loop back to login with `oidc_state` | The `tracker_oidc` cookie was not returned: `AUTH_PUBLIC_URL` is `http` while the site is served over `https` (or the reverse), the host used in the browser differs from the one in `AUTH_PUBLIC_URL`, a proxy strips cookies, or two tabs started a login. Retry with a single tab. | +| `oidc_unavailable` | Discovery of the IdP failed (network, DNS, TLS trust). An unreachable token endpoint ends on `oidc_failed` instead. Sign in with `admin`, fix the network, and retry: discovery is retried on the next login. | +| `403` "not registered in Tracker" | `AUTH_OIDC_USER_PROVISIONING=false` and the user has never signed in. | + ## Teams A team carries a list of permissions, an optional list of catalog services (empty means every service; per-service filtering is enforced in a later -release) and optional OIDC group names (used once OIDC lands). Users belong +release) and optional OIDC group names (see +[Single Sign-On](#single-sign-on-openid-connect)). Users belong to any number of teams and get the union of their rights. The built-in `Administrators` team cannot be renamed, deleted or stripped of permissions. @@ -187,6 +478,9 @@ decisions, with `principal` in `anonymous`, `user`, `apikey` and `result` in `allowed`, `unauthenticated`, `denied`. `tracker_auth_logins_total{method,result}` counts login attempts, with -`method` in `local` (`oidc` once it lands) and `result` in `success`, -`failure`, `rate_limited`. Malformed bodies, cross-site refusals and internal -errors are not login attempts and are not counted. +`method` in `local`, `oidc` and `result` in `success`, `failure`, +`rate_limited`. Malformed bodies, cross-site refusals and internal errors are +not login attempts and are not counted. For `oidc`, a callback is counted when +it carries a valid login transaction (a callback with a missing or invalid +transaction is not), and so is a login start that fails +because the identity provider is unreachable. diff --git a/docs/CONFIGURATION.md b/docs/CONFIGURATION.md index 2023cbb..f812844 100644 --- a/docs/CONFIGURATION.md +++ b/docs/CONFIGURATION.md @@ -66,13 +66,22 @@ BUY_ME_COFFEE_URL=https://www.buymeacoffee.com/yourname | `AUTH_ADMIN_PASSWORD` | generated | Password of the initial `admin` account. Only used when no user exists yet. When unset, a random password is printed once in the logs. | | `AUTH_SESSION_SECRET` | persisted in MongoDB | Base64 secret (32 bytes minimum) signing session cookies. Set it explicitly when running several replicas without a shared database secret. | | `AUTH_SESSION_TTL` | `12h` | Session lifetime. | -| `AUTH_PUBLIC_URL` | - | Public URL of the UI. An `https` URL makes cookies `Secure`. | +| `AUTH_PUBLIC_URL` | - | Public URL of the UI. An `https` URL makes cookies `Secure`. Required with OIDC, where it is the base of the redirect URI (`scheme://host[:port]`, no path). | | `AUTH_COOKIE_SECURE` | `false` | Force the `Secure` flag on cookies. | | `AUTH_TRUST_PROXY` | `false` | Use the last entry of `X-Forwarded-For` as client IP for login rate limiting, and `X-Forwarded-Proto` to decide the request scheme. Only enable it behind a reverse proxy that appends the peer address to the header. | +| `AUTH_OIDC_ISSUER` | - | OpenID Connect issuer URL. Setting it enables single sign-on. `https` only (`http` for loopback). Must match the token `iss` exactly. | +| `AUTH_OIDC_CLIENT_ID` | - | OIDC client ID. Required with an issuer. | +| `AUTH_OIDC_CLIENT_SECRET` | - | OIDC client secret. Required with an issuer. Keep it in a secret store. | +| `AUTH_OIDC_SCOPES` | `openid profile email` | Requested scopes, space or comma separated. | +| `AUTH_OIDC_GROUPS_CLAIM` | `groups` | `id_token` claim carrying the groups used for team mapping. | +| `AUTH_OIDC_USERNAME_CLAIM` | `preferred_username` | `id_token` claim used as username, `email` as fallback. | +| `AUTH_OIDC_USER_PROVISIONING` | `true` | Create accounts at first OIDC login. | +| `AUTH_OIDC_TEAM_SYNC` | `true` | Synchronize teams from the groups claim at each login. | +| `AUTH_OIDC_BUTTON_LABEL` | `Single Sign-On` | Label of the login button (64 characters max). | When `AUTH_ANONYMOUS_PERMISSIONS` is set, its value is used as is, even when empty. When it is unset, the default is the read-only set `event:read,catalog:read,lock:read,links:read` if `DEMO_MODE=true`, otherwise every permission except `access:manage` (transitional default, with a startup warning). -See [AUTHENTICATION.md](AUTHENTICATION.md) for permissions, teams and API keys. +See [AUTHENTICATION.md](AUTHENTICATION.md) for permissions, teams and API keys, and [Single Sign-On](AUTHENTICATION.md#single-sign-on-openid-connect) for the OpenID Connect setup, redirect URI and identity provider recipes. **Example:** ```bash @@ -81,6 +90,15 @@ AUTH_ADMIN_PASSWORD=change-me-at-first-login AUTH_PUBLIC_URL=https://tracker.example.com ``` +**Example with OpenID Connect:** +```bash +AUTH_PUBLIC_URL=https://tracker.example.com +AUTH_OIDC_ISSUER=https://keycloak.example.com/realms/main +AUTH_OIDC_CLIENT_ID=tracker +AUTH_OIDC_CLIENT_SECRET= +# Redirect URI to register: https://tracker.example.com/api/v1alpha1/auth/oidc/callback +``` + ### Slack Integration | Variable | Default | Description | diff --git a/docs/index.md b/docs/index.md index 7057f1c..40d309d 100644 --- a/docs/index.md +++ b/docs/index.md @@ -9,7 +9,7 @@ Tracker est une API de gestion d'événements, de catalogues et de verrous const ### 📚 Documentation générale - [README](./README.md) - Vue d'ensemble et architecture - [Spécification API](./api-specification.md) - OpenAPI et Protobuf -- [Authentication](./AUTHENTICATION.md) - Users, teams, permissions and API keys +- [Authentication](./AUTHENTICATION.md) - Users, teams, permissions, API keys and single sign-on ### 🔧 APIs par service - [Events API](./events.md) - Gestion des événements diff --git a/go.mod b/go.mod index 3ea8634..8388b48 100644 --- a/go.mod +++ b/go.mod @@ -3,9 +3,11 @@ module github.com/bananaops/tracker go 1.26.1 require ( + github.com/coreos/go-oidc/v3 v3.21.0 github.com/go-openapi/runtime v0.29.5 github.com/golang-jwt/jwt/v5 v5.3.0 golang.org/x/crypto v0.56.0 + golang.org/x/oauth2 v0.37.0 google.golang.org/grpc v1.83.2 google.golang.org/protobuf v1.36.11 gopkg.in/yaml.v3 v3.0.1 @@ -15,6 +17,7 @@ require ( github.com/beorn7/perks v1.0.1 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/davecgh/go-spew v1.1.1 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-openapi/analysis v0.25.3 // indirect github.com/go-openapi/errors v0.22.8 // indirect github.com/go-openapi/jsonpointer v0.23.2 // indirect diff --git a/go.sum b/go.sum index 552d420..a2c0e4b 100644 --- a/go.sum +++ b/go.sum @@ -2,12 +2,16 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/coreos/go-oidc/v3 v3.21.0 h1:wZo4Q9Pum8dYEj0eMUPrqR+kvuGkeUplbLpNCkBqoWM= +github.com/coreos/go-oidc/v3 v3.21.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4= github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/envoyproxy/protoc-gen-validate v1.3.3 h1:MVQghNeW+LZcmXe7SY1V36Z+WFMDjpqGAGacLe2T0ds= github.com/envoyproxy/protoc-gen-validate v1.3.3/go.mod h1:TsndJ/ngyIdQRhMcVVGDDHINPLWB7C82oDArY51KfB0= +github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= +github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= @@ -144,6 +148,8 @@ golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/oauth2 v0.37.0 h1:JUlcxA8oAtauLfiH8FX2/FkAWHAdi0QtGCGc+hofE98= +golang.org/x/oauth2 v0.37.0/go.mod h1:IxwZNxUULJmpBFf9K/9NTMSIfZZuvuTy1gGxhigP/58= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= diff --git a/internal/auth/authz/authz.go b/internal/auth/authz/authz.go index 586810a..062f66c 100644 --- a/internal/auth/authz/authz.go +++ b/internal/auth/authz/authz.go @@ -24,7 +24,7 @@ var authRequests = prometheus.NewCounterVec( ) // AuthLogins counts login attempts. The method label names the authentication -// method (local for now, oidc once it lands) and the result label is one of +// method (local or oidc) and the result label is one of // LoginSuccess, LoginFailure or LoginRateLimited. Malformed bodies, cross-site // refusals and internal errors are not login attempts and are not counted. // It is exported so the login handler, which lives in the server package, can @@ -40,6 +40,7 @@ var AuthLogins = prometheus.NewCounterVec( // Values of the AuthLogins labels. const ( LoginMethodLocal = "local" + LoginMethodOIDC = "oidc" LoginSuccess = "success" LoginFailure = "failure" LoginRateLimited = "rate_limited" diff --git a/internal/auth/config.go b/internal/auth/config.go index 317d545..1a33fb7 100644 --- a/internal/auth/config.go +++ b/internal/auth/config.go @@ -2,9 +2,26 @@ package auth import ( "encoding/base64" + "errors" "fmt" + "log/slog" + "net" + "net/url" + "strconv" "strings" "time" + "unicode/utf8" +) + +const ( + OIDCLoginPath = "/api/v1alpha1/auth/oidc/login" + OIDCCallbackPath = "/api/v1alpha1/auth/oidc/callback" + + defaultOIDCScopes = "openid profile email" + defaultOIDCGroupsClaim = "groups" + defaultOIDCUsernameClaim = "preferred_username" + defaultOIDCButtonLabel = "Single Sign-On" + maxOIDCButtonLabelLength = 64 ) // Config is the authentication configuration read from the environment. @@ -22,6 +39,47 @@ type Config struct { CookieSecure bool TrustProxy bool DemoMode bool + // OIDC is the OpenID Connect login, disabled when Issuer is empty. + OIDC OIDCConfig +} + +// OIDCConfig is the OpenID Connect login configuration read from the +// AUTH_OIDC_* environment variables. +type OIDCConfig struct { + Issuer string + ClientID string + ClientSecret string + Scopes []string + GroupsClaim string + UsernameClaim string + UserProvisioning bool + TeamSync bool + ButtonLabel string +} + +// Enabled reports whether OIDC login is configured. +func (c OIDCConfig) Enabled() bool { + return c.Issuer != "" +} + +// LogValue implements slog.LogValuer, omitting ClientSecret from logs. +func (c OIDCConfig) LogValue() slog.Value { + return slog.GroupValue( + slog.String("issuer", c.Issuer), + slog.String("client_id", c.ClientID), + slog.Any("scopes", c.Scopes), + slog.String("groups_claim", c.GroupsClaim), + slog.String("username_claim", c.UsernameClaim), + slog.Bool("user_provisioning", c.UserProvisioning), + slog.Bool("team_sync", c.TeamSync), + slog.String("button_label", c.ButtonLabel), + ) +} + +// OIDCRedirectURL is the OIDC callback URL registered with the identity +// provider. +func (c Config) OIDCRedirectURL() string { + return c.PublicURL + OIDCCallbackPath } // LookupEnv has the signature of os.LookupEnv. @@ -96,5 +154,139 @@ func LoadConfig(lookup LookupEnv) (Config, error) { cfg.PublicURL = strings.TrimRight(get("AUTH_PUBLIC_URL"), "/") cfg.CookieSecure = strings.HasPrefix(cfg.PublicURL, "https://") || get("AUTH_COOKIE_SECURE") == "true" cfg.TrustProxy = get("AUTH_TRUST_PROXY") == "true" + + oidc, err := loadOIDCConfig(get, cfg.PublicURL) + if err != nil { + return Config{}, err + } + cfg.OIDC = oidc + + return cfg, nil +} + +// loadOIDCConfig reads the AUTH_OIDC_* variables. OIDC is off without an +// issuer. Error messages name the variable, never the client secret. +func loadOIDCConfig(get func(string) string, publicURL string) (OIDCConfig, error) { + cfg := OIDCConfig{ + Issuer: get("AUTH_OIDC_ISSUER"), + ClientID: get("AUTH_OIDC_CLIENT_ID"), + ClientSecret: get("AUTH_OIDC_CLIENT_SECRET"), + } + if cfg.Issuer == "" { + if cfg.ClientID != "" || cfg.ClientSecret != "" { + return OIDCConfig{}, errors.New("AUTH_OIDC_CLIENT_ID or AUTH_OIDC_CLIENT_SECRET is set but AUTH_OIDC_ISSUER is empty") + } + return OIDCConfig{}, nil + } + if err := validateOIDCIssuer(cfg.Issuer); err != nil { + return OIDCConfig{}, err + } + if cfg.ClientID == "" { + return OIDCConfig{}, errors.New("AUTH_OIDC_CLIENT_ID is required when AUTH_OIDC_ISSUER is set") + } + if cfg.ClientSecret == "" { + return OIDCConfig{}, errors.New("AUTH_OIDC_CLIENT_SECRET is required when AUTH_OIDC_ISSUER is set") + } + if err := validateOIDCPublicURL(publicURL); err != nil { + return OIDCConfig{}, err + } + cfg.Scopes = parseOIDCScopes(getOr(get, "AUTH_OIDC_SCOPES", defaultOIDCScopes)) + cfg.GroupsClaim = getOr(get, "AUTH_OIDC_GROUPS_CLAIM", defaultOIDCGroupsClaim) + cfg.UsernameClaim = getOr(get, "AUTH_OIDC_USERNAME_CLAIM", defaultOIDCUsernameClaim) + var err error + if cfg.UserProvisioning, err = boolOr(get, "AUTH_OIDC_USER_PROVISIONING", true); err != nil { + return OIDCConfig{}, err + } + if cfg.TeamSync, err = boolOr(get, "AUTH_OIDC_TEAM_SYNC", true); err != nil { + return OIDCConfig{}, err + } + cfg.ButtonLabel = getOr(get, "AUTH_OIDC_BUTTON_LABEL", defaultOIDCButtonLabel) + if utf8.RuneCountInString(cfg.ButtonLabel) > maxOIDCButtonLabelLength { + return OIDCConfig{}, fmt.Errorf("AUTH_OIDC_BUTTON_LABEL must be at most %d characters", maxOIDCButtonLabelLength) + } return cfg, nil } + +// getOr returns def when key is unset or blank. +func getOr(get func(string) string, key, def string) string { + if v := get(key); v != "" { + return v + } + return def +} + +// boolOr returns def when key is unset or blank, otherwise parses it as a +// bool. The error never repeats the value. +func boolOr(get func(string) string, key string, def bool) (bool, error) { + v := get(key) + if v == "" { + return def, nil + } + b, err := strconv.ParseBool(v) + if err != nil { + return false, fmt.Errorf("%s must be true or false", key) + } + return b, nil +} + +// parseOIDCScopes splits on spaces and commas, deduplicates while keeping +// order, and puts openid first. +func parseOIDCScopes(raw string) []string { + fields := strings.FieldsFunc(raw, func(r rune) bool { + return r == ' ' || r == ',' + }) + seen := make(map[string]bool, len(fields)+1) + scopes := make([]string, 0, len(fields)+1) + scopes = append(scopes, "openid") + seen["openid"] = true + for _, f := range fields { + if f == "" || seen[f] { + continue + } + seen[f] = true + scopes = append(scopes, f) + } + return scopes +} + +// validateOIDCIssuer requires an absolute https URL without query or +// fragment; http is accepted for loopback hosts only. +func validateOIDCIssuer(issuer string) error { + const msg = "AUTH_OIDC_ISSUER must be an absolute https URL without query or fragment (http is accepted for loopback hosts only)" + u, err := url.Parse(issuer) + if err != nil || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" { + return errors.New(msg) + } + switch u.Scheme { + case "https": + return nil + case "http": + if isLoopbackHost(u.Hostname()) { + return nil + } + } + return errors.New(msg) +} + +// isLoopbackHost reports whether h is localhost or a loopback IP. +func isLoopbackHost(h string) bool { + if h == "localhost" { + return true + } + ip := net.ParseIP(h) + return ip != nil && ip.IsLoopback() +} + +// validateOIDCPublicURL requires an absolute scheme://host[:port] URL +// without a path, since it is the base of the OIDC redirect URI. +func validateOIDCPublicURL(publicURL string) error { + if publicURL == "" { + return errors.New("AUTH_PUBLIC_URL is required when AUTH_OIDC_ISSUER is set: it is the base of the OIDC redirect URI") + } + const msg = "AUTH_PUBLIC_URL must be scheme://host[:port] without path when OIDC is enabled" + u, err := url.Parse(publicURL) + if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.Path != "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" { + return errors.New(msg) + } + return nil +} diff --git a/internal/auth/config_test.go b/internal/auth/config_test.go index 10f10a9..29c733e 100644 --- a/internal/auth/config_test.go +++ b/internal/auth/config_test.go @@ -3,6 +3,8 @@ package auth import ( "bytes" "encoding/base64" + "log/slog" + "strings" "testing" "time" @@ -66,3 +68,120 @@ func TestLoadConfigErrors(t *testing.T) { _, err = LoadConfig(envOf(map[string]string{"AUTH_ADMIN_PASSWORD": "short"})) assert.ErrorIs(t, err, ErrPasswordPolicy) } + +func oidcEnv(extra map[string]string) map[string]string { + env := map[string]string{ + "AUTH_PUBLIC_URL": "https://tracker.example.com", + "AUTH_OIDC_ISSUER": "https://idp.example.com/realms/main", + "AUTH_OIDC_CLIENT_ID": "tracker", + "AUTH_OIDC_CLIENT_SECRET": "s3cr3t-value-never-logged", + } + for k, v := range extra { + env[k] = v + } + return env +} + +func TestLoadConfigOIDCDisabledByDefault(t *testing.T) { + cfg, err := LoadConfig(envOf(map[string]string{})) + require.NoError(t, err) + assert.False(t, cfg.OIDC.Enabled()) + assert.Empty(t, cfg.OIDC.ButtonLabel) +} + +func TestLoadConfigOIDCDefaults(t *testing.T) { + cfg, err := LoadConfig(envOf(oidcEnv(nil))) + require.NoError(t, err) + o := cfg.OIDC + assert.True(t, o.Enabled()) + assert.Equal(t, "https://idp.example.com/realms/main", o.Issuer) + assert.Equal(t, "tracker", o.ClientID) + assert.Equal(t, "s3cr3t-value-never-logged", o.ClientSecret) + assert.Equal(t, []string{"openid", "profile", "email"}, o.Scopes) + assert.Equal(t, "groups", o.GroupsClaim) + assert.Equal(t, "preferred_username", o.UsernameClaim) + assert.True(t, o.UserProvisioning) + assert.True(t, o.TeamSync) + assert.Equal(t, "Single Sign-On", o.ButtonLabel) + assert.Equal(t, "https://tracker.example.com/api/v1alpha1/auth/oidc/callback", cfg.OIDCRedirectURL()) +} + +func TestLoadConfigOIDCExplicit(t *testing.T) { + cfg, err := LoadConfig(envOf(oidcEnv(map[string]string{ + "AUTH_PUBLIC_URL": "https://tracker.example.com/", + "AUTH_OIDC_ISSUER": "https://tenant.example.com/", + "AUTH_OIDC_SCOPES": "profile, groups email", + "AUTH_OIDC_GROUPS_CLAIM": "roles", + "AUTH_OIDC_USERNAME_CLAIM": "email", + "AUTH_OIDC_USER_PROVISIONING": "false", + "AUTH_OIDC_TEAM_SYNC": "false", + "AUTH_OIDC_BUTTON_LABEL": "Sign in with Okta", + }))) + require.NoError(t, err) + o := cfg.OIDC + assert.Equal(t, "https://tenant.example.com/", o.Issuer, "the trailing slash is significant for issuer matching") + assert.Equal(t, []string{"openid", "profile", "groups", "email"}, o.Scopes, "openid is added first, separators are commas or spaces") + assert.Equal(t, "roles", o.GroupsClaim) + assert.Equal(t, "email", o.UsernameClaim) + assert.False(t, o.UserProvisioning) + assert.False(t, o.TeamSync) + assert.Equal(t, "Sign in with Okta", o.ButtonLabel) + assert.Equal(t, "https://tracker.example.com/api/v1alpha1/auth/oidc/callback", cfg.OIDCRedirectURL()) +} + +func TestLoadConfigOIDCLoopbackHTTPIssuer(t *testing.T) { + for _, issuer := range []string{"http://127.0.0.1:5556/dex", "http://localhost:8081", "http://[::1]:5556"} { + _, err := LoadConfig(envOf(oidcEnv(map[string]string{"AUTH_OIDC_ISSUER": issuer}))) + assert.NoError(t, err, issuer) + } +} + +func TestLoadConfigOIDCInvalid(t *testing.T) { + cases := map[string]map[string]string{ + "http issuer off loopback": {"AUTH_OIDC_ISSUER": "http://idp.example.com"}, + "relative issuer": {"AUTH_OIDC_ISSUER": "idp.example.com"}, + "issuer with query": {"AUTH_OIDC_ISSUER": "https://idp.example.com?x=1"}, + "missing client id": {"AUTH_OIDC_CLIENT_ID": ""}, + "missing client secret": {"AUTH_OIDC_CLIENT_SECRET": ""}, + "missing public url": {"AUTH_PUBLIC_URL": ""}, + "public url with path": {"AUTH_PUBLIC_URL": "https://example.com/tracker"}, + "public url not absolute": {"AUTH_PUBLIC_URL": "tracker.example.com"}, + "bad provisioning bool": {"AUTH_OIDC_USER_PROVISIONING": "yes please"}, + "bad team sync bool": {"AUTH_OIDC_TEAM_SYNC": "maybe"}, + "label too long": {"AUTH_OIDC_BUTTON_LABEL": strings.Repeat("x", 65)}, + "empty groups claim name": {"AUTH_OIDC_GROUPS_CLAIM": " "}, // blank falls back to the default: must NOT error, see below + } + for name, extra := range cases { + t.Run(name, func(t *testing.T) { + _, err := LoadConfig(envOf(oidcEnv(extra))) + if name == "empty groups claim name" { + require.NoError(t, err) + return + } + require.Error(t, err) + assert.NotContains(t, err.Error(), "s3cr3t-value-never-logged", "the client secret never appears in errors") + }) + } +} + +func TestLoadConfigOIDCClientWithoutIssuer(t *testing.T) { + _, err := LoadConfig(envOf(map[string]string{"AUTH_OIDC_CLIENT_ID": "tracker", "AUTH_OIDC_CLIENT_SECRET": "s3cr3t-value-never-logged"})) + require.Error(t, err) + assert.Contains(t, err.Error(), "AUTH_OIDC_ISSUER") + assert.NotContains(t, err.Error(), "s3cr3t-value-never-logged") +} + +func TestLoadConfigOIDCMissingPublicURLMessage(t *testing.T) { + _, err := LoadConfig(envOf(oidcEnv(map[string]string{"AUTH_PUBLIC_URL": ""}))) + require.Error(t, err) + assert.Contains(t, err.Error(), "AUTH_PUBLIC_URL") +} + +func TestOIDCConfigLogValueHidesSecret(t *testing.T) { + cfg, err := LoadConfig(envOf(oidcEnv(nil))) + require.NoError(t, err) + var buf bytes.Buffer + slog.New(slog.NewJSONHandler(&buf, nil)).Info("oidc", "config", cfg.OIDC) + assert.Contains(t, buf.String(), "https://idp.example.com/realms/main") + assert.NotContains(t, buf.String(), "s3cr3t-value-never-logged") +} diff --git a/internal/auth/identity/mongo_testing_test.go b/internal/auth/identity/mongo_testing_test.go new file mode 100644 index 0000000..cb5e885 --- /dev/null +++ b/internal/auth/identity/mongo_testing_test.go @@ -0,0 +1,37 @@ +package identity + +import ( + "context" + "fmt" + "os" + "testing" + "time" + + store "github.com/bananaops/tracker/internal/stores" + "github.com/stretchr/testify/require" + "go.mongodb.org/mongo-driver/mongo" + "go.mongodb.org/mongo-driver/mongo/options" +) + +// mongoStores returns real stores on a throwaway database, or skips without MONGO_TEST_URI. +func mongoStores(t *testing.T) (*store.AuthUserStore, *store.AuthTeamStore) { + t.Helper() + uri := os.Getenv("MONGO_TEST_URI") + if uri == "" { + t.Skip("MONGO_TEST_URI not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + client, err := mongo.Connect(ctx, options.Client().ApplyURI(uri)) + require.NoError(t, err) + db := client.Database(fmt.Sprintf("tracker_test_%d", time.Now().UnixNano())) + t.Cleanup(func() { + c, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + _ = db.Drop(c) + _ = client.Disconnect(c) + }) + require.NoError(t, store.EnsureIndexes(ctx, db)) + return store.NewAuthUserStoreFromCollection(db.Collection("auth_users")), + store.NewAuthTeamStoreFromCollection(db.Collection("auth_teams")) +} diff --git a/internal/auth/identity/oidc.go b/internal/auth/identity/oidc.go new file mode 100644 index 0000000..c35b38e --- /dev/null +++ b/internal/auth/identity/oidc.go @@ -0,0 +1,302 @@ +package identity + +import ( + "context" + "errors" + "fmt" + "time" + + store "github.com/bananaops/tracker/internal/stores" + "go.mongodb.org/mongo-driver/bson/primitive" +) + +const ( + maxUsernameLength = 64 + maxUsernameAttempts = 100 +) + +// OIDCUserStore is the subset of store.AuthUserStore used to resolve OpenID Connect users. +type OIDCUserStore interface { + GetByOIDCIdentity(ctx context.Context, issuer, subject string) (*store.User, error) + Create(ctx context.Context, u *store.User) error + UpdateOIDCProfile(ctx context.Context, id primitive.ObjectID, email, displayName string, at time.Time) error +} + +// OIDCIdentity is what the identity provider asserts about a user. +type OIDCIdentity struct { + Issuer, Subject, Username, Email, DisplayName string +} + +var ( + ErrOIDCNotProvisioned = errors.New("oidc user is not provisioned") + ErrOIDCUserDisabled = errors.New("oidc user is disabled") + ErrOIDCNoUsername = errors.New("oidc identity carries no usable username") + ErrOIDCUsernameExhausted = errors.New("no free username for the oidc identity") + ErrOIDCInvalidIdentity = errors.New("oidc identity has no issuer or subject") + ErrOIDCNotOIDCUser = errors.New("user bound to the oidc identity is not an oidc account") +) + +// ResolveOIDCUser finds the user bound to (issuer, subject), or creates it +// when provisioning is on. It never binds an existing account by username or +// email. created reports a new account. +func ResolveOIDCUser(ctx context.Context, users OIDCUserStore, id OIDCIdentity, provisioning bool, now time.Time) (user *store.User, created bool, err error) { + if id.Issuer == "" || id.Subject == "" { + return nil, false, ErrOIDCInvalidIdentity + } + existing, err := users.GetByOIDCIdentity(ctx, id.Issuer, id.Subject) + if err == nil { + u, err := refreshOIDCUser(ctx, users, existing, id, now) + return u, false, err + } + if !errors.Is(err, store.ErrNotFound) { + return nil, false, fmt.Errorf("lookup oidc user: %w", err) + } + if !provisioning { + return nil, false, ErrOIDCNotProvisioned + } + if id.Username == "" { + return nil, false, ErrOIDCNoUsername + } + + for attempt := 1; attempt <= maxUsernameAttempts; attempt++ { + name := candidateUsername(id.Username, attempt) + displayName := id.DisplayName + if displayName == "" { + displayName = name + } + u := &store.User{ + Username: name, + Email: id.Email, + DisplayName: displayName, + Source: store.UserSourceOIDC, + OIDCIssuer: id.Issuer, + OIDCSubject: id.Subject, + Teams: []primitive.ObjectID{}, + LastLoginAt: &now, + } + err := users.Create(ctx, u) + if err == nil { + return u, true, nil + } + if !errors.Is(err, store.ErrAlreadyExists) { + return nil, false, fmt.Errorf("create oidc user: %w", err) + } + // Either a concurrent first login won the race on (issuer, subject), + // or the username is taken. + existing, lookupErr := users.GetByOIDCIdentity(ctx, id.Issuer, id.Subject) + if lookupErr == nil { + u, err := refreshOIDCUser(ctx, users, existing, id, now) + return u, false, err + } + if !errors.Is(lookupErr, store.ErrNotFound) { + return nil, false, fmt.Errorf("lookup oidc user: %w", lookupErr) + } + } + return nil, false, ErrOIDCUsernameExhausted +} + +func refreshOIDCUser(ctx context.Context, users OIDCUserStore, u *store.User, id OIDCIdentity, now time.Time) (*store.User, error) { + if u.Disabled { + return nil, ErrOIDCUserDisabled + } + if u.Source != store.UserSourceOIDC { + return nil, ErrOIDCNotOIDCUser + } + displayName := id.DisplayName + if displayName == "" { + displayName = u.DisplayName + } + if displayName == "" { + displayName = u.Username + } + if err := users.UpdateOIDCProfile(ctx, u.ID, id.Email, displayName, now); err != nil { + return nil, fmt.Errorf("update oidc profile: %w", err) + } + u.Email, u.DisplayName, u.LastLoginAt = id.Email, displayName, &now + return u, nil +} + +// candidateUsername returns base for the first attempt, then base suffixed +// with -attempt, truncated so the result fits the username length limit. The +// base is ASCII by construction. +func candidateUsername(base string, attempt int) string { + if attempt <= 1 { + return base + } + suffix := fmt.Sprintf("-%d", attempt) + if max := maxUsernameLength - len(suffix); len(base) > max { + base = base[:max] + } + return base + suffix +} + +// OIDCTeamStore is the subset of store.AuthTeamStore used to sync memberships. +type OIDCTeamStore interface { + ListWithOIDCGroups(ctx context.Context) ([]*store.Team, error) +} + +// OIDCMembershipStore is the subset of store.AuthUserStore used to sync memberships. +type OIDCMembershipStore interface { + SyncTeams(ctx context.Context, id primitive.ObjectID, add, remove []primitive.ObjectID) error + CountEnabledInTeam(ctx context.Context, teamID, excludeUser primitive.ObjectID) (int64, error) +} + +// TeamSyncResult names the teams changed by a sync, for logging. +type TeamSyncResult struct { + Added []string + Removed []string + // Kept lists teams the user should have left but kept: the last enabled + // administrator is never removed from Administrators. + Kept []string +} + +// ErrOIDCGroupsClaimMissing is returned when teams are mapped to OIDC groups, +// the identity provider sent no groups claim and the user already holds a +// mapped membership: nothing is granted or removed, so a broken mapper cannot +// strip a user of its rights. A user holding no mapped membership has nothing +// to lose and is treated as having no groups. +var ErrOIDCGroupsClaimMissing = errors.New("oidc groups claim is missing while the user holds mapped team memberships") + +// CheckOIDCGroupsClaim is the precondition of SyncOIDCTeams, exposed so a +// caller can refuse a login before writing anything. user is the account +// already bound to the identity, read-only, or nil for a new subject. It +// returns ErrOIDCGroupsClaimMissing when the claim was not sent, at least one +// team is mapped to OIDC groups and user holds one of the mapped teams. +func CheckOIDCGroupsClaim(ctx context.Context, teams OIDCTeamStore, user *store.User, groupsPresent bool) error { + if groupsPresent { + return nil + } + mapped, err := teams.ListWithOIDCGroups(ctx) + if err != nil { + return fmt.Errorf("list mapped teams: %w", err) + } + return checkGroupsClaim(mapped, user, groupsPresent) +} + +func checkGroupsClaim(mapped []*store.Team, user *store.User, groupsPresent bool) error { + if groupsPresent || user == nil { + return nil + } + held := make(map[primitive.ObjectID]bool, len(user.Teams)) + for _, id := range user.Teams { + held[id] = true + } + for _, t := range mapped { + if held[t.ID] { + return ErrOIDCGroupsClaimMissing + } + } + return nil +} + +// SyncOIDCTeams makes the user a member of every team whose OIDC groups +// intersect groups, and removes it from the other mapped teams. Teams without +// OIDC groups are left alone. Comparison is exact and case sensitive. +// groupsPresent tells whether the claim was sent at all: absent fails without +// any write when the user holds a mapped membership and otherwise means no +// group, like a present but empty claim. +// Both directions are targeted and idempotent over every mapped team, so a +// membership changed elsewhere since the user was loaded is corrected too; +// Added and Removed report only the changes relative to the loaded user. +func SyncOIDCTeams(ctx context.Context, users OIDCMembershipStore, teams OIDCTeamStore, user *store.User, groups []string, groupsPresent bool) (TeamSyncResult, error) { + res := TeamSyncResult{} + if user.Source != store.UserSourceOIDC { + return res, ErrOIDCNotOIDCUser + } + mapped, err := teams.ListWithOIDCGroups(ctx) + if err != nil { + return res, fmt.Errorf("list mapped teams: %w", err) + } + if err := checkGroupsClaim(mapped, user, groupsPresent); err != nil { + return res, err + } + if !groupsPresent { + // An absent claim says nothing about the groups: never add or remove. + return res, nil + } + if len(mapped) == 0 { + return res, nil + } + member := make(map[primitive.ObjectID]bool, len(user.Teams)) + for _, id := range user.Teams { + member[id] = true + } + + add, stale := planTeamSync(mapped, groups) + var remove []*store.Team + for _, t := range stale { + if t.Builtin && t.Name == store.AdministratorsTeamName && !user.Disabled { + others, err := users.CountEnabledInTeam(ctx, t.ID, user.ID) + if err != nil { + return res, fmt.Errorf("count enabled administrators: %w", err) + } + if others == 0 { + if member[t.ID] { + res.Kept = append(res.Kept, t.Name) + } + continue + } + } + remove = append(remove, t) + } + + if len(add) == 0 && len(remove) == 0 { + return res, nil + } + if err := users.SyncTeams(ctx, user.ID, teamIDs(add), teamIDs(remove)); err != nil { + return res, fmt.Errorf("sync user teams: %w", err) + } + + removed := map[primitive.ObjectID]bool{} + for _, t := range remove { + removed[t.ID] = true + if member[t.ID] { + res.Removed = append(res.Removed, t.Name) + } + } + next := make([]primitive.ObjectID, 0, len(user.Teams)+len(add)) + for _, id := range user.Teams { + if !removed[id] { + next = append(next, id) + } + } + for _, t := range add { + if !member[t.ID] { + next = append(next, t.ID) + res.Added = append(res.Added, t.Name) + } + } + user.Teams = next + return res, nil +} + +func teamMatches(t *store.Team, groups []string) bool { + for _, tg := range t.OIDCGroups { + for _, g := range groups { + if tg == g { + return true + } + } + } + return false +} + +func teamIDs(teams []*store.Team) []primitive.ObjectID { + ids := make([]primitive.ObjectID, 0, len(teams)) + for _, t := range teams { + ids = append(ids, t.ID) + } + return ids +} + +// planTeamSync is the pure decision behind SyncOIDCTeams. +func planTeamSync(mapped []*store.Team, groups []string) (match, stale []*store.Team) { + for _, t := range mapped { + if teamMatches(t, groups) { + match = append(match, t) + } else { + stale = append(stale, t) + } + } + return match, stale +} diff --git a/internal/auth/identity/oidc_test.go b/internal/auth/identity/oidc_test.go new file mode 100644 index 0000000..a8a184c --- /dev/null +++ b/internal/auth/identity/oidc_test.go @@ -0,0 +1,587 @@ +package identity + +import ( + "context" + "strings" + "sync" + "testing" + "time" + + store "github.com/bananaops/tracker/internal/stores" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.mongodb.org/mongo-driver/bson/primitive" +) + +const testIssuer = "https://idp" + +func oidcID(subject, username string) OIDCIdentity { + return OIDCIdentity{Issuer: testIssuer, Subject: subject, Username: username, Email: username + "@x.io", DisplayName: username} +} + +func TestResolveOIDCUserProvisions(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + id := OIDCIdentity{Issuer: testIssuer, Subject: "s1", Username: "alice", Email: "alice@x.io", DisplayName: "Alice"} + u, created, err := ResolveOIDCUser(ctx, users, id, true, now) + require.NoError(t, err) + assert.True(t, created) + assert.False(t, u.ID.IsZero()) + assert.Equal(t, store.UserSourceOIDC, u.Source) + assert.Equal(t, testIssuer, u.OIDCIssuer) + assert.Equal(t, "s1", u.OIDCSubject) + assert.Empty(t, u.PasswordHash) + assert.False(t, u.MustChangePassword) + assert.NotNil(t, u.Teams) + assert.Empty(t, u.Teams) + require.NotNil(t, u.LastLoginAt) + assert.False(t, u.Disabled) + + stored, err := users.GetByOIDCIdentity(ctx, testIssuer, "s1") + require.NoError(t, err) + assert.Equal(t, u.ID, stored.ID) + assert.Equal(t, "Alice", stored.DisplayName) +} + +func TestResolveOIDCUserReturnsExisting(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + first, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + require.NoError(t, err) + + again := OIDCIdentity{Issuer: testIssuer, Subject: "s1", Username: "renamed", Email: "new@x.io", DisplayName: "Alice New"} + u, created, err := ResolveOIDCUser(ctx, users, again, true, now.Add(time.Minute)) + require.NoError(t, err) + assert.False(t, created) + assert.Equal(t, first.ID, u.ID) + assert.Equal(t, "new@x.io", u.Email) + assert.Equal(t, "Alice New", u.DisplayName) + + stored, err := users.GetByID(ctx, first.ID) + require.NoError(t, err) + assert.Equal(t, "new@x.io", stored.Email) + assert.Equal(t, "Alice New", stored.DisplayName) + assert.Equal(t, "alice", stored.Username) +} + +func TestResolveOIDCUserUsernameCollision(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + local := &store.User{Username: "alice", Source: store.UserSourceLocal, PasswordHash: "x"} + require.NoError(t, users.Create(ctx, local)) + + u1, created, err := ResolveOIDCUser(ctx, users, oidcID("s1", "Alice"), true, now) + require.NoError(t, err) + assert.True(t, created) + assert.Equal(t, "Alice-2", u1.Username) + + u2, _, err := ResolveOIDCUser(ctx, users, oidcID("s2", "alice"), true, now) + require.NoError(t, err) + assert.Equal(t, "alice-3", u2.Username) + + got, err := users.GetByID(ctx, local.ID) + require.NoError(t, err) + assert.Equal(t, "x", got.PasswordHash) + assert.Empty(t, got.OIDCIssuer) + assert.Empty(t, got.OIDCSubject) +} + +func TestResolveOIDCUserNeverBindsAdmin(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + team := &store.Team{Name: "admins", Permissions: []string{}} + require.NoError(t, teams.Create(ctx, team)) + admin := &store.User{Username: "admin", Source: store.UserSourceLocal, PasswordHash: "x", Teams: []primitive.ObjectID{team.ID}} + require.NoError(t, users.Create(ctx, admin)) + + u, created, err := ResolveOIDCUser(ctx, users, oidcID("s1", "admin"), true, time.Now().UTC()) + require.NoError(t, err) + assert.True(t, created) + assert.NotEqual(t, admin.ID, u.ID) + assert.Equal(t, "admin-2", u.Username) + assert.Equal(t, store.UserSourceOIDC, u.Source) + assert.Empty(t, u.Teams) +} + +func TestResolveOIDCUserNeverBindsByEmail(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + + local := &store.User{Username: "carol", Email: "victim@x.io", DisplayName: "Carol", Source: store.UserSourceLocal, PasswordHash: "x"} + require.NoError(t, users.Create(ctx, local)) + before, err := users.GetByID(ctx, local.ID) + require.NoError(t, err) + + id := OIDCIdentity{Issuer: testIssuer, Subject: "s1", Username: "mallory", Email: "victim@x.io", DisplayName: "Mallory"} + u, created, err := ResolveOIDCUser(ctx, users, id, true, time.Now().UTC()) + require.NoError(t, err) + assert.True(t, created) + assert.NotEqual(t, local.ID, u.ID) + + after, err := users.GetByID(ctx, local.ID) + require.NoError(t, err) + assert.Equal(t, before, after) +} + +func TestResolveOIDCUserLongUsername(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + long := strings.Repeat("a", 64) + + require.NoError(t, users.Create(ctx, &store.User{Username: long, Source: store.UserSourceLocal, PasswordHash: "x"})) + u, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", long), true, time.Now().UTC()) + require.NoError(t, err) + assert.LessOrEqual(t, len(u.Username), 64) + assert.True(t, strings.HasSuffix(u.Username, "-2")) +} + +func TestResolveOIDCUserProvisioningDisabled(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + _, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), false, now) + assert.ErrorIs(t, err, ErrOIDCNotProvisioned) + n, err := users.Count(ctx) + require.NoError(t, err) + assert.Zero(t, n) + + _, _, err = ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + require.NoError(t, err) + u, created, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), false, now) + require.NoError(t, err) + assert.False(t, created) + assert.Equal(t, "alice", u.Username) +} + +func TestResolveOIDCUserDisabled(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + u, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + require.NoError(t, err) + u.Disabled = true + require.NoError(t, users.Update(ctx, u)) + + changed := OIDCIdentity{Issuer: testIssuer, Subject: "s1", Email: "other@x.io", DisplayName: "Other"} + _, _, err = ResolveOIDCUser(ctx, users, changed, true, now.Add(time.Hour)) + assert.ErrorIs(t, err, ErrOIDCUserDisabled) + + stored, err := users.GetByID(ctx, u.ID) + require.NoError(t, err) + assert.Equal(t, "alice@x.io", stored.Email) + assert.Equal(t, "alice", stored.DisplayName) +} + +func TestResolveOIDCUserNoUsername(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + _, _, err := ResolveOIDCUser(ctx, users, OIDCIdentity{Issuer: testIssuer, Subject: "s1"}, true, now) + assert.ErrorIs(t, err, ErrOIDCNoUsername) + + _, _, err = ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + require.NoError(t, err) + u, created, err := ResolveOIDCUser(ctx, users, OIDCIdentity{Issuer: testIssuer, Subject: "s1", Email: "a@x.io"}, true, now) + require.NoError(t, err) + assert.False(t, created) + assert.Equal(t, "alice", u.Username) +} + +func TestResolveOIDCUserConcurrentFirstLogin(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + const workers = 2 + ids := make([]primitive.ObjectID, workers) + errs := make([]error, workers) + start := make(chan struct{}) + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + <-start + u, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + errs[i] = err + if u != nil { + ids[i] = u.ID + } + }() + } + close(start) + wg.Wait() + + for i := 0; i < workers; i++ { + require.NoError(t, errs[i]) + } + assert.Equal(t, ids[0], ids[1]) + assert.False(t, ids[0].IsZero()) + n, err := users.Count(ctx) + require.NoError(t, err) + assert.EqualValues(t, 1, n) +} + +func TestCandidateUsername(t *testing.T) { + assert.Equal(t, "bob", candidateUsername("bob", 1)) + assert.Equal(t, "bob-2", candidateUsername("bob", 2)) + got := candidateUsername(strings.Repeat("a", 64), 12) + assert.Len(t, got, 64) + assert.True(t, strings.HasSuffix(got, "-12")) +} + +func TestResolveOIDCUserKeepsDisplayNameWhenClaimEmpty(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + first, _, err := ResolveOIDCUser(ctx, users, oidcID("s1", "alice"), true, now) + require.NoError(t, err) + u, _, err := ResolveOIDCUser(ctx, users, OIDCIdentity{Issuer: testIssuer, Subject: "s1", Email: "a@x.io"}, true, now) + require.NoError(t, err) + assert.Equal(t, "alice", u.DisplayName) + stored, err := users.GetByID(ctx, first.ID) + require.NoError(t, err) + assert.Equal(t, "alice", stored.DisplayName) + assert.Equal(t, "a@x.io", stored.Email) +} + +func TestResolveOIDCUserRequiresIssuerAndSubject(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + now := time.Now().UTC() + + _, _, err := ResolveOIDCUser(ctx, users, OIDCIdentity{Subject: "s1", Username: "alice"}, true, now) + assert.ErrorIs(t, err, ErrOIDCInvalidIdentity) + _, _, err = ResolveOIDCUser(ctx, users, OIDCIdentity{Issuer: testIssuer, Username: "alice"}, true, now) + assert.ErrorIs(t, err, ErrOIDCInvalidIdentity) + n, err := users.Count(ctx) + require.NoError(t, err) + assert.Zero(t, n) +} + +func TestResolveOIDCUserRefusesNonOIDCAccount(t *testing.T) { + users, _ := mongoStores(t) + ctx := context.Background() + + local := &store.User{Username: "eve", Email: "eve@x.io", Source: store.UserSourceLocal, PasswordHash: "x", OIDCIssuer: testIssuer, OIDCSubject: "s1"} + require.NoError(t, users.Create(ctx, local)) + + _, _, err := ResolveOIDCUser(ctx, users, OIDCIdentity{Issuer: testIssuer, Subject: "s1", Email: "evil@x.io", DisplayName: "Evil"}, true, time.Now().UTC()) + assert.ErrorIs(t, err, ErrOIDCNotOIDCUser) + stored, err := users.GetByID(ctx, local.ID) + require.NoError(t, err) + assert.Equal(t, "eve@x.io", stored.Email) + assert.Empty(t, stored.DisplayName) +} + +func TestPlanTeamSync(t *testing.T) { + p := &store.Team{ID: primitive.NewObjectID(), Name: "P", OIDCGroups: []string{"platform-eng"}} + o := &store.Team{ID: primitive.NewObjectID(), Name: "O", OIDCGroups: []string{"ops"}} + a := &store.Team{ID: primitive.NewObjectID(), Name: store.AdministratorsTeamName, Builtin: true, OIDCGroups: []string{"tracker-admins"}} + mapped := []*store.Team{p, o, a} + ids := func(ts ...*store.Team) []primitive.ObjectID { + out := []primitive.ObjectID{} + for _, x := range ts { + out = append(out, x.ID) + } + return out + } + + tests := []struct { + name string + current []*store.Team + groups []string + add []string + remove []string + }{ + {"join", nil, []string{"platform-eng"}, []string{"P"}, []string{}}, + {"unchanged", []*store.Team{p}, []string{"platform-eng"}, []string{}, []string{}}, + {"switch", []*store.Team{p}, []string{"ops"}, []string{"O"}, []string{"P"}}, + {"no groups", []*store.Team{p, o}, nil, []string{}, []string{"P", "O"}}, + {"case differs", []*store.Team{p}, []string{"Platform-Eng"}, []string{}, []string{"P"}}, + {"no trim", nil, []string{" platform-eng"}, []string{}, []string{}}, + {"admins left", []*store.Team{a}, nil, []string{}, []string{store.AdministratorsTeamName}}, + {"admins and ops", nil, []string{"tracker-admins", "ops"}, []string{store.AdministratorsTeamName, "O"}, []string{}}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + match, stale := planTeamSync(mapped, tc.groups) + cur := map[primitive.ObjectID]bool{} + for _, id := range ids(tc.current...) { + cur[id] = true + } + add, remove := []string{}, []string{} + for _, m := range match { + if !cur[m.ID] { + add = append(add, m.Name) + } + } + for _, st := range stale { + if cur[st.ID] { + remove = append(remove, st.Name) + } + } + assert.ElementsMatch(t, tc.add, add) + assert.ElementsMatch(t, tc.remove, remove) + }) + } +} + +func TestSyncOIDCTeams(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + _, err := Bootstrap(ctx, users, teams, "initial-admin-password") + require.NoError(t, err) + admins, err := teams.GetByName(ctx, store.AdministratorsTeamName) + require.NoError(t, err) + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + ops := &store.Team{Name: "Ops", OIDCGroups: []string{"ops"}} + manual := &store.Team{Name: "Manual", OIDCGroups: []string{}} + for _, tm := range []*store.Team{platform, ops, manual} { + require.NoError(t, teams.Create(ctx, tm)) + } + local, err := users.GetByUsername(ctx, "admin") + require.NoError(t, err) + initialAdminTeams := append([]primitive.ObjectID(nil), local.Teams...) + + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob", Teams: []primitive.ObjectID{manual.ID}} + require.NoError(t, users.Create(ctx, bob)) + reload := func() *store.User { + u, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + return u + } + + res, err := SyncOIDCTeams(ctx, users, teams, bob, []string{"platform-eng"}, true) + require.NoError(t, err) + assert.Equal(t, []string{"Platform"}, res.Added) + assert.ElementsMatch(t, []primitive.ObjectID{manual.ID, platform.ID}, reload().Teams) + assert.ElementsMatch(t, reload().Teams, bob.Teams) + + res, err = SyncOIDCTeams(ctx, users, teams, bob, []string{"ops"}, true) + require.NoError(t, err) + assert.Equal(t, []string{"Ops"}, res.Added) + assert.Equal(t, []string{"Platform"}, res.Removed) + assert.ElementsMatch(t, []primitive.ObjectID{manual.ID, ops.ID}, reload().Teams) + assert.ElementsMatch(t, reload().Teams, bob.Teams) + + // Administrators without oidcGroups is not mapped: left alone. + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{admins.ID}, nil)) + bob = reload() + _, err = SyncOIDCTeams(ctx, users, teams, bob, nil, true) + require.NoError(t, err) + assert.Contains(t, reload().Teams, admins.ID) + + // Mapped Administrators follows the claim while another admin is active. + admins.OIDCGroups = []string{"tracker-admins"} + require.NoError(t, teams.Update(ctx, admins)) + res, err = SyncOIDCTeams(ctx, users, teams, bob, []string{"tracker-admins"}, true) + require.NoError(t, err) + assert.Empty(t, res.Removed) + res, err = SyncOIDCTeams(ctx, users, teams, bob, nil, true) + require.NoError(t, err) + assert.Equal(t, []string{store.AdministratorsTeamName}, res.Removed) + assert.NotContains(t, reload().Teams, admins.ID) + + // Last enabled administrator is kept. + local, err = users.GetByUsername(ctx, "admin") + require.NoError(t, err) + local.Disabled = true + require.NoError(t, users.Update(ctx, local)) + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{admins.ID}, nil)) + bob = reload() + res, err = SyncOIDCTeams(ctx, users, teams, bob, nil, true) + require.NoError(t, err) + assert.Equal(t, []string{store.AdministratorsTeamName}, res.Kept) + assert.Empty(t, res.Removed) + assert.Contains(t, reload().Teams, admins.ID) + + after, err := users.GetByUsername(ctx, "admin") + require.NoError(t, err) + assert.Equal(t, initialAdminTeams, after.Teams) +} + +func TestSyncOIDCTeamsMissingClaim(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob"} + require.NoError(t, users.Create(ctx, bob)) + + // No mapped team: an absent claim is harmless. + _, err := SyncOIDCTeams(ctx, users, teams, bob, nil, false) + require.NoError(t, err) + + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + require.NoError(t, teams.Create(ctx, platform)) + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{platform.ID}, nil)) + bob, err = users.GetByID(ctx, bob.ID) + require.NoError(t, err) + + // Mapped teams, absent claim and a mapped membership: refused, nothing written. + _, err = SyncOIDCTeams(ctx, users, teams, bob, nil, false) + assert.ErrorIs(t, err, ErrOIDCGroupsClaimMissing) + got, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Equal(t, []primitive.ObjectID{platform.ID}, got.Teams) + + // Present but empty: mapped memberships are removed. + res, err := SyncOIDCTeams(ctx, users, teams, bob, []string{}, true) + require.NoError(t, err) + assert.Equal(t, []string{"Platform"}, res.Removed) + got, err = users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Empty(t, got.Teams) + + // Absent claim and no mapped membership: treated as no groups, no error. + manual := &store.Team{Name: "Manual"} + require.NoError(t, teams.Create(ctx, manual)) + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{manual.ID}, nil)) + bob, err = users.GetByID(ctx, bob.ID) + require.NoError(t, err) + res, err = SyncOIDCTeams(ctx, users, teams, bob, nil, false) + require.NoError(t, err) + assert.Empty(t, res.Added) + assert.Empty(t, res.Removed) + got, err = users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Equal(t, []primitive.ObjectID{manual.ID}, got.Teams) +} + +func TestSyncOIDCTeamsRemovesStaleMembership(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + require.NoError(t, teams.Create(ctx, platform)) + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob"} + require.NoError(t, users.Create(ctx, bob)) + + // Added in the database after bob was loaded. + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{platform.ID}, nil)) + res, err := SyncOIDCTeams(ctx, users, teams, bob, []string{}, true) + require.NoError(t, err) + assert.Empty(t, res.Removed) + got, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Empty(t, got.Teams) +} + +func TestSyncOIDCTeamsRefusesNonOIDCUser(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + require.NoError(t, teams.Create(ctx, platform)) + local := &store.User{Username: "loc", Source: store.UserSourceLocal, PasswordHash: "x"} + require.NoError(t, users.Create(ctx, local)) + + _, err := SyncOIDCTeams(ctx, users, teams, local, []string{"platform-eng"}, true) + assert.ErrorIs(t, err, ErrOIDCNotOIDCUser) + got, err := users.GetByID(ctx, local.ID) + require.NoError(t, err) + assert.Empty(t, got.Teams) +} + +func TestSyncOIDCTeamsReAddsMatchingMembership(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + require.NoError(t, teams.Create(ctx, platform)) + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob", Teams: []primitive.ObjectID{platform.ID}} + require.NoError(t, users.Create(ctx, bob)) + + // Removed in the database after bob was loaded. + require.NoError(t, users.SyncTeams(ctx, bob.ID, nil, []primitive.ObjectID{platform.ID})) + res, err := SyncOIDCTeams(ctx, users, teams, bob, []string{"platform-eng"}, true) + require.NoError(t, err) + assert.Empty(t, res.Added) + assert.Equal(t, []primitive.ObjectID{platform.ID}, bob.Teams) + got, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Equal(t, []primitive.ObjectID{platform.ID}, got.Teams) +} + +func TestSyncOIDCTeamsKeepsLastAdminAddedAfterLoad(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + _, err := Bootstrap(ctx, users, teams, "initial-admin-password") + require.NoError(t, err) + admins, err := teams.GetByName(ctx, store.AdministratorsTeamName) + require.NoError(t, err) + admins.OIDCGroups = []string{"tracker-admins"} + require.NoError(t, teams.Update(ctx, admins)) + local, err := users.GetByUsername(ctx, "admin") + require.NoError(t, err) + local.Disabled = true + require.NoError(t, users.Update(ctx, local)) + + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob"} + require.NoError(t, users.Create(ctx, bob)) + // Added in the database after bob was loaded: absent from bob.Teams. + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{admins.ID}, nil)) + + _, err = SyncOIDCTeams(ctx, users, teams, bob, []string{"other"}, true) + require.NoError(t, err) + got, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Contains(t, got.Teams, admins.ID) +} + +type fakeMappedTeams struct{ teams []*store.Team } + +func (f fakeMappedTeams) ListWithOIDCGroups(context.Context) ([]*store.Team, error) { + return f.teams, nil +} + +func TestCheckOIDCGroupsClaim(t *testing.T) { + ctx := context.Background() + platform := &store.Team{ID: primitive.NewObjectID(), Name: "Platform", OIDCGroups: []string{"g"}} + mapped := fakeMappedTeams{teams: []*store.Team{platform}} + holder := &store.User{Teams: []primitive.ObjectID{platform.ID}} + other := &store.User{Teams: []primitive.ObjectID{primitive.NewObjectID()}} + + assert.ErrorIs(t, CheckOIDCGroupsClaim(ctx, mapped, holder, false), ErrOIDCGroupsClaimMissing) + assert.NoError(t, CheckOIDCGroupsClaim(ctx, mapped, holder, true)) + assert.NoError(t, CheckOIDCGroupsClaim(ctx, mapped, other, false), "no mapped membership") + assert.NoError(t, CheckOIDCGroupsClaim(ctx, mapped, nil, false), "new subject") + assert.NoError(t, CheckOIDCGroupsClaim(ctx, fakeMappedTeams{}, holder, false), "no mapped team") +} + +func TestSyncOIDCTeamsAbsentClaimNeverRemoves(t *testing.T) { + users, teams := mongoStores(t) + ctx := context.Background() + + bob := &store.User{Username: "bob", Source: store.UserSourceOIDC, OIDCIssuer: testIssuer, OIDCSubject: "bob"} + require.NoError(t, users.Create(ctx, bob)) + platform := &store.Team{Name: "Platform", OIDCGroups: []string{"platform-eng"}} + require.NoError(t, teams.Create(ctx, platform)) + + // bob was loaded without any team; a mapped membership appears afterwards. + require.NoError(t, users.SyncTeams(ctx, bob.ID, []primitive.ObjectID{platform.ID}, nil)) + before, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + + res, err := SyncOIDCTeams(ctx, users, teams, bob, nil, false) + require.NoError(t, err) + assert.Empty(t, res.Added) + assert.Empty(t, res.Removed) + got, err := users.GetByID(ctx, bob.ID) + require.NoError(t, err) + assert.Equal(t, []primitive.ObjectID{platform.ID}, got.Teams) + assert.Equal(t, before.UpdatedAt, got.UpdatedAt, "no write") +} diff --git a/internal/auth/sso/claims.go b/internal/auth/sso/claims.go new file mode 100644 index 0000000..b463b8b --- /dev/null +++ b/internal/auth/sso/claims.go @@ -0,0 +1,128 @@ +// Package sso talks to an external OpenID Connect identity provider: it +// verifies id_tokens and extracts the claims Tracker keeps. +package sso + +import ( + "fmt" + "regexp" + + "github.com/bananaops/tracker/internal/auth" +) + +// usernamePattern is what a claim value must match to be used as a +// username: it must start with an alphanumeric character, 2 to 64 +// characters long overall. +var usernamePattern = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9@._+-]{1,63}$`) + +// Claims is what Tracker keeps from a verified id_token. +type Claims struct { + Issuer string + Subject string + Username string // first valid of the configured claim then email, "" when none + Email string + DisplayName string + Groups []string + // GroupsPresent is true when the groups claim exists in the token, even + // when it carries no usable group. + GroupsPresent bool + // GroupsUnexpectedType is true when the groups claim exists but is + // neither a string nor an array. It is then treated as absent + // (GroupsPresent is false), so it never removes mapped teams. + GroupsUnexpectedType bool +} + +// claimsFrom builds Claims from the raw id_token payload, using cfg to know +// which claim carries the username and the groups. +func claimsFrom(issuer, subject string, raw map[string]any, cfg auth.OIDCConfig) (Claims, error) { + if subject == "" { + return Claims{}, fmt.Errorf("%w: id_token has no subject", ErrClaims) + } + + email := stringClaim(raw, "email") + username := validUsername(stringClaim(raw, cfg.UsernameClaim)) + if username == "" { + username = validUsername(email) + } + + groups, present, unexpected := groupsClaim(raw, cfg.GroupsClaim) + + return Claims{ + Issuer: issuer, + Subject: subject, + Username: username, + Email: email, + DisplayName: displayName(raw, username), + Groups: groups, + GroupsPresent: present, + GroupsUnexpectedType: unexpected, + }, nil +} + +// validUsername returns v when it matches usernamePattern, "" otherwise. +func validUsername(v string) string { + if usernamePattern.MatchString(v) { + return v + } + return "" +} + +// displayName prefers the name claim, then given_name and family_name +// joined, then falls back to username. +func displayName(raw map[string]any, username string) string { + if name := stringClaim(raw, "name"); name != "" { + return name + } + + given := stringClaim(raw, "given_name") + family := stringClaim(raw, "family_name") + switch { + case given != "" && family != "": + return given + " " + family + case given != "": + return given + case family != "": + return family + } + + return username +} + +// stringClaim returns raw[name] when it is a string, "" otherwise. +func stringClaim(raw map[string]any, name string) string { + s, _ := raw[name].(string) + return s +} + +// groupsClaim reads a claim that is either a single string or an array of +// strings, deduplicating while keeping order and dropping non-string and +// empty entries. The second result is false when the claim is absent or of +// an unexpected type (object, number, bool, null), which the third result +// tells apart: such a claim carries no information about the groups. +func groupsClaim(raw map[string]any, name string) (groups []string, present, unexpected bool) { + v, ok := raw[name] + if !ok { + return nil, false, false + } + + switch t := v.(type) { + case string: + if t == "" { + return []string{}, true, false + } + return []string{t}, true, false + case []any: + seen := make(map[string]bool, len(t)) + out := make([]string, 0, len(t)) + for _, e := range t { + s, ok := e.(string) + if !ok || s == "" || seen[s] { + continue + } + seen[s] = true + out = append(out, s) + } + return out, true, false + default: + return nil, false, true + } +} diff --git a/internal/auth/sso/claims_test.go b/internal/auth/sso/claims_test.go new file mode 100644 index 0000000..5ea33c7 --- /dev/null +++ b/internal/auth/sso/claims_test.go @@ -0,0 +1,223 @@ +package sso + +import ( + "errors" + "testing" + + "github.com/bananaops/tracker/internal/auth" +) + +func defaultTestOIDCConfig() auth.OIDCConfig { + return auth.OIDCConfig{ + GroupsClaim: "groups", + UsernameClaim: "preferred_username", + } +} + +func TestClaimsFromUsername(t *testing.T) { + cases := []struct { + name string + raw map[string]any + want string + }{ + { + name: "preferred username", + raw: map[string]any{"preferred_username": "alice", "email": "a@x.io"}, + want: "alice", + }, + { + name: "email fallback", + raw: map[string]any{"email": "alice@example.com"}, + want: "alice@example.com", + }, + { + name: "invalid preferred falls back to email", + raw: map[string]any{"preferred_username": "bad name!", "email": "a@x.io"}, + want: "a@x.io", + }, + { + name: "none", + raw: map[string]any{}, + want: "", + }, + { + name: "preferred not a string", + raw: map[string]any{"preferred_username": 42, "email": "a@x.io"}, + want: "a@x.io", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := claimsFrom("https://issuer.example", "sub-1", tc.raw, defaultTestOIDCConfig()) + if err != nil { + t.Fatalf("claimsFrom: %v", err) + } + if got.Username != tc.want { + t.Fatalf("Username = %q, want %q", got.Username, tc.want) + } + }) + } +} + +func TestClaimsFromGroups(t *testing.T) { + cases := []struct { + name string + raw map[string]any + cfg *auth.OIDCConfig + wantGroups []string + wantPresent bool + wantUnexp bool + checkGroupLen bool // when true, also assert len(Groups) == len(wantGroups) for the empty-slice case + }{ + { + name: "array deduplicated", + raw: map[string]any{"groups": []any{"a", "b", "a"}}, + wantGroups: []string{"a", "b"}, + wantPresent: true, + }, + { + name: "single string", + raw: map[string]any{"groups": "a"}, + wantGroups: []string{"a"}, + wantPresent: true, + }, + { + name: "mixed types filtered", + raw: map[string]any{"groups": []any{"a", 3, nil, ""}}, + wantGroups: []string{"a"}, + wantPresent: true, + }, + { + name: "absent", + raw: map[string]any{}, + wantGroups: nil, + wantPresent: false, + }, + { + name: "empty array", + raw: map[string]any{"groups": []any{}}, + wantGroups: []string{}, + wantPresent: true, + checkGroupLen: true, + }, + { + name: "object treated as absent", + raw: map[string]any{"groups": map[string]any{"a": true}}, + wantGroups: nil, + wantPresent: false, + wantUnexp: true, + }, + { + name: "number treated as absent", + raw: map[string]any{"groups": float64(3)}, + wantGroups: nil, + wantPresent: false, + wantUnexp: true, + }, + { + name: "bool treated as absent", + raw: map[string]any{"groups": true}, + wantGroups: nil, + wantPresent: false, + wantUnexp: true, + }, + { + name: "null treated as absent", + raw: map[string]any{"groups": nil}, + wantGroups: nil, + wantPresent: false, + wantUnexp: true, + }, + { + name: "custom claim", + cfg: &auth.OIDCConfig{GroupsClaim: "roles", UsernameClaim: "preferred_username"}, + raw: map[string]any{"roles": []any{"x"}}, + wantGroups: []string{"x"}, + wantPresent: true, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := defaultTestOIDCConfig() + if tc.cfg != nil { + cfg = *tc.cfg + } + got, err := claimsFrom("https://issuer.example", "sub-1", tc.raw, cfg) + if err != nil { + t.Fatalf("claimsFrom: %v", err) + } + if got.GroupsUnexpectedType != tc.wantUnexp { + t.Fatalf("GroupsUnexpectedType = %v, want %v", got.GroupsUnexpectedType, tc.wantUnexp) + } + if got.GroupsPresent != tc.wantPresent { + t.Fatalf("GroupsPresent = %v, want %v", got.GroupsPresent, tc.wantPresent) + } + if tc.checkGroupLen { + if len(got.Groups) != len(tc.wantGroups) { + t.Fatalf("Groups = %v, want length %d", got.Groups, len(tc.wantGroups)) + } + return + } + if !equalStrings(got.Groups, tc.wantGroups) { + t.Fatalf("Groups = %v, want %v", got.Groups, tc.wantGroups) + } + }) + } +} + +func TestClaimsFromDisplayName(t *testing.T) { + cases := []struct { + name string + raw map[string]any + want string + }{ + { + name: "name claim", + raw: map[string]any{"name": "Alice A"}, + want: "Alice A", + }, + { + name: "given plus family name", + raw: map[string]any{"given_name": "Alice", "family_name": "Example"}, + want: "Alice Example", + }, + { + name: "falls back to username", + raw: map[string]any{"preferred_username": "alice"}, + want: "alice", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := claimsFrom("https://issuer.example", "sub-1", tc.raw, defaultTestOIDCConfig()) + if err != nil { + t.Fatalf("claimsFrom: %v", err) + } + if got.DisplayName != tc.want { + t.Fatalf("DisplayName = %q, want %q", got.DisplayName, tc.want) + } + }) + } +} + +func TestClaimsFromEmptySubject(t *testing.T) { + _, err := claimsFrom("https://issuer.example", "", map[string]any{}, defaultTestOIDCConfig()) + if !errors.Is(err, ErrClaims) { + t.Fatalf("err = %v, want ErrClaims", err) + } +} + +func equalStrings(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/internal/auth/sso/provider.go b/internal/auth/sso/provider.go new file mode 100644 index 0000000..ca17faf --- /dev/null +++ b/internal/auth/sso/provider.go @@ -0,0 +1,194 @@ +package sso + +import ( + "context" + "crypto/subtle" + "errors" + "fmt" + "net/http" + "sync" + "time" + + gooidc "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" + + "github.com/bananaops/tracker/internal/auth" +) + +const ( + defaultHTTPTimeout = 10 * time.Second + defaultRetryInterval = 5 * time.Second +) + +var ( + // ErrUnavailable means discovery could not reach or parse the identity + // provider's configuration. + ErrUnavailable = errors.New("identity provider unavailable") + // ErrExchange means the authorization code exchange failed. + ErrExchange = errors.New("authorization code exchange failed") + // ErrIDToken means the id_token was rejected: bad signature, issuer, + // audience, expiry or nonce. + ErrIDToken = errors.New("id_token rejected") + // ErrClaims means the id_token claims could not be turned into usable + // Claims. + ErrClaims = errors.New("id_token claims unusable") +) + +// Provider is the part of the identity provider used by the HTTP handlers. +type Provider interface { + AuthCodeURL(ctx context.Context, state, nonce, verifier string) (string, error) + Exchange(ctx context.Context, code, verifier, nonce string) (Claims, error) +} + +// ProviderOption configures an OIDCProvider. +type ProviderOption func(*OIDCProvider) + +// WithHTTPClient sets the HTTP client used for discovery, JWKS fetches and +// the token exchange. +func WithHTTPClient(c *http.Client) ProviderOption { + return func(p *OIDCProvider) { + p.httpClient = c + } +} + +// WithClock overrides the time source, for tests. +func WithClock(now func() time.Time) ProviderOption { + return func(p *OIDCProvider) { + p.now = now + } +} + +// WithRetryInterval overrides the minimum delay between two discovery +// attempts after a failure. +func WithRetryInterval(d time.Duration) ProviderOption { + return func(p *OIDCProvider) { + p.retry = d + } +} + +// OIDCProvider talks to the identity provider. Discovery is lazy: it runs on +// the first use and is retried at most once per retry interval after a +// failure, so an unreachable provider never prevents Tracker from starting. +type OIDCProvider struct { + cfg auth.OIDCConfig + redirectURL string + httpClient *http.Client + now func() time.Time + retry time.Duration + + mu sync.Mutex + oauth *oauth2.Config + verifier *gooidc.IDTokenVerifier + lastErr error + lastAttempt time.Time +} + +var _ Provider = (*OIDCProvider)(nil) + +// NewOIDCProvider builds a client for cfg's issuer. redirectURL is the +// Tracker callback URL registered with the identity provider. +func NewOIDCProvider(cfg auth.OIDCConfig, redirectURL string, opts ...ProviderOption) *OIDCProvider { + p := &OIDCProvider{ + cfg: cfg, + redirectURL: redirectURL, + httpClient: &http.Client{Timeout: defaultHTTPTimeout}, + now: time.Now, + retry: defaultRetryInterval, + } + for _, o := range opts { + o(p) + } + return p +} + +// Discover runs discovery now; used as a startup warm-up. Safe to call +// concurrently. +func (p *OIDCProvider) Discover(_ context.Context) error { + _, _, err := p.ready() + return err +} + +// ready returns the oauth2 config and id_token verifier, running discovery +// on first use and retrying at most once per retry interval after a +// failure. +func (p *OIDCProvider) ready() (*oauth2.Config, *gooidc.IDTokenVerifier, error) { + p.mu.Lock() + defer p.mu.Unlock() + + if p.oauth != nil { + return p.oauth, p.verifier, nil + } + if p.lastErr != nil && p.now().Sub(p.lastAttempt) < p.retry { + return nil, nil, fmt.Errorf("%w: %w", ErrUnavailable, p.lastErr) + } + p.lastAttempt = p.now() + + // No deadline on this context: the remote key set may keep it for later + // JWKS refreshes. The HTTP client timeout bounds every request instead. + dctx := gooidc.ClientContext(context.Background(), p.httpClient) + provider, err := gooidc.NewProvider(dctx, p.cfg.Issuer) + if err != nil { + p.lastErr = err + return nil, nil, fmt.Errorf("%w: %w", ErrUnavailable, err) + } + p.lastErr = nil + + p.oauth = &oauth2.Config{ + ClientID: p.cfg.ClientID, + ClientSecret: p.cfg.ClientSecret, + Endpoint: provider.Endpoint(), + RedirectURL: p.redirectURL, + Scopes: p.cfg.Scopes, + } + p.verifier = provider.Verifier(&gooidc.Config{ClientID: p.cfg.ClientID, Now: p.now}) + + return p.oauth, p.verifier, nil +} + +// AuthCodeURL builds the authorization URL: PKCE S256 challenge derived +// from verifier, plus the nonce that Exchange will check against the +// id_token. +func (p *OIDCProvider) AuthCodeURL(_ context.Context, state, nonce, verifier string) (string, error) { + oc, _, err := p.ready() + if err != nil { + return "", err + } + return oc.AuthCodeURL(state, gooidc.Nonce(nonce), oauth2.S256ChallengeOption(verifier)), nil +} + +// Exchange trades an authorization code for a verified id_token and returns +// its claims. verifier is the PKCE code_verifier generated for AuthCodeURL, +// nonce is the value AuthCodeURL sent. +func (p *OIDCProvider) Exchange(ctx context.Context, code, verifier, nonce string) (Claims, error) { + oc, idv, err := p.ready() + if err != nil { + return Claims{}, err + } + + ctx = gooidc.ClientContext(ctx, p.httpClient) + tok, err := oc.Exchange(ctx, code, oauth2.VerifierOption(verifier)) + if err != nil { + return Claims{}, fmt.Errorf("%w: %w", ErrExchange, err) + } + + rawIDToken, ok := tok.Extra("id_token").(string) + if !ok || rawIDToken == "" { + return Claims{}, fmt.Errorf("%w: token response carries no id_token", ErrIDToken) + } + + idt, err := idv.Verify(ctx, rawIDToken) + if err != nil { + return Claims{}, fmt.Errorf("%w: %w", ErrIDToken, err) + } + + if nonce == "" || subtle.ConstantTimeCompare([]byte(idt.Nonce), []byte(nonce)) != 1 { + return Claims{}, fmt.Errorf("%w: nonce mismatch", ErrIDToken) + } + + var raw map[string]any + if err := idt.Claims(&raw); err != nil { + return Claims{}, fmt.Errorf("%w: %w", ErrClaims, err) + } + + return claimsFrom(idt.Issuer, idt.Subject, raw, p.cfg) +} diff --git a/internal/auth/sso/provider_test.go b/internal/auth/sso/provider_test.go new file mode 100644 index 0000000..0d2d6ac --- /dev/null +++ b/internal/auth/sso/provider_test.go @@ -0,0 +1,458 @@ +package sso + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "golang.org/x/oauth2" + + "github.com/bananaops/tracker/internal/auth" + "github.com/bananaops/tracker/internal/auth/sso/ssotest" +) + +func newTestProvider(t *testing.T, idp *ssotest.IdP, opts ...ProviderOption) *OIDCProvider { + t.Helper() + cfg := auth.OIDCConfig{ + Issuer: idp.URL, ClientID: idp.ClientID, ClientSecret: idp.ClientSecret, + Scopes: []string{"openid", "profile", "email"}, GroupsClaim: "groups", + UsernameClaim: "preferred_username", UserProvisioning: true, TeamSync: true, + } + return NewOIDCProvider(cfg, "http://tracker.test"+auth.OIDCCallbackPath, opts...) +} + +// flow runs AuthCodeURL, the IdP authorization and Exchange with the given +// nonce. AuthCodeURL always uses "n-1" as its own nonce; exchangeNonce is +// what the caller then presents to Exchange, letting tests simulate a +// mismatch. +func flow(t *testing.T, p *OIDCProvider, idp *ssotest.IdP, exchangeNonce string) (Claims, error) { + t.Helper() + + verifier := oauth2.GenerateVerifier() + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + + return p.Exchange(context.Background(), code, verifier, exchangeNonce) +} + +func scopeContains(scope, want string) bool { + for _, s := range strings.Fields(scope) { + if s == want { + return true + } + } + return false +} + +func TestAuthCodeURL(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + verifier := oauth2.GenerateVerifier() + + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + if !strings.HasPrefix(authURL, idp.URL+"/authorize") { + t.Fatalf("authURL = %q, want prefix %q", authURL, idp.URL+"/authorize") + } + + u, err := url.Parse(authURL) + if err != nil { + t.Fatalf("parse authURL: %v", err) + } + q := u.Query() + + if q.Get("response_type") != "code" { + t.Fatalf("response_type = %q, want code", q.Get("response_type")) + } + if q.Get("client_id") != idp.ClientID { + t.Fatalf("client_id = %q, want %q", q.Get("client_id"), idp.ClientID) + } + if want := "http://tracker.test" + auth.OIDCCallbackPath; q.Get("redirect_uri") != want { + t.Fatalf("redirect_uri = %q, want %q", q.Get("redirect_uri"), want) + } + if !scopeContains(q.Get("scope"), "openid") { + t.Fatalf("scope = %q, want it to contain openid", q.Get("scope")) + } + if q.Get("state") != "st" { + t.Fatalf("state = %q, want st", q.Get("state")) + } + if q.Get("nonce") != "n-1" { + t.Fatalf("nonce = %q, want n-1", q.Get("nonce")) + } + if q.Get("code_challenge_method") != "S256" { + t.Fatalf("code_challenge_method = %q, want S256", q.Get("code_challenge_method")) + } + challenge := q.Get("code_challenge") + if challenge == "" { + t.Fatal("code_challenge is empty") + } + if challenge == verifier { + t.Fatal("code_challenge equals the verifier") + } +} + +func TestExchangeSuccess(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + + claims, err := flow(t, p, idp, "n-1") + if err != nil { + t.Fatalf("flow: %v", err) + } + if claims.Issuer != idp.URL { + t.Fatalf("Issuer = %q, want %q", claims.Issuer, idp.URL) + } + if claims.Subject != "user-1" { + t.Fatalf("Subject = %q, want user-1", claims.Subject) + } + if claims.Username != "alice" { + t.Fatalf("Username = %q, want alice", claims.Username) + } + if claims.Email == "" { + t.Fatal("Email is empty") + } + if claims.DisplayName != "Alice Example" { + t.Fatalf("DisplayName = %q, want Alice Example", claims.DisplayName) + } + if !claims.GroupsPresent { + t.Fatal("GroupsPresent = false, want true") + } +} + +func TestExchangeGroups(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + + idp.SetUser(ssotest.User{ + Subject: "user-2", + Claims: map[string]any{ + "preferred_username": "bob", + "email": "bob@example.com", + "groups": []string{"platform-eng", "ops"}, + }, + }) + + claims, err := flow(t, p, idp, "n-1") + if err != nil { + t.Fatalf("flow: %v", err) + } + if !equalStrings(claims.Groups, []string{"platform-eng", "ops"}) { + t.Fatalf("Groups = %v, want [platform-eng ops]", claims.Groups) + } +} + +func TestExchangeNonceMismatch(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + + _, err := flow(t, p, idp, "other-nonce") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeWrongAudience(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { + claims["aud"] = "someone-else" + }) + + _, err := flow(t, p, idp, "n-1") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeExpiredToken(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { + claims["exp"] = time.Now().Add(-time.Hour).Unix() + }) + + _, err := flow(t, p, idp, "n-1") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeWrongIssuer(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { + claims["iss"] = "https://evil.example" + }) + + _, err := flow(t, p, idp, "n-1") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeAlgNone(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + idp.SetSigning(ssotest.SignNone) + + _, err := flow(t, p, idp, "n-1") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeUnknownKey(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + idp.SetSigning(ssotest.SignForeignKey) + + _, err := flow(t, p, idp, "n-1") + if !errors.Is(err, ErrIDToken) { + t.Fatalf("err = %v, want ErrIDToken", err) + } +} + +func TestExchangeWrongVerifier(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + + verifier := oauth2.GenerateVerifier() + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + + otherVerifier := oauth2.GenerateVerifier() + _, err = p.Exchange(context.Background(), code, otherVerifier, "n-1") + if !errors.Is(err, ErrExchange) { + t.Fatalf("err = %v, want ErrExchange", err) + } +} + +func TestExchangeReusedCode(t *testing.T) { + idp := ssotest.New(t) + p := newTestProvider(t, idp) + + verifier := oauth2.GenerateVerifier() + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + + if _, err := p.Exchange(context.Background(), code, verifier, "n-1"); err != nil { + t.Fatalf("first Exchange: %v", err) + } + if _, err := p.Exchange(context.Background(), code, verifier, "n-1"); !errors.Is(err, ErrExchange) { + t.Fatalf("second Exchange err = %v, want ErrExchange", err) + } +} + +// flakyTransport fails the first request, then delegates every later one. +// It lets a test simulate a provider that is briefly unreachable without +// changing the issuer's URL, which would trip the issuer-equality check. +type flakyTransport struct { + mu sync.Mutex + failed bool + inner http.RoundTripper +} + +func (t *flakyTransport) RoundTrip(req *http.Request) (*http.Response, error) { + t.mu.Lock() + shouldFail := !t.failed + t.failed = true + t.mu.Unlock() + if shouldFail { + return nil, errors.New("simulated network failure") + } + return t.inner.RoundTrip(req) +} + +func TestDiscoveryUnavailableThenRetry(t *testing.T) { + var calls int32 + badServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + atomic.AddInt32(&calls, 1) + w.WriteHeader(http.StatusInternalServerError) + })) + defer badServer.Close() + + clock := time.Now() + now := func() time.Time { return clock } + + cfg := auth.OIDCConfig{Issuer: badServer.URL, ClientID: "client", Scopes: []string{"openid"}} + p := NewOIDCProvider(cfg, "http://tracker.test"+auth.OIDCCallbackPath, WithRetryInterval(time.Minute), WithClock(now)) + + if _, err := p.AuthCodeURL(context.Background(), "st", "n-1", "verifier"); !errors.Is(err, ErrUnavailable) { + t.Fatalf("first AuthCodeURL err = %v, want ErrUnavailable", err) + } + if got := atomic.LoadInt32(&calls); got != 1 { + t.Fatalf("calls = %d, want 1", got) + } + + if _, err := p.AuthCodeURL(context.Background(), "st", "n-1", "verifier"); !errors.Is(err, ErrUnavailable) { + t.Fatalf("second AuthCodeURL err = %v, want ErrUnavailable", err) + } + if got := atomic.LoadInt32(&calls); got != 1 { + t.Fatalf("calls after immediate retry = %d, want 1", got) + } + + clock = clock.Add(2 * time.Minute) + + if _, err := p.AuthCodeURL(context.Background(), "st", "n-1", "verifier"); !errors.Is(err, ErrUnavailable) { + t.Fatalf("third AuthCodeURL err = %v, want ErrUnavailable", err) + } + if got := atomic.LoadInt32(&calls); got != 2 { + t.Fatalf("calls after clock advance = %d, want 2", got) + } + + // A real IdP behind a transport that fails once: after the error, + // advancing the clock lets the next attempt succeed. + idp := ssotest.New(t) + transport := &flakyTransport{inner: http.DefaultTransport} + client := &http.Client{Transport: transport} + + p2 := NewOIDCProvider(auth.OIDCConfig{ + Issuer: idp.URL, ClientID: idp.ClientID, ClientSecret: idp.ClientSecret, Scopes: []string{"openid"}, + }, "http://tracker.test"+auth.OIDCCallbackPath, WithHTTPClient(client), WithRetryInterval(time.Minute), WithClock(now)) + + if err := p2.Discover(context.Background()); !errors.Is(err, ErrUnavailable) { + t.Fatalf("first Discover err = %v, want ErrUnavailable", err) + } + + clock = clock.Add(2 * time.Minute) + + if err := p2.Discover(context.Background()); err != nil { + t.Fatalf("second Discover err = %v, want nil", err) + } +} + +func TestDiscoveryIssuerMismatch(t *testing.T) { + idp := ssotest.New(t) + cfg := auth.OIDCConfig{Issuer: idp.URL + "/", ClientID: idp.ClientID, ClientSecret: idp.ClientSecret, Scopes: []string{"openid"}} + p := NewOIDCProvider(cfg, "http://tracker.test"+auth.OIDCCallbackPath) + + if err := p.Discover(context.Background()); !errors.Is(err, ErrUnavailable) { + t.Fatalf("err = %v, want ErrUnavailable", err) + } +} + +func TestErrorsDoNotLeakSecrets(t *testing.T) { + idp := ssotest.New(t) + + forbidden := []string{idp.ClientSecret} + + check := func(t *testing.T, err error) { + t.Helper() + if err == nil { + t.Fatal("err is nil, want a rejection") + } + msg := err.Error() + for _, s := range forbidden { + if s != "" && strings.Contains(msg, s) { + t.Fatalf("error %q leaks %q", msg, s) + } + } + if last := idp.LastIDToken(); last != "" && strings.Contains(msg, last) { + t.Fatalf("error %q leaks the id_token", msg) + } + } + + t.Run("NonceMismatch", func(t *testing.T) { + p := newTestProvider(t, idp) + _, err := flow(t, p, idp, "other-nonce") + check(t, err) + }) + + t.Run("WrongAudience", func(t *testing.T) { + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { claims["aud"] = "someone-else" }) + defer idp.SetTokenMutator(nil) + _, err := flow(t, p, idp, "n-1") + check(t, err) + }) + + t.Run("ExpiredToken", func(t *testing.T) { + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { claims["exp"] = time.Now().Add(-time.Hour).Unix() }) + defer idp.SetTokenMutator(nil) + _, err := flow(t, p, idp, "n-1") + check(t, err) + }) + + t.Run("WrongIssuer", func(t *testing.T) { + p := newTestProvider(t, idp) + idp.SetTokenMutator(func(claims map[string]any) { claims["iss"] = "https://evil.example" }) + defer idp.SetTokenMutator(nil) + _, err := flow(t, p, idp, "n-1") + check(t, err) + }) + + t.Run("AlgNone", func(t *testing.T) { + p := newTestProvider(t, idp) + idp.SetSigning(ssotest.SignNone) + defer idp.SetSigning(ssotest.SignRS256) + _, err := flow(t, p, idp, "n-1") + check(t, err) + }) + + t.Run("UnknownKey", func(t *testing.T) { + p := newTestProvider(t, idp) + idp.SetSigning(ssotest.SignForeignKey) + defer idp.SetSigning(ssotest.SignRS256) + _, err := flow(t, p, idp, "n-1") + check(t, err) + }) + + t.Run("WrongVerifier", func(t *testing.T) { + p := newTestProvider(t, idp) + verifier := oauth2.GenerateVerifier() + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + otherVerifier := oauth2.GenerateVerifier() + _, err = p.Exchange(context.Background(), code, otherVerifier, "n-1") + check(t, err) + if strings.Contains(err.Error(), verifier) || strings.Contains(err.Error(), otherVerifier) { + t.Fatalf("error %q leaks the verifier", err.Error()) + } + }) + + t.Run("ReusedCode", func(t *testing.T) { + p := newTestProvider(t, idp) + verifier := oauth2.GenerateVerifier() + authURL, err := p.AuthCodeURL(context.Background(), "st", "n-1", verifier) + if err != nil { + t.Fatalf("AuthCodeURL: %v", err) + } + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + if _, err := p.Exchange(context.Background(), code, verifier, "n-1"); err != nil { + t.Fatalf("first Exchange: %v", err) + } + _, err = p.Exchange(context.Background(), code, verifier, "n-1") + check(t, err) + if strings.Contains(err.Error(), code) { + t.Fatalf("error %q leaks the code", err.Error()) + } + }) +} diff --git a/internal/auth/sso/redirect.go b/internal/auth/sso/redirect.go new file mode 100644 index 0000000..caa1e5d --- /dev/null +++ b/internal/auth/sso/redirect.go @@ -0,0 +1,29 @@ +package sso + +import ( + "net/url" + "strings" + "unicode" +) + +const maxRedirectLength = 1024 + +// SafeRedirect returns raw when it is a local absolute path, "/" otherwise. +// It is the server side twin of the web safeRedirect helper and blocks open +// redirects such as //evil.example or /\evil.example. +func SafeRedirect(raw string) string { + if raw == "" || len(raw) > maxRedirectLength || raw[0] != '/' || + strings.HasPrefix(raw, "//") || strings.ContainsRune(raw, '\\') { + return "/" + } + for _, r := range raw { + if unicode.IsControl(r) { + return "/" + } + } + u, err := url.Parse(raw) + if err != nil || u.Scheme != "" || u.Host != "" || u.User != nil { + return "/" + } + return raw +} diff --git a/internal/auth/sso/redirect_test.go b/internal/auth/sso/redirect_test.go new file mode 100644 index 0000000..c4a8d15 --- /dev/null +++ b/internal/auth/sso/redirect_test.go @@ -0,0 +1,40 @@ +package sso + +import ( + "strings" + "testing" +) + +func TestSafeRedirect(t *testing.T) { + cases := []struct { + name string + in string + want string + }{ + {name: "empty", in: "", want: "/"}, + {name: "root", in: "/", want: "/"}, + {name: "simple path", in: "/locks", want: "/locks"}, + {name: "path with query and fragment", in: "/events?service=api&tab=1#top", want: "/events?service=api&tab=1#top"}, + {name: "encoded local path", in: "/%2F%2Fevil.example", want: "/%2F%2Fevil.example"}, + {name: "protocol relative", in: "//evil.example", want: "/"}, + {name: "triple slash", in: "///evil.example", want: "/"}, + {name: "backslash escape", in: "/\\evil.example", want: "/"}, + {name: "embedded backslash", in: "/a\\b", want: "/"}, + {name: "absolute url", in: "https://evil.example", want: "/"}, + {name: "no leading slash", in: "evil.example", want: "/"}, + {name: "javascript scheme", in: "javascript:alert(1)", want: "/"}, + {name: "header injection", in: "/ok\r\nSet-Cookie: x=1", want: "/"}, + {name: "tab control char", in: "/tab\there", want: "/"}, + {name: "C1 control char", in: "/a\u0085b", want: "/"}, + {name: "leading space", in: " /locks", want: "/"}, + {name: "too long", in: "/" + strings.Repeat("a", 1024), want: "/"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := SafeRedirect(tc.in); got != tc.want { + t.Fatalf("SafeRedirect(%q) = %q, want %q", tc.in, got, tc.want) + } + }) + } +} diff --git a/internal/auth/sso/ssotest/idp.go b/internal/auth/sso/ssotest/idp.go new file mode 100644 index 0000000..d7dadf3 --- /dev/null +++ b/internal/auth/sso/ssotest/idp.go @@ -0,0 +1,493 @@ +// Package ssotest runs an in-process OpenID Connect provider for tests. It +// is imported by tests only and never linked into the tracker binary. +package ssotest + +import ( + "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "crypto/subtle" + "crypto/x509" + "encoding/base64" + "encoding/json" + "encoding/pem" + "fmt" + "math/big" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// KeyID is the kid published in the JWKS and set on every signed token +// header, including tokens signed with the foreign, unpublished key. +const KeyID = "ssotest-key" + +// SigningMode selects how the next id_token is signed. +type SigningMode int + +const ( + // SignRS256 produces a valid signature with the published key. + SignRS256 SigningMode = iota + // SignNone produces an unsigned token, alg "none". + SignNone + // SignForeignKey produces an RS256 token signed with a key absent + // from the JWKS, reusing the same kid. + SignForeignKey + // SignHS256PublicKey produces an HS256 token whose HMAC secret is the + // PEM encoded published public key: the algorithm confusion attack. + SignHS256PublicKey +) + +// User is the identity returned by the next authorizations. +type User struct { + Subject string + // Claims is merged into the id_token: preferred_username, email, + // name, groups... + Claims map[string]any +} + +// authState is the server-side memory of a pending authorization code. +type authState struct { + challenge string + nonce string + redirectURI string + user User +} + +// IdP is an in-process OpenID Connect provider for tests. +type IdP struct { + // URL is the issuer, http://127.0.0.1:, no trailing slash. + URL string + ClientID string + ClientSecret string + + server *httptest.Server + + key *rsa.PrivateKey // published in the JWKS + foreign *rsa.PrivateKey // never published, same kid + + mu sync.Mutex + user User + signing SigningMode + mutator func(claims map[string]any) + authErrCode string + authErrDesc string + codes map[string]authState + tokenRequests int + lastIDToken string +} + +// New starts the server and registers t.Cleanup(Close). +func New(t testing.TB) *IdP { + t.Helper() + + key, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("generate signing key: %v", err) + } + foreign, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("generate foreign key: %v", err) + } + + secret := make([]byte, 32) + if _, err := rand.Read(secret); err != nil { + t.Fatalf("generate client secret: %v", err) + } + + idp := &IdP{ + ClientID: "tracker-test", + ClientSecret: base64.RawURLEncoding.EncodeToString(secret), + key: key, + foreign: foreign, + user: User{ + Subject: "user-1", + Claims: map[string]any{ + "preferred_username": "alice", + "email": "alice@example.com", + "name": "Alice Example", + "groups": []string{}, + }, + }, + codes: make(map[string]authState), + } + + mux := http.NewServeMux() + mux.HandleFunc("GET /.well-known/openid-configuration", idp.handleDiscovery) + mux.HandleFunc("GET /keys", idp.handleKeys) + mux.HandleFunc("GET /authorize", idp.handleAuthorize) + mux.HandleFunc("POST /token", idp.handleToken) + + idp.server = httptest.NewServer(mux) + idp.URL = idp.server.URL + t.Cleanup(idp.Close) + + return idp +} + +// Close shuts the server down. +func (i *IdP) Close() { + i.server.Close() +} + +// SetUser sets the identity returned by the next authorizations. +func (i *IdP) SetUser(u User) { + i.mu.Lock() + defer i.mu.Unlock() + i.user = u +} + +// SetSigning selects how the next id_token is signed. +func (i *IdP) SetSigning(m SigningMode) { + i.mu.Lock() + defer i.mu.Unlock() + i.signing = m +} + +// SetTokenMutator sets a function applied to the id_token claims last, +// right before signing. +func (i *IdP) SetTokenMutator(f func(claims map[string]any)) { + i.mu.Lock() + defer i.mu.Unlock() + i.mutator = f +} + +// SetAuthorizeError makes /authorize redirect with error=code instead of +// running the normal flow. An empty code disables it. +func (i *IdP) SetAuthorizeError(code, description string) { + i.mu.Lock() + defer i.mu.Unlock() + i.authErrCode = code + i.authErrDesc = description +} + +// TokenRequests returns how many requests /token has received. +func (i *IdP) TokenRequests() int { + i.mu.Lock() + defer i.mu.Unlock() + return i.tokenRequests +} + +// LastIDToken returns the id_token issued by the last successful /token +// request. +func (i *IdP) LastIDToken() string { + i.mu.Lock() + defer i.mu.Unlock() + return i.lastIDToken +} + +// Authorize plays the browser at the IdP: it issues a GET on authURL +// without following the redirect and returns the parsed Location (the +// Tracker callback URL). +func (i *IdP) Authorize(t testing.TB, authURL string) *url.URL { + t.Helper() + + client := &http.Client{ + Timeout: 5 * time.Second, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + } + + req, err := http.NewRequest(http.MethodGet, authURL, nil) //nolint:noctx // test-only, one-shot request against a loopback test server + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + + res, err := client.Do(req) // #nosec G107 -- test identity provider on loopback, never linked into the binary + if err != nil { + t.Fatalf("GET %s: %v", authURL, err) + } + defer res.Body.Close() + + if res.StatusCode != http.StatusFound { + t.Fatalf("status = %d, want %d", res.StatusCode, http.StatusFound) + } + + loc, err := url.Parse(res.Header.Get("Location")) + if err != nil { + t.Fatalf("parse Location: %v", err) + } + return loc +} + +func (i *IdP) handleDiscovery(w http.ResponseWriter, _ *http.Request) { + doc := map[string]any{ + "issuer": i.URL, + "authorization_endpoint": i.URL + "/authorize", + "token_endpoint": i.URL + "/token", + "jwks_uri": i.URL + "/keys", + "response_types_supported": []string{"code"}, + "subject_types_supported": []string{"public"}, + "id_token_signing_alg_values_supported": []string{"RS256"}, + "code_challenge_methods_supported": []string{"S256"}, + "token_endpoint_auth_methods_supported": []string{"client_secret_basic", "client_secret_post"}, + "scopes_supported": []string{"openid", "profile", "email", "groups"}, + } + writeJSON(w, http.StatusOK, doc) +} + +func (i *IdP) handleKeys(w http.ResponseWriter, _ *http.Request) { + n := base64.RawURLEncoding.EncodeToString(i.key.N.Bytes()) + e := base64.RawURLEncoding.EncodeToString(big.NewInt(int64(i.key.E)).Bytes()) + + jwks := map[string]any{ + "keys": []map[string]any{ + { + "kty": "RSA", + "use": "sig", + "alg": "RS256", + "kid": KeyID, + "n": n, + "e": e, + }, + }, + } + writeJSON(w, http.StatusOK, jwks) +} + +func (i *IdP) handleAuthorize(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + redirectURI := q.Get("redirect_uri") + state := q.Get("state") + + i.mu.Lock() + errCode, errDesc := i.authErrCode, i.authErrDesc + i.mu.Unlock() + + if errCode != "" { + redirectWithError(w, r, redirectURI, state, errCode, errDesc) + return + } + + valid := q.Get("response_type") == "code" && + q.Get("client_id") == i.ClientID && + redirectURI != "" && + scopeContains(q.Get("scope"), "openid") && + q.Get("code_challenge_method") == "S256" && + q.Get("code_challenge") != "" + + if !valid { + redirectWithError(w, r, redirectURI, state, "invalid_request", "missing or invalid authorization request parameter") + return + } + + code, err := randomToken(32) + if err != nil { + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + + i.mu.Lock() + i.codes[code] = authState{ + challenge: q.Get("code_challenge"), + nonce: q.Get("nonce"), + redirectURI: redirectURI, + user: i.user, + } + i.mu.Unlock() + + dest, err := url.Parse(redirectURI) + if err != nil { + http.Error(w, "invalid redirect_uri", http.StatusBadRequest) + return + } + dq := dest.Query() + dq.Set("code", code) + dq.Set("state", state) + dest.RawQuery = dq.Encode() + + http.Redirect(w, r, dest.String(), http.StatusFound) // #nosec G710 -- test identity provider on loopback, never linked into the binary; redirects to the client-supplied redirect_uri as a real IdP authorize endpoint does +} + +func redirectWithError(w http.ResponseWriter, r *http.Request, redirectURI, state, code, description string) { + if redirectURI == "" { + http.Error(w, code, http.StatusBadRequest) + return + } + dest, err := url.Parse(redirectURI) + if err != nil { + http.Error(w, "invalid redirect_uri", http.StatusBadRequest) + return + } + q := dest.Query() + q.Set("error", code) + q.Set("error_description", description) + q.Set("state", state) + dest.RawQuery = q.Encode() + http.Redirect(w, r, dest.String(), http.StatusFound) // #nosec G710 -- test identity provider on loopback, never linked into the binary; redirects to the client-supplied redirect_uri as a real IdP authorize endpoint does +} + +func (i *IdP) handleToken(w http.ResponseWriter, r *http.Request) { + i.mu.Lock() + i.tokenRequests++ + i.mu.Unlock() + + if err := r.ParseForm(); err != nil { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid_request"}) + return + } + + if !i.authenticateClient(r) { + writeJSON(w, http.StatusUnauthorized, map[string]string{"error": "invalid_client"}) + return + } + + if r.PostForm.Get("grant_type") != "authorization_code" { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "unsupported_grant_type"}) + return + } + + code := r.PostForm.Get("code") + + i.mu.Lock() + state, ok := i.codes[code] + delete(i.codes, code) // single use, even if a later check fails + i.mu.Unlock() + + if !ok { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid_grant"}) + return + } + + if r.PostForm.Get("redirect_uri") != state.redirectURI { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid_grant"}) + return + } + + sum := sha256.Sum256([]byte(r.PostForm.Get("code_verifier"))) + computed := base64.RawURLEncoding.EncodeToString(sum[:]) + if subtle.ConstantTimeCompare([]byte(computed), []byte(state.challenge)) != 1 { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid_grant"}) + return + } + + idToken, err := i.signIDToken(state) + if err != nil { + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + + accessToken, err := randomToken(32) + if err != nil { + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + + i.mu.Lock() + i.lastIDToken = idToken + i.mu.Unlock() + + writeJSON(w, http.StatusOK, map[string]any{ + "access_token": accessToken, + "token_type": "Bearer", + "expires_in": 300, + "id_token": idToken, + }) +} + +// authenticateClient checks the client_id/client_secret pair, either from +// HTTP Basic auth or from the request body, with a constant-time +// comparison of the secret. +func (i *IdP) authenticateClient(r *http.Request) bool { + var clientID, clientSecret string + + if basicID, basicSecret, ok := r.BasicAuth(); ok { + unescapedID, err := url.QueryUnescape(basicID) + if err != nil { + return false + } + unescapedSecret, err := url.QueryUnescape(basicSecret) + if err != nil { + return false + } + clientID, clientSecret = unescapedID, unescapedSecret + } else { + clientID = r.PostForm.Get("client_id") + clientSecret = r.PostForm.Get("client_secret") + } + + if subtle.ConstantTimeCompare([]byte(clientID), []byte(i.ClientID)) != 1 { + return false + } + return subtle.ConstantTimeCompare([]byte(clientSecret), []byte(i.ClientSecret)) == 1 +} + +func (i *IdP) signIDToken(state authState) (string, error) { + now := time.Now() + + i.mu.Lock() + signing := i.signing + mutator := i.mutator + i.mu.Unlock() + + claims := map[string]any{ + "iss": i.URL, + "sub": state.user.Subject, + "aud": i.ClientID, + "iat": now.Unix(), + "exp": now.Add(5 * time.Minute).Unix(), + } + if state.nonce != "" { + claims["nonce"] = state.nonce + } + for k, v := range state.user.Claims { + claims[k] = v + } + if mutator != nil { + mutator(claims) + } + + switch signing { + case SignNone: + token := jwt.NewWithClaims(jwt.SigningMethodNone, jwt.MapClaims(claims)) + token.Header["kid"] = KeyID + return token.SignedString(jwt.UnsafeAllowNoneSignatureType) + case SignForeignKey: + token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(claims)) + token.Header["kid"] = KeyID + return token.SignedString(i.foreign) + case SignHS256PublicKey: + der, err := x509.MarshalPKIXPublicKey(&i.key.PublicKey) + if err != nil { + return "", fmt.Errorf("marshal public key: %w", err) + } + secret := pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der}) + token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims(claims)) + token.Header["kid"] = KeyID + return token.SignedString(secret) + default: + token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(claims)) + token.Header["kid"] = KeyID + return token.SignedString(i.key) + } +} + +func writeJSON(w http.ResponseWriter, status int, v any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(v) +} + +func scopeContains(scope, want string) bool { + for _, s := range strings.Fields(scope) { + if s == want { + return true + } + } + return false +} + +func randomToken(n int) (string, error) { + b := make([]byte, n) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("generate random token: %w", err) + } + return base64.RawURLEncoding.EncodeToString(b), nil +} diff --git a/internal/auth/sso/ssotest/idp_test.go b/internal/auth/sso/ssotest/idp_test.go new file mode 100644 index 0000000..38e2d6d --- /dev/null +++ b/internal/auth/sso/ssotest/idp_test.go @@ -0,0 +1,386 @@ +package ssotest + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "net/http" + "net/url" + "strings" + "testing" + + "github.com/golang-jwt/jwt/v5" +) + +// randomVerifier returns a 43-character base64url string suitable as a PKCE +// code_verifier. +func randomVerifier(t testing.TB) string { + t.Helper() + b := make([]byte, 32) + if _, err := rand.Read(b); err != nil { + t.Fatalf("rand.Read: %v", err) + } + v := base64.RawURLEncoding.EncodeToString(b) + if len(v) != 43 { + t.Fatalf("verifier length = %d, want 43", len(v)) + } + return v +} + +func challengeFor(verifier string) string { + sum := sha256.Sum256([]byte(verifier)) + return base64.RawURLEncoding.EncodeToString(sum[:]) +} + +// getJSON issues a GET request via http.NewRequest+Do (rather than http.Get +// with a variable URL) and decodes the JSON response body into v. +func getJSON(t testing.TB, rawURL string, v any) { + t.Helper() + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, rawURL, nil) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET %s: %v", rawURL, err) + } + defer res.Body.Close() + if err := json.NewDecoder(res.Body).Decode(v); err != nil { + t.Fatalf("decode %s: %v", rawURL, err) + } +} + +func TestDiscoveryAndKeys(t *testing.T) { + idp := New(t) + + var doc struct { + Issuer string `json:"issuer"` + JWKSURI string `json:"jwks_uri"` + } + getJSON(t, idp.URL+"/.well-known/openid-configuration", &doc) + if doc.Issuer != idp.URL { + t.Fatalf("issuer = %q, want %q", doc.Issuer, idp.URL) + } + + var jwks struct { + Keys []struct { + Kid string `json:"kid"` + } `json:"keys"` + } + getJSON(t, doc.JWKSURI, &jwks) + if len(jwks.Keys) != 1 || jwks.Keys[0].Kid != KeyID { + t.Fatalf("jwks keys = %+v, want one key with kid %q", jwks.Keys, KeyID) + } +} + +// authorizeAndToken drives the full authorization code flow with PKCE and +// returns the token endpoint's raw JSON response and status code. +func authorizeAndToken(t testing.TB, idp *IdP, verifier string, extra url.Values) (*http.Response, map[string]any) { + t.Helper() + + challenge := challengeFor(verifier) + + authURL := idp.URL + "/authorize?" + url.Values{ + "client_id": {idp.ClientID}, + "redirect_uri": {"http://tracker.test/cb"}, + "response_type": {"code"}, + "scope": {"openid profile"}, + "state": {"st"}, + "nonce": {"nn"}, + "code_challenge": {challenge}, + "code_challenge_method": {"S256"}, + }.Encode() + + loc := idp.Authorize(t, authURL) + + if loc.Query().Get("state") != "st" { + t.Fatalf("state = %q, want %q", loc.Query().Get("state"), "st") + } + code := loc.Query().Get("code") + if code == "" { + t.Fatalf("no code in redirect: %s", loc) + } + + form := url.Values{ + "grant_type": {"authorization_code"}, + "code": {code}, + "redirect_uri": {"http://tracker.test/cb"}, + } + if extra.Get("code_verifier") != "" { + form.Set("code_verifier", extra.Get("code_verifier")) + } else { + form.Set("code_verifier", verifier) + } + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, idp.URL+"/token", strings.NewReader(form.Encode())) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + if extra.Get("client_secret") != "" { + req.SetBasicAuth(idp.ClientID, extra.Get("client_secret")) + } else { + req.SetBasicAuth(idp.ClientID, idp.ClientSecret) + } + + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST token: %v", err) + } + defer res.Body.Close() + + var body map[string]any + if res.Header.Get("Content-Type") != "" { + if err := json.NewDecoder(res.Body).Decode(&body); err != nil { + t.Fatalf("decode token response: %v", err) + } + } + return res, body +} + +func TestAuthorizationCodeFlowWithPKCE(t *testing.T) { + idp := New(t) + verifier := randomVerifier(t) + + res, body := authorizeAndToken(t, idp, verifier, url.Values{}) + if res.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200, body = %v", res.StatusCode, body) + } + + idToken, _ := body["id_token"].(string) + if idToken == "" { + t.Fatalf("no id_token in response: %v", body) + } + + claims := jwt.MapClaims{} + _, err := jwt.ParseWithClaims(idToken, claims, func(*jwt.Token) (interface{}, error) { + return &idp.key.PublicKey, nil + }) + if err != nil { + t.Fatalf("parse id_token: %v", err) + } + + if claims["iss"] != idp.URL { + t.Errorf("iss = %v, want %v", claims["iss"], idp.URL) + } + if claims["aud"] != idp.ClientID { + t.Errorf("aud = %v, want %v", claims["aud"], idp.ClientID) + } + if claims["nonce"] != "nn" { + t.Errorf("nonce = %v, want %q", claims["nonce"], "nn") + } + if claims["preferred_username"] != "alice" { + t.Errorf("preferred_username = %v, want %q", claims["preferred_username"], "alice") + } + + if got := idp.TokenRequests(); got != 1 { + t.Errorf("TokenRequests() = %d, want 1", got) + } +} + +func TestTokenRejectsWrongVerifier(t *testing.T) { + idp := New(t) + verifier := randomVerifier(t) + + res, body := authorizeAndToken(t, idp, verifier, url.Values{"code_verifier": {randomVerifier(t)}}) + if res.StatusCode != http.StatusBadRequest { + t.Fatalf("status = %d, want 400, body = %v", res.StatusCode, body) + } + if body["error"] != "invalid_grant" { + t.Errorf("error = %v, want %q", body["error"], "invalid_grant") + } +} + +func TestCodeIsSingleUse(t *testing.T) { + idp := New(t) + verifier := randomVerifier(t) + challenge := challengeFor(verifier) + + authURL := idp.URL + "/authorize?" + url.Values{ + "client_id": {idp.ClientID}, + "redirect_uri": {"http://tracker.test/cb"}, + "response_type": {"code"}, + "scope": {"openid"}, + "state": {"st"}, + "code_challenge": {challenge}, + "code_challenge_method": {"S256"}, + }.Encode() + loc := idp.Authorize(t, authURL) + code := loc.Query().Get("code") + + form := url.Values{ + "grant_type": {"authorization_code"}, + "code": {code}, + "redirect_uri": {"http://tracker.test/cb"}, + "code_verifier": {verifier}, + } + post := func() *http.Response { + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, idp.URL+"/token", strings.NewReader(form.Encode())) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + req.SetBasicAuth(idp.ClientID, idp.ClientSecret) + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST token: %v", err) + } + return res + } + + first := post() + first.Body.Close() + if first.StatusCode != http.StatusOK { + t.Fatalf("first status = %d, want 200", first.StatusCode) + } + + second := post() + defer second.Body.Close() + if second.StatusCode != http.StatusBadRequest { + t.Fatalf("second status = %d, want 400", second.StatusCode) + } + var body map[string]any + if err := json.NewDecoder(second.Body).Decode(&body); err != nil { + t.Fatalf("decode: %v", err) + } + if body["error"] != "invalid_grant" { + t.Errorf("error = %v, want %q", body["error"], "invalid_grant") + } +} + +func TestTokenRejectsWrongClientSecret(t *testing.T) { + idp := New(t) + verifier := randomVerifier(t) + + res, body := authorizeAndToken(t, idp, verifier, url.Values{"client_secret": {"wrong-secret"}}) + if res.StatusCode != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401, body = %v", res.StatusCode, body) + } + if body["error"] != "invalid_client" { + t.Errorf("error = %v, want %q", body["error"], "invalid_client") + } +} + +func TestAuthorizeRequiresS256(t *testing.T) { + idp := New(t) + + authURL := idp.URL + "/authorize?" + url.Values{ + "client_id": {idp.ClientID}, + "redirect_uri": {"http://tracker.test/cb"}, + "response_type": {"code"}, + "scope": {"openid"}, + "state": {"st"}, + }.Encode() + + loc := idp.Authorize(t, authURL) + if loc.Query().Get("error") != "invalid_request" { + t.Errorf("error = %v, want %q", loc.Query().Get("error"), "invalid_request") + } +} + +func TestAuthorizeError(t *testing.T) { + idp := New(t) + idp.SetAuthorizeError("access_denied", "nope") + + authURL := idp.URL + "/authorize?" + url.Values{ + "client_id": {idp.ClientID}, + "redirect_uri": {"http://tracker.test/cb"}, + "response_type": {"code"}, + "scope": {"openid"}, + "state": {"st"}, + "code_challenge": {"c"}, + "code_challenge_method": {"S256"}, + }.Encode() + + loc := idp.Authorize(t, authURL) + if loc.Query().Get("error") != "access_denied" { + t.Errorf("error = %v, want %q", loc.Query().Get("error"), "access_denied") + } + if loc.Query().Get("state") != "st" { + t.Errorf("state = %v, want %q", loc.Query().Get("state"), "st") + } +} + +func TestSigningModes(t *testing.T) { + idp := New(t) + verifier := randomVerifier(t) + + idp.SetSigning(SignNone) + _, body := authorizeAndToken(t, idp, verifier, url.Values{}) + idToken, _ := body["id_token"].(string) + if idToken == "" { + t.Fatalf("no id_token: %v", body) + } + parts := strings.Split(idToken, ".") + if len(parts) != 3 { + t.Fatalf("id_token has %d parts, want 3", len(parts)) + } + headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + t.Fatalf("decode header: %v", err) + } + var header struct { + Alg string `json:"alg"` + } + if err := json.Unmarshal(headerJSON, &header); err != nil { + t.Fatalf("unmarshal header: %v", err) + } + if header.Alg != "none" { + t.Errorf("alg = %q, want %q", header.Alg, "none") + } + + idp.SetSigning(SignForeignKey) + _, body2 := authorizeAndToken(t, idp, verifier, url.Values{}) + idToken2, _ := body2["id_token"].(string) + if idToken2 == "" { + t.Fatalf("no id_token: %v", body2) + } + _, err = jwt.Parse(idToken2, func(*jwt.Token) (interface{}, error) { + return &idp.key.PublicKey, nil + }) + if err == nil { + t.Fatal("expected signature verification to fail for a foreign key token") + } +} + +func TestMutatorAndUser(t *testing.T) { + idp := New(t) + idp.SetUser(User{ + Subject: "user-42", + Claims: map[string]any{ + "preferred_username": "bob", + "email": "bob@example.com", + "name": "Bob Example", + "groups": []string{"admins"}, + }, + }) + idp.SetTokenMutator(func(claims map[string]any) { + claims["aud"] = "other" + }) + + verifier := randomVerifier(t) + res, body := authorizeAndToken(t, idp, verifier, url.Values{}) + if res.StatusCode != http.StatusOK { + t.Fatalf("status = %d, want 200, body = %v", res.StatusCode, body) + } + idToken, _ := body["id_token"].(string) + + claims := jwt.MapClaims{} + _, err := jwt.ParseWithClaims(idToken, claims, func(*jwt.Token) (interface{}, error) { + return &idp.key.PublicKey, nil + }) + if err != nil { + t.Fatalf("parse id_token: %v", err) + } + if claims["sub"] != "user-42" { + t.Errorf("sub = %v, want %q", claims["sub"], "user-42") + } + if claims["preferred_username"] != "bob" { + t.Errorf("preferred_username = %v, want %q", claims["preferred_username"], "bob") + } + if claims["aud"] != "other" { + t.Errorf("aud = %v, want %q (mutator applied last)", claims["aud"], "other") + } +} diff --git a/internal/auth/sso/transaction.go b/internal/auth/sso/transaction.go new file mode 100644 index 0000000..fecea2b --- /dev/null +++ b/internal/auth/sso/transaction.go @@ -0,0 +1,237 @@ +package sso + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/hkdf" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + "golang.org/x/oauth2" + + "github.com/bananaops/tracker/internal/auth" +) + +const ( + // TransactionCookieName carries the sealed OIDC login transaction. + TransactionCookieName = "tracker_oidc" + // TransactionTTL is how long a transaction is accepted after issuance. + TransactionTTL = 10 * time.Minute + + // hkdfInfo binds the derived key to this specific use, separating it + // from any other key derived from the same session secret. + hkdfInfo = "tracker oidc transaction v1" + // maxTransactionCookieLength bounds Decode's input before any decoding + // work happens. + maxTransactionCookieLength = 2048 + // futureSkew is how far into the future IssuedAt may be before a + // transaction is rejected as invalid, to tolerate minor clock drift. + futureSkew = time.Minute + + randomTokenBytes = 32 +) + +var ( + // ErrTransactionInvalid covers every rejection except a well-formed, + // correctly decrypted transaction that is simply too old: tampering, + // truncation, wrong key, wrong AAD, malformed input, and an issuance + // timestamp too far in the future. Kept generic so decoding never gives + // an attacker an oracle. + ErrTransactionInvalid = errors.New("invalid oidc transaction") + // ErrTransactionExpired means the transaction decrypted and parsed + // correctly but is older than TransactionTTL. + ErrTransactionExpired = errors.New("expired oidc transaction") + // ErrTransactionTooLarge means the encoded value would exceed the cookie + // size cap. Callers must retry with Redirect set to "/", which always fits. + ErrTransactionTooLarge = errors.New("oidc transaction too large for a cookie") +) + +// Transaction is the state carried across the redirect to the identity +// provider and back: the CSRF state, the id_token nonce, the PKCE code +// verifier, where to send the browser after login, and when it was issued. +// JSON tags are kept short since the marshaled form is encrypted, not +// displayed. +type Transaction struct { + State string `json:"s"` + Nonce string `json:"n"` + Verifier string `json:"v"` + Redirect string `json:"r"` + IssuedAt int64 `json:"t"` +} + +// NewTransaction builds a fresh transaction: random state and nonce, a PKCE +// code verifier, and redirect sanitized through SafeRedirect. +func NewTransaction(redirect string, now time.Time) (Transaction, error) { + state, err := randomToken() + if err != nil { + return Transaction{}, fmt.Errorf("generate state: %w", err) + } + nonce, err := randomToken() + if err != nil { + return Transaction{}, fmt.Errorf("generate nonce: %w", err) + } + return Transaction{ + State: state, + Nonce: nonce, + Verifier: oauth2.GenerateVerifier(), + Redirect: SafeRedirect(redirect), + IssuedAt: now.Unix(), + }, nil +} + +// randomToken returns a 32 byte crypto/rand value, base64url encoded +// without padding. +func randomToken() (string, error) { + b := make([]byte, randomTokenBytes) + if _, err := rand.Read(b); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(b), nil +} + +// TransactionCodec seals and opens the transaction cookie with AES-256-GCM, +// under a key derived from the session secret so no new secret needs to be +// provisioned. +type TransactionCodec struct { + // Now is overridable in tests. + Now func() time.Time + + aead cipher.AEAD +} + +// NewTransactionCodec derives the encryption key from sessionSecret with +// HKDF-SHA256. sessionSecret must be at least auth.SessionSecretLength +// bytes, the same requirement as the session HMAC secret. +func NewTransactionCodec(sessionSecret []byte) (*TransactionCodec, error) { + if len(sessionSecret) < auth.SessionSecretLength { + return nil, fmt.Errorf("session secret must be at least %d bytes", auth.SessionSecretLength) + } + + key, err := hkdf.Key(sha256.New, sessionSecret, nil, hkdfInfo, 32) + if err != nil { + return nil, fmt.Errorf("derive transaction key: %w", err) + } + block, err := aes.NewCipher(key) + if err != nil { + return nil, fmt.Errorf("build cipher: %w", err) + } + aead, err := cipher.NewGCM(block) + if err != nil { + return nil, fmt.Errorf("build AEAD: %w", err) + } + + return &TransactionCodec{Now: time.Now, aead: aead}, nil +} + +// Encode seals t into an opaque, base64url value suitable for a cookie. +// It returns ErrTransactionTooLarge when the value would exceed the size Decode +// accepts; callers must then retry with Redirect "/". +// The cookie name is bound in as associated data, so a value cannot be +// replayed under a different cookie. +func (c *TransactionCodec) Encode(t Transaction) (string, error) { + // HTML escaping is off: json.Marshal would expand <, > and & to 6 bytes + // each, inflating a long redirect past the cookie size limit. + var buf bytes.Buffer + enc := json.NewEncoder(&buf) + enc.SetEscapeHTML(false) + if err := enc.Encode(t); err != nil { + return "", fmt.Errorf("marshal transaction: %w", err) + } + plain := bytes.TrimSuffix(buf.Bytes(), []byte("\n")) + + nonce := make([]byte, c.aead.NonceSize()) + if _, err := rand.Read(nonce); err != nil { + return "", fmt.Errorf("generate nonce: %w", err) + } + + sealed := c.aead.Seal(nonce, nonce, plain, []byte(TransactionCookieName)) + encoded := base64.RawURLEncoding.EncodeToString(sealed) + if len(encoded) > maxTransactionCookieLength { + return "", ErrTransactionTooLarge + } + return encoded, nil +} + +// Decode opens a value produced by Encode. Every rejection short of a +// verified-but-too-old transaction collapses to ErrTransactionInvalid, with +// no detail that would let an attacker distinguish tampering from a bad key +// from a malformed value; nothing decrypted or partially decrypted is ever +// included in the error. +func (c *TransactionCodec) Decode(value string) (Transaction, error) { + if value == "" || len(value) > maxTransactionCookieLength { + return Transaction{}, ErrTransactionInvalid + } + + sealed, err := base64.RawURLEncoding.DecodeString(value) + if err != nil { + return Transaction{}, ErrTransactionInvalid + } + + nonceSize := c.aead.NonceSize() + if len(sealed) < nonceSize+c.aead.Overhead() { + return Transaction{}, ErrTransactionInvalid + } + nonce, ciphertext := sealed[:nonceSize], sealed[nonceSize:] + + plain, err := c.aead.Open(nil, nonce, ciphertext, []byte(TransactionCookieName)) + if err != nil { + return Transaction{}, ErrTransactionInvalid + } + + var t Transaction + if err := json.Unmarshal(plain, &t); err != nil { + return Transaction{}, ErrTransactionInvalid + } + + if t.State == "" || t.Nonce == "" || t.Verifier == "" { + return Transaction{}, ErrTransactionInvalid + } + + issued := time.Unix(t.IssuedAt, 0) + now := c.Now() + if issued.After(now.Add(futureSkew)) { + return Transaction{}, ErrTransactionInvalid + } + if now.Sub(issued) > TransactionTTL { + return Transaction{}, ErrTransactionExpired + } + + return t, nil +} + +// TransactionCookie builds the cookie carrying a sealed transaction value. +// Scoped to the callback path since only the callback handler needs it. +func TransactionCookie(value string, secure bool) *http.Cookie { + // Lax, not Strict: the identity provider sends the browser back with a + // top level cross-site GET, which Strict would strip the cookie from. + return &http.Cookie{ // #nosec G124 -- Secure is configuration driven, HttpOnly and SameSite are set + Name: TransactionCookieName, + Value: value, + Path: auth.OIDCCallbackPath, + HttpOnly: true, + Secure: secure, + SameSite: http.SameSiteLaxMode, + MaxAge: int(TransactionTTL.Seconds()), + } +} + +// ClearTransactionCookie builds the cookie that removes the transaction. +func ClearTransactionCookie(secure bool) *http.Cookie { + return &http.Cookie{ // #nosec G124 -- Secure is configuration driven, HttpOnly and SameSite are set + Name: TransactionCookieName, + Value: "", + Path: auth.OIDCCallbackPath, + HttpOnly: true, + Secure: secure, + SameSite: http.SameSiteLaxMode, + Expires: time.Unix(0, 0), + MaxAge: -1, + } +} diff --git a/internal/auth/sso/transaction_test.go b/internal/auth/sso/transaction_test.go new file mode 100644 index 0000000..35fd6fd --- /dev/null +++ b/internal/auth/sso/transaction_test.go @@ -0,0 +1,336 @@ +package sso + +import ( + "bytes" + "encoding/base64" + "errors" + "net/http" + "regexp" + "strings" + "testing" + "time" + + "github.com/bananaops/tracker/internal/auth" +) + +var base64URLPattern = regexp.MustCompile(`^[A-Za-z0-9_-]+$`) + +func testSecret(b byte) []byte { + return bytes.Repeat([]byte{b}, auth.SessionSecretLength) +} + +func TestNewTransaction(t *testing.T) { + now := time.Now() + + tx1, err := NewTransaction("/locks", now) + if err != nil { + t.Fatalf("NewTransaction: %v", err) + } + tx2, err := NewTransaction("/locks", now) + if err != nil { + t.Fatalf("NewTransaction: %v", err) + } + + for _, tc := range []struct { + name string + v string + }{ + {"tx1 state", tx1.State}, + {"tx1 nonce", tx1.Nonce}, + {"tx1 verifier", tx1.Verifier}, + {"tx2 state", tx2.State}, + {"tx2 nonce", tx2.Nonce}, + {"tx2 verifier", tx2.Verifier}, + } { + if len(tc.v) != 43 { + t.Fatalf("%s length = %d, want 43", tc.name, len(tc.v)) + } + if !base64URLPattern.MatchString(tc.v) { + t.Fatalf("%s = %q, want base64url characters only", tc.name, tc.v) + } + } + + if tx1.State == tx2.State { + t.Fatal("State is the same across two calls") + } + if tx1.Nonce == tx2.Nonce { + t.Fatal("Nonce is the same across two calls") + } + + if tx1.IssuedAt != now.Unix() { + t.Fatalf("IssuedAt = %d, want %d", tx1.IssuedAt, now.Unix()) + } + + txRedirect, err := NewTransaction("//evil", now) + if err != nil { + t.Fatalf("NewTransaction: %v", err) + } + if txRedirect.Redirect != "/" { + t.Fatalf("Redirect = %q, want /", txRedirect.Redirect) + } +} + +func TestCodecRoundTrip(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + tx := Transaction{ + State: "state-value-0123456789012345678901234", + Nonce: "nonce-value-0123456789012345678901234", + Verifier: "verifier-value-01234567890123456789012", + Redirect: "/locks", + IssuedAt: time.Now().Unix(), + } + + encoded, err := codec.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + + for _, secret := range []string{tx.State, tx.Nonce, tx.Verifier} { + if strings.Contains(encoded, secret) { + t.Fatalf("encoded value contains a cleartext secret %q", secret) + } + std := base64.StdEncoding.EncodeToString([]byte(secret)) + if strings.Contains(encoded, std) { + t.Fatalf("encoded value contains a standard-base64 secret %q", std) + } + } + + got, err := codec.Decode(encoded) + if err != nil { + t.Fatalf("Decode: %v", err) + } + if got != tx { + t.Fatalf("Decode = %+v, want %+v", got, tx) + } +} + +func TestCodecRejectsTampering(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + tx := Transaction{ + State: "state-value-0123456789012345678901234", + Nonce: "nonce-value-0123456789012345678901234", + Verifier: "verifier-value-01234567890123456789012", + Redirect: "/locks", + IssuedAt: time.Now().Unix(), + } + encoded, err := codec.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + + mutateMiddle := func(s string) string { + mid := len(s) / 2 + b := []byte(s) + if b[mid] == 'A' { + b[mid] = 'B' + } else { + b[mid] = 'A' + } + return string(b) + } + + cases := map[string]string{ + "modified middle char": mutateMiddle(encoded), + "truncated": encoded[:len(encoded)-4], + "appended char": encoded + "A", + "empty": "", + "too long": strings.Repeat("A", 2049), + "invalid base64": "not-valid-!!!base64!!!", + } + + for name, value := range cases { + t.Run(name, func(t *testing.T) { + _, err := codec.Decode(value) + if !errors.Is(err, ErrTransactionInvalid) { + t.Fatalf("Decode(%s) err = %v, want ErrTransactionInvalid", name, err) + } + }) + } +} + +func TestCodecRejectsOtherSecret(t *testing.T) { + encoder, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + decoder, err := NewTransactionCodec(testSecret(2)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + tx := Transaction{ + State: "state-value-0123456789012345678901234", + Nonce: "nonce-value-0123456789012345678901234", + Verifier: "verifier-value-01234567890123456789012", + Redirect: "/locks", + IssuedAt: time.Now().Unix(), + } + encoded, err := encoder.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + + if _, err := decoder.Decode(encoded); !errors.Is(err, ErrTransactionInvalid) { + t.Fatalf("Decode with the wrong secret err = %v, want ErrTransactionInvalid", err) + } +} + +func TestCodecExpiry(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + issuedAt := time.Now() + tx := Transaction{ + State: "state-value-0123456789012345678901234", + Nonce: "nonce-value-0123456789012345678901234", + Verifier: "verifier-value-01234567890123456789012", + Redirect: "/locks", + IssuedAt: issuedAt.Unix(), + } + encoded, err := codec.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + + codec.Now = func() time.Time { return issuedAt.Add(9*time.Minute + 59*time.Second) } + if _, err := codec.Decode(encoded); err != nil { + t.Fatalf("Decode within TTL: %v, want nil", err) + } + + codec.Now = func() time.Time { return issuedAt.Add(10*time.Minute + time.Second) } + if _, err := codec.Decode(encoded); !errors.Is(err, ErrTransactionExpired) { + t.Fatalf("Decode past TTL err = %v, want ErrTransactionExpired", err) + } + + codec.Now = func() time.Time { return issuedAt.Add(-2 * time.Minute) } + if _, err := codec.Decode(encoded); !errors.Is(err, ErrTransactionInvalid) { + t.Fatalf("Decode with issuance in the future err = %v, want ErrTransactionInvalid", err) + } +} + +func TestCodecRejectsIncomplete(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + tx := Transaction{ + State: "state-value-0123456789012345678901234", + Nonce: "nonce-value-0123456789012345678901234", + Verifier: "", + Redirect: "/locks", + IssuedAt: time.Now().Unix(), + } + encoded, err := codec.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + + if _, err := codec.Decode(encoded); !errors.Is(err, ErrTransactionInvalid) { + t.Fatalf("Decode with no verifier err = %v, want ErrTransactionInvalid", err) + } +} + +func TestNewTransactionCodecShortSecret(t *testing.T) { + _, err := NewTransactionCodec(bytes.Repeat([]byte{1}, auth.SessionSecretLength-1)) + if err == nil { + t.Fatal("NewTransactionCodec with a short secret succeeded, want an error") + } +} + +func TestTransactionCookieAttributes(t *testing.T) { + cookie := TransactionCookie("v", true) + if cookie.Name != TransactionCookieName { + t.Fatalf("Name = %q, want %q", cookie.Name, TransactionCookieName) + } + if cookie.Path != auth.OIDCCallbackPath { + t.Fatalf("Path = %q, want %q", cookie.Path, auth.OIDCCallbackPath) + } + if !cookie.HttpOnly { + t.Fatal("HttpOnly = false, want true") + } + if !cookie.Secure { + t.Fatal("Secure = false, want true") + } + if cookie.SameSite != http.SameSiteLaxMode { + t.Fatalf("SameSite = %v, want SameSiteLaxMode", cookie.SameSite) + } + if cookie.MaxAge != 600 { + t.Fatalf("MaxAge = %d, want 600", cookie.MaxAge) + } + + insecure := TransactionCookie("v", false) + if insecure.Secure { + t.Fatal("Secure = true, want false") + } + + clear := ClearTransactionCookie(true) + if clear.Name != TransactionCookieName { + t.Fatalf("Name = %q, want %q", clear.Name, TransactionCookieName) + } + if clear.Path != auth.OIDCCallbackPath { + t.Fatalf("Path = %q, want %q", clear.Path, auth.OIDCCallbackPath) + } + if clear.MaxAge != -1 { + t.Fatalf("MaxAge = %d, want -1", clear.MaxAge) + } + if clear.Value != "" { + t.Fatalf("Value = %q, want empty", clear.Value) + } +} + +func TestCodecWorstCaseRedirect(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + + for name, filler := range map[string]string{"html": "<&>", "non-ascii": "é"} { + t.Run(name, func(t *testing.T) { + redirect := SafeRedirect("/" + strings.Repeat(filler, 1023/len(filler))) + if redirect == "/" { + t.Fatal("redirect rejected by SafeRedirect, test is meaningless") + } + tx, err := NewTransaction(redirect, time.Now()) + if err != nil { + t.Fatalf("NewTransaction: %v", err) + } + encoded, err := codec.Encode(tx) + if err != nil { + t.Fatalf("Encode: %v", err) + } + if len(encoded) > maxTransactionCookieLength { + t.Fatalf("encoded length = %d, want <= %d", len(encoded), maxTransactionCookieLength) + } + got, err := codec.Decode(encoded) + if err != nil || got != tx { + t.Fatalf("Decode = %+v, %v; want %+v", got, err, tx) + } + }) + } +} + +func TestCodecEncodeTooLarge(t *testing.T) { + codec, err := NewTransactionCodec(testSecret(1)) + if err != nil { + t.Fatalf("NewTransactionCodec: %v", err) + } + tx, err := NewTransaction("/", time.Now()) + if err != nil { + t.Fatalf("NewTransaction: %v", err) + } + tx.Redirect = "/" + strings.Repeat("a", maxTransactionCookieLength) + if _, err := codec.Encode(tx); !errors.Is(err, ErrTransactionTooLarge) { + t.Fatalf("Encode err = %v, want ErrTransactionTooLarge", err) + } +} diff --git a/internal/stores/auth_teams.go b/internal/stores/auth_teams.go index 02f5d5d..0b5f906 100644 --- a/internal/stores/auth_teams.go +++ b/internal/stores/auth_teams.go @@ -123,3 +123,8 @@ func (s *AuthTeamStore) Delete(ctx context.Context, id primitive.ObjectID) error } return nil } + +// ListWithOIDCGroups returns the teams mapped to at least one OIDC group. +func (s *AuthTeamStore) ListWithOIDCGroups(ctx context.Context) ([]*Team, error) { + return s.find(ctx, bson.M{"oidcGroups.0": bson.M{"$exists": true}}) +} diff --git a/internal/stores/auth_teams_test.go b/internal/stores/auth_teams_test.go index 25f3bde..08bec08 100644 --- a/internal/stores/auth_teams_test.go +++ b/internal/stores/auth_teams_test.go @@ -47,3 +47,21 @@ func TestAuthTeamStoreCRUD(t *testing.T) { require.NoError(t, err) assert.Len(t, list, 1) } + +func TestAuthTeamStoreListWithOIDCGroups(t *testing.T) { + db := testDatabase(t) + s := NewAuthTeamStoreFromCollection(db.Collection(authTeamsCollection)) + ctx := context.Background() + + require.NoError(t, s.Create(ctx, &Team{Name: "P", OIDCGroups: []string{"platform-eng"}})) + require.NoError(t, s.Create(ctx, &Team{Name: "N", OIDCGroups: []string{}})) + require.NoError(t, s.Create(ctx, &Team{Name: "Q", OIDCGroups: []string{"ops", "x"}})) + + got, err := s.ListWithOIDCGroups(ctx) + require.NoError(t, err) + names := []string{} + for _, tm := range got { + names = append(names, tm.Name) + } + assert.Equal(t, []string{"P", "Q"}, names) +} diff --git a/internal/stores/auth_users.go b/internal/stores/auth_users.go index e601d67..edbc690 100644 --- a/internal/stores/auth_users.go +++ b/internal/stores/auth_users.go @@ -53,6 +53,28 @@ func (s *AuthUserStore) GetByUsername(ctx context.Context, username string) (*Us return s.findOne(ctx, bson.M{"usernameLower": strings.ToLower(strings.TrimSpace(username))}) } +// GetByOIDCIdentity finds the user bound to an identity provider subject. +func (s *AuthUserStore) GetByOIDCIdentity(ctx context.Context, issuer, subject string) (*User, error) { + return s.findOne(ctx, bson.M{"oidcIssuer": issuer, "oidcSubject": subject}) +} + +// UpdateOIDCProfile refreshes the identity provider fields and the last login time. +func (s *AuthUserStore) UpdateOIDCProfile(ctx context.Context, id primitive.ObjectID, email, displayName string, at time.Time) error { + res, err := s.coll.UpdateByID(ctx, id, bson.M{"$set": bson.M{ + "email": email, + "displayName": displayName, + "lastLoginAt": at, + "updatedAt": time.Now().UTC(), + }}) + if err != nil { + return err + } + if res.MatchedCount == 0 { + return ErrNotFound + } + return nil +} + func (s *AuthUserStore) findOne(ctx context.Context, filter bson.M) (*User, error) { var u User err := s.coll.FindOne(ctx, filter).Decode(&u) @@ -122,3 +144,34 @@ func (s *AuthUserStore) CountEnabledInTeam(ctx context.Context, teamID, excludeU } return s.coll.CountDocuments(ctx, filter) } + +// SyncTeams removes then adds team memberships without touching the others. +// Two targeted updates are needed because MongoDB refuses $pull and $addToSet +// on the same field in one update. Both are idempotent. +func (s *AuthUserStore) SyncTeams(ctx context.Context, id primitive.ObjectID, add, remove []primitive.ObjectID) error { + if len(remove) > 0 { + res, err := s.coll.UpdateByID(ctx, id, bson.M{ + "$pull": bson.M{"teams": bson.M{"$in": remove}}, + "$set": bson.M{"updatedAt": time.Now().UTC()}, + }) + if err != nil { + return err + } + if res.MatchedCount == 0 { + return ErrNotFound + } + } + if len(add) > 0 { + res, err := s.coll.UpdateByID(ctx, id, bson.M{ + "$addToSet": bson.M{"teams": bson.M{"$each": add}}, + "$set": bson.M{"updatedAt": time.Now().UTC()}, + }) + if err != nil { + return err + } + if res.MatchedCount == 0 { + return ErrNotFound + } + } + return nil +} diff --git a/internal/stores/auth_users_test.go b/internal/stores/auth_users_test.go index 1e8bfe3..79b9f4a 100644 --- a/internal/stores/auth_users_test.go +++ b/internal/stores/auth_users_test.go @@ -66,3 +66,63 @@ func TestAuthUserStoreCRUD(t *testing.T) { assert.ErrorIs(t, s.Update(ctx, &User{ID: primitive.NewObjectID(), Username: "ghost"}), ErrNotFound) } + +func TestAuthUserStoreOIDC(t *testing.T) { + db := testDatabase(t) + s := NewAuthUserStoreFromCollection(db.Collection(authUsersCollection)) + ctx := context.Background() + + team := primitive.NewObjectID() + oidcUser := &User{Username: "bob", Source: UserSourceOIDC, OIDCIssuer: "https://idp", OIDCSubject: "sub-1", Teams: []primitive.ObjectID{team}} + require.NoError(t, s.Create(ctx, oidcUser)) + + got, err := s.GetByOIDCIdentity(ctx, "https://idp", "sub-1") + require.NoError(t, err) + assert.Equal(t, oidcUser.ID, got.ID) + _, err = s.GetByOIDCIdentity(ctx, "https://idp", "sub-2") + assert.ErrorIs(t, err, ErrNotFound) + _, err = s.GetByOIDCIdentity(ctx, "https://other", "sub-1") + assert.ErrorIs(t, err, ErrNotFound) + + dup := &User{Username: "bob2", Source: UserSourceOIDC, OIDCIssuer: "https://idp", OIDCSubject: "sub-1"} + assert.ErrorIs(t, s.Create(ctx, dup), ErrAlreadyExists) + + require.NoError(t, s.Create(ctx, &User{Username: "l1", Source: UserSourceLocal, PasswordHash: "x"})) + require.NoError(t, s.Create(ctx, &User{Username: "l2", Source: UserSourceLocal, PasswordHash: "x"})) + + at := time.Now().UTC().Add(time.Hour) + require.NoError(t, s.UpdateOIDCProfile(ctx, oidcUser.ID, "b@x.io", "Bob B", at)) + after, err := s.GetByID(ctx, oidcUser.ID) + require.NoError(t, err) + assert.Equal(t, "b@x.io", after.Email) + assert.Equal(t, "Bob B", after.DisplayName) + require.NotNil(t, after.LastLoginAt) + assert.WithinDuration(t, at, *after.LastLoginAt, time.Second) + assert.True(t, after.UpdatedAt.After(got.UpdatedAt)) + assert.Equal(t, "bob", after.Username) + assert.Equal(t, []primitive.ObjectID{team}, after.Teams) + assert.Equal(t, UserSourceOIDC, after.Source) + assert.Equal(t, got.SessionVersion, after.SessionVersion) + + assert.ErrorIs(t, s.UpdateOIDCProfile(ctx, primitive.NewObjectID(), "a", "b", at), ErrNotFound) +} + +func TestAuthUserStoreSyncTeams(t *testing.T) { + db := testDatabase(t) + s := NewAuthUserStoreFromCollection(db.Collection(authUsersCollection)) + ctx := context.Background() + + a, m, b := primitive.NewObjectID(), primitive.NewObjectID(), primitive.NewObjectID() + u := &User{Username: "sync", Source: UserSourceOIDC, Teams: []primitive.ObjectID{a, m}} + require.NoError(t, s.Create(ctx, u)) + + for range 2 { + require.NoError(t, s.SyncTeams(ctx, u.ID, []primitive.ObjectID{b}, []primitive.ObjectID{a})) + got, err := s.GetByID(ctx, u.ID) + require.NoError(t, err) + assert.ElementsMatch(t, []primitive.ObjectID{m, b}, got.Teams) + } + + require.NoError(t, s.SyncTeams(ctx, u.ID, nil, nil)) + assert.ErrorIs(t, s.SyncTeams(ctx, primitive.NewObjectID(), []primitive.ObjectID{b}, nil), ErrNotFound) +} diff --git a/server/auth.go b/server/auth.go index ae92516..3511911 100644 --- a/server/auth.go +++ b/server/auth.go @@ -43,9 +43,14 @@ func (a *Auth) GetAuthConfig(ctx context.Context, _ *authv1.GetAuthConfigRequest if err := authz.Authorize(ctx); err != nil { return nil, err } + label := "" + if a.cfg.OIDC.Enabled() { + label = a.cfg.OIDC.ButtonLabel + } return &authv1.GetAuthConfigResponse{ LocalLoginEnabled: true, - OidcEnabled: false, + OidcEnabled: a.cfg.OIDC.Enabled(), + OidcButtonLabel: label, AnonymousPermissions: permissionStrings(a.cfg.AnonymousPermissions), DemoMode: a.cfg.DemoMode, }, nil diff --git a/server/auth_http.go b/server/auth_http.go index 56e2bc1..f6ad5ca 100644 --- a/server/auth_http.go +++ b/server/auth_http.go @@ -141,11 +141,16 @@ func (h *AuthHTTP) handleLogin(w http.ResponseWriter, r *http.Request, _ map[str } func (h *AuthHTTP) issueSession(w http.ResponseWriter, user *store.User) error { - token, expires, err := h.sessions.Issue(user.ID.Hex(), user.SessionVersion) + return setSessionCookie(w, h.sessions, h.cfg.CookieSecure, user) +} + +// setSessionCookie issues a session for user and sets the tracker_session cookie. +func setSessionCookie(w http.ResponseWriter, sessions *auth.SessionManager, secure bool, user *store.User) error { + token, expires, err := sessions.Issue(user.ID.Hex(), user.SessionVersion) if err != nil { return err } - http.SetCookie(w, auth.SessionCookie(token, expires, h.cfg.CookieSecure)) + http.SetCookie(w, auth.SessionCookie(token, expires, secure)) return nil } diff --git a/server/auth_oidc.go b/server/auth_oidc.go new file mode 100644 index 0000000..8228397 --- /dev/null +++ b/server/auth_oidc.go @@ -0,0 +1,360 @@ +package server + +import ( + "context" + "crypto/subtle" + "errors" + "fmt" + "html" + "log/slog" + "net/http" + "time" + "unicode/utf8" + + "github.com/bananaops/tracker/internal/auth" + "github.com/bananaops/tracker/internal/auth/authz" + "github.com/bananaops/tracker/internal/auth/identity" + "github.com/bananaops/tracker/internal/auth/sso" + store "github.com/bananaops/tracker/internal/stores" + "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" + "golang.org/x/oauth2" +) + +// Error codes carried by the /login?error= redirect. They are constants: a +// value received from the identity provider never reaches the redirect. +const ( + oidcErrDenied = "oidc_denied" + oidcErrState = "oidc_state" + oidcErrFailed = "oidc_failed" + oidcErrUnavailable = "oidc_unavailable" + + oidcExchangeTimeout = 15 * time.Second + + maxIdPErrorLength = 64 + maxIdPErrorDescriptionLength = 200 + + msgNotProvisioned = "Your identity provider account is not registered in Tracker. Ask a Tracker administrator for access." + msgUserDisabled = "Your Tracker account is disabled. Ask a Tracker administrator." +) + +const oidcRefusalPage = ` +Sign-in refused +

Sign-in refused

%s

Back to sign in

+` + +var errTransactionMissing = errors.New("oidc transaction cookie missing") + +// OIDCHTTP serves the OpenID Connect login and callback routes. +type OIDCHTTP struct { + users *store.AuthUserStore + teams *store.AuthTeamStore + sessions *auth.SessionManager + provider sso.Provider + codec *sso.TransactionCodec + cfg auth.Config + logger *slog.Logger + now func() time.Time +} + +func NewOIDCHTTP(users *store.AuthUserStore, teams *store.AuthTeamStore, sessions *auth.SessionManager, provider sso.Provider, codec *sso.TransactionCodec, cfg auth.Config) *OIDCHTTP { + return &OIDCHTTP{ + users: users, + teams: teams, + sessions: sessions, + provider: provider, + codec: codec, + cfg: cfg, + logger: slog.Default(), + now: time.Now, + } +} + +// Register mounts GET login and callback. Call it only when cfg.OIDC.Enabled(). +func (h *OIDCHTTP) Register(mux *runtime.ServeMux) { + routes := []struct { + path string + handler runtime.HandlerFunc + }{ + {auth.OIDCLoginPath, h.handleLogin}, + {auth.OIDCCallbackPath, h.handleCallback}, + } + for _, r := range routes { + if err := mux.HandlePath(http.MethodGet, r.path, authz.RequireHTTP(auth.PermPublic, r.handler)); err != nil { + h.logger.Error("Failed to register OIDC route", "path", r.path, "error", err) + } + } +} + +func (h *OIDCHTTP) count(result string) { + authz.AuthLogins.WithLabelValues(authz.LoginMethodOIDC, result).Inc() +} + +func (h *OIDCHTTP) handleLogin(w http.ResponseWriter, r *http.Request, _ map[string]string) { + w.Header().Set("Cache-Control", "no-store") + + tx, err := sso.NewTransaction(r.URL.Query().Get("redirect"), h.now()) + if err != nil { + h.logger.Error("auth.login", "method", "oidc", "reason", "transaction_failed", "error", err) + writeJSONError(w, http.StatusInternalServerError, "internal error") + return + } + target, err := h.provider.AuthCodeURL(r.Context(), tx.State, tx.Nonce, tx.Verifier) + if err != nil { + h.logger.Error("auth.login", "method", "oidc", "result", "failure", "reason", "provider_unavailable", "ip", auth.ClientIP(r, h.cfg.TrustProxy)) + h.count(authz.LoginFailure) + h.redirectError(w, r, oidcErrUnavailable) + return + } + value, err := h.codec.Encode(tx) + if errors.Is(err, sso.ErrTransactionTooLarge) { + // The redirect is what made it too big: drop it rather than fail. + tx.Redirect = "/" + value, err = h.codec.Encode(tx) + } + if err != nil { + h.logger.Error("auth.login", "method", "oidc", "reason", "transaction_encode_failed", "error", err) + writeJSONError(w, http.StatusInternalServerError, "internal error") + return + } + http.SetCookie(w, sso.TransactionCookie(value, h.cfg.CookieSecure)) + http.Redirect(w, r, target, http.StatusFound) +} + +func (h *OIDCHTTP) handleCallback(w http.ResponseWriter, r *http.Request, _ map[string]string) { + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Referrer-Policy", "no-referrer") + http.SetCookie(w, sso.ClearTransactionCookie(h.cfg.CookieSecure)) + + ip := auth.ClientIP(r, h.cfg.TrustProxy) + q := r.URL.Query() + tx, txErr := h.readTransaction(r) + + if q.Get("error") != "" { + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "idp_error", + "idp_error", truncate(q.Get("error"), maxIdPErrorLength), + "idp_error_description", truncate(q.Get("error_description"), maxIdPErrorDescriptionLength), + "ip", ip) + if txErr == nil { + h.count(authz.LoginFailure) + } + h.redirectError(w, r, oidcErrDenied) + return + } + if txErr != nil { + reason := "transaction_invalid" + switch { + case errors.Is(txErr, errTransactionMissing): + reason = "transaction_missing" + case errors.Is(txErr, sso.ErrTransactionExpired): + reason = "transaction_expired" + } + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", reason, "ip", ip) + h.redirectError(w, r, oidcErrState) + return + } + if subtle.ConstantTimeCompare([]byte(q.Get("state")), []byte(tx.State)) != 1 { + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "state_mismatch", "ip", ip) + h.count(authz.LoginFailure) + h.redirectError(w, r, oidcErrState) + return + } + code := q.Get("code") + if code == "" { + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "code_missing", "ip", ip) + h.count(authz.LoginFailure) + h.redirectError(w, r, oidcErrFailed) + return + } + + ctx, cancel := context.WithTimeout(r.Context(), oidcExchangeTimeout) + defer cancel() + claims, err := h.provider.Exchange(ctx, code, tx.Verifier, tx.Nonce) + if err != nil { + h.logExchangeError(err, ip) + h.count(authz.LoginFailure) + if errors.Is(err, sso.ErrUnavailable) { + h.redirectError(w, r, oidcErrUnavailable) + return + } + h.redirectError(w, r, oidcErrFailed) + return + } + + if h.cfg.OIDC.TeamSync && claims.GroupsUnexpectedType { + h.logger.Warn("auth.oidc.sync", "method", "oidc", "reason", "groups_claim_unexpected_type", + "claim", h.cfg.OIDC.GroupsClaim, "username", claims.Username, "ip", ip) + } + + // Refuse before any write: resolving the user would create it or refresh + // its profile and last login although the sync is then going to refuse. + if h.cfg.OIDC.TeamSync && !claims.GroupsPresent { + // Read-only lookup: the rule needs the memberships the user holds now. + var existing *store.User + found, err := h.users.GetByOIDCIdentity(ctx, claims.Issuer, claims.Subject) + switch { + case err == nil: + existing = found + case !errors.Is(err, store.ErrNotFound): + h.logger.Error("auth.login", "method", "oidc", "result", "failure", "reason", "resolve_failed", + "issuer", claims.Issuer, "subject", claims.Subject, "ip", ip, "error", err) + h.count(authz.LoginFailure) + h.redirectError(w, r, oidcErrFailed) + return + } + if err := identity.CheckOIDCGroupsClaim(ctx, h.teams, existing, claims.GroupsPresent); err != nil { + h.count(authz.LoginFailure) + h.logSyncFailure(err, claims, ip) + h.redirectError(w, r, oidcErrFailed) + return + } + h.warnGroupsClaimMissingAccepted(ctx, claims, ip) + } + + // Only issuer and subject identify a user; the email claim is data, never + // a key to find or link an account. + user, created, err := identity.ResolveOIDCUser(ctx, h.users, identity.OIDCIdentity{ + Issuer: claims.Issuer, + Subject: claims.Subject, + Username: claims.Username, + Email: claims.Email, + DisplayName: claims.DisplayName, + }, h.cfg.OIDC.UserProvisioning, h.now().UTC()) + if err != nil { + h.count(authz.LoginFailure) + switch { + case errors.Is(err, identity.ErrOIDCNotProvisioned): + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "not_provisioned", + "issuer", claims.Issuer, "subject", claims.Subject, "ip", ip) + writeOIDCRefusal(w, msgNotProvisioned) + case errors.Is(err, identity.ErrOIDCUserDisabled): + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "user_disabled", + "issuer", claims.Issuer, "subject", claims.Subject, "ip", ip) + writeOIDCRefusal(w, msgUserDisabled) + case errors.Is(err, identity.ErrOIDCNoUsername), errors.Is(err, identity.ErrOIDCUsernameExhausted): + h.logger.Warn("auth.login", "method", "oidc", "result", "failure", "reason", "username_unusable", + "issuer", claims.Issuer, "subject", claims.Subject, "ip", ip) + h.redirectError(w, r, oidcErrFailed) + default: + // Includes ErrOIDCInvalidIdentity and ErrOIDCNotOIDCUser: neither + // should happen with a verified token, so they are errors. + h.logger.Error("auth.login", "method", "oidc", "result", "failure", "reason", "resolve_failed", + "issuer", claims.Issuer, "subject", claims.Subject, "ip", ip, "error", err) + h.redirectError(w, r, oidcErrFailed) + } + return + } + + var sync identity.TeamSyncResult + if h.cfg.OIDC.TeamSync { + sync, err = identity.SyncOIDCTeams(ctx, h.users, h.teams, user, claims.Groups, claims.GroupsPresent) + if err != nil { + h.count(authz.LoginFailure) + h.logSyncFailure(err, claims, ip) + h.redirectError(w, r, oidcErrFailed) + return + } + if len(sync.Kept) > 0 { + h.logger.Error("kept the last enabled administrator in Administrators", + "event", "auth.oidc.sync", "username", user.Username, "teams", sync.Kept) + } + } + + if err := setSessionCookie(w, h.sessions, h.cfg.CookieSecure, user); err != nil { + h.logger.Error("auth.login", "method", "oidc", "result", "failure", "reason", "session_failed", "username", user.Username, "error", err) + h.count(authz.LoginFailure) + h.redirectError(w, r, oidcErrFailed) + return + } + h.logger.Info("auth.login", "method", "oidc", "result", "success", "username", user.Username, + "created", created, "teams_added", sync.Added, "teams_removed", sync.Removed, "ip", ip) + h.count(authz.LoginSuccess) + http.Redirect(w, r, sso.SafeRedirect(tx.Redirect), http.StatusSeeOther) +} + +// warnGroupsClaimMissingAccepted logs, once per login, that the groups claim +// is absent, teams are mapped to OIDC groups and the login is accepted anyway +// because the user holds no mapped membership. +func (h *OIDCHTTP) warnGroupsClaimMissingAccepted(ctx context.Context, claims sso.Claims, ip string) { + mapped, err := h.teams.ListWithOIDCGroups(ctx) + if err != nil || len(mapped) == 0 { + return + } + h.logger.Warn("auth.oidc.sync", "method", "oidc", "result", "accepted", "reason", "groups_claim_missing_accepted", + "claim", h.cfg.OIDC.GroupsClaim, "username", claims.Username, "ip", ip) +} + +// logSyncFailure logs a refused or failed team sync. The username is the one +// asserted by the token: the user may not exist yet. +func (h *OIDCHTTP) logSyncFailure(err error, claims sso.Claims, ip string) { + if errors.Is(err, identity.ErrOIDCGroupsClaimMissing) { + h.logger.Error("auth.oidc.sync", "method", "oidc", "result", "failure", "reason", "groups_claim_missing", + "claim", h.cfg.OIDC.GroupsClaim, "username", claims.Username, "ip", ip) + return + } + h.logger.Error("auth.oidc.sync", "method", "oidc", "result", "failure", "reason", "sync_failed", + "username", claims.Username, "ip", ip, "error", err) +} + +// readTransaction decrypts the transaction cookie and re-applies SafeRedirect +// to the stored redirect, whatever the codec accepted. +func (h *OIDCHTTP) readTransaction(r *http.Request) (sso.Transaction, error) { + c, err := r.Cookie(sso.TransactionCookieName) + if err != nil { + return sso.Transaction{}, errTransactionMissing + } + tx, err := h.codec.Decode(c.Value) + if err != nil { + return sso.Transaction{}, fmt.Errorf("decode oidc transaction: %w", err) + } + tx.Redirect = sso.SafeRedirect(tx.Redirect) + return tx, nil +} + +// logExchangeError logs why the code exchange failed without ever printing +// the error text of a token endpoint failure: oauth2 embeds the raw response +// body in it, which may be HTML or carry tokens. +func (h *OIDCHTTP) logExchangeError(err error, ip string) { + attrs := []any{"method", "oidc", "result", "failure", "ip", ip} + var re *oauth2.RetrieveError + switch { + case errors.As(err, &re): + attrs = append(attrs, "reason", "exchange_failed", + "idp_error", truncate(re.ErrorCode, maxIdPErrorLength), + "idp_error_description", truncate(re.ErrorDescription, maxIdPErrorDescriptionLength)) + if re.Response != nil { + attrs = append(attrs, "status", re.Response.StatusCode) + } + case errors.Is(err, sso.ErrUnavailable): + attrs = append(attrs, "reason", "provider_unavailable") + case errors.Is(err, sso.ErrIDToken): + attrs = append(attrs, "reason", "id_token_verification_failed") + case errors.Is(err, sso.ErrClaims): + attrs = append(attrs, "reason", "claims_unusable") + default: + attrs = append(attrs, "reason", "exchange_failed") + } + h.logger.Warn("auth.login", attrs...) +} + +func (h *OIDCHTTP) redirectError(w http.ResponseWriter, r *http.Request, code string) { + http.Redirect(w, r, "/login?error="+code, http.StatusSeeOther) +} + +func writeOIDCRefusal(w http.ResponseWriter, message string) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.Header().Set("Content-Security-Policy", "default-src 'none'") + w.Header().Set("X-Content-Type-Options", "nosniff") + w.WriteHeader(http.StatusForbidden) + _, _ = fmt.Fprintf(w, oidcRefusalPage, html.EscapeString(message)) +} + +// truncate cuts s to at most n bytes on a rune boundary. +func truncate(s string, n int) string { + if len(s) <= n { + return s + } + for n > 0 && !utf8.RuneStart(s[n]) { + n-- + } + return s[:n] +} diff --git a/server/auth_oidc_security_test.go b/server/auth_oidc_security_test.go new file mode 100644 index 0000000..21680f6 --- /dev/null +++ b/server/auth_oidc_security_test.go @@ -0,0 +1,608 @@ +package server + +import ( + "bytes" + "context" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/bananaops/tracker/internal/auth" + "github.com/bananaops/tracker/internal/auth/authz" + "github.com/bananaops/tracker/internal/auth/sso" + "github.com/bananaops/tracker/internal/auth/sso/ssotest" + store "github.com/bananaops/tracker/internal/stores" + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// requireTransactionCleared checks the callback expired the transaction cookie. +func requireTransactionCleared(t *testing.T, rec *httptest.ResponseRecorder) { + t.Helper() + c := cookieNamed(rec, sso.TransactionCookieName) + require.NotNil(t, c, "transaction cookie must be cleared") + assert.Equal(t, -1, c.MaxAge) + assert.Empty(t, c.Value) +} + +// requireRefusedRedirect checks a failed callback: 303 to loc, no session, transaction cookie cleared. +func requireRefusedRedirect(t *testing.T, rec *httptest.ResponseRecorder, loc string) { + t.Helper() + assert.Equal(t, http.StatusSeeOther, rec.Code) + assert.Equal(t, loc, rec.Header().Get("Location")) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + requireTransactionCleared(t, rec) +} + +// requireRefusedPage checks a 403 HTML refusal: no session, CSP, transaction cookie cleared. +func requireRefusedPage(t *testing.T, rec *httptest.ResponseRecorder, contains string) { + t.Helper() + assert.Equal(t, http.StatusForbidden, rec.Code) + assert.Contains(t, rec.Header().Get("Content-Type"), "text/html") + assert.NotEmpty(t, rec.Header().Get("Content-Security-Policy")) + assert.Contains(t, rec.Body.String(), contains) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + requireTransactionCleared(t, rec) +} + +func (h *oidcHarness) userNamed(t *testing.T, name string) *store.User { + t.Helper() + u, err := h.f.users.GetByUsername(context.Background(), name) + require.NoError(t, err) + return u +} + +func (h *oidcHarness) teamNamed(t *testing.T, name string) *store.Team { + t.Helper() + tm, err := h.f.teams.GetByName(context.Background(), name) + require.NoError(t, err) + return tm +} + +func (h *oidcHarness) createTeam(t *testing.T, name string, groups ...string) *store.Team { + t.Helper() + tm := &store.Team{Name: name, Permissions: []string{"event:read"}, OIDCGroups: groups} + require.NoError(t, h.f.teams.Create(context.Background(), tm)) + return tm +} + +func (h *oidcHarness) setGroups(groups any) { + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "alice", "groups": groups, + }}) +} + +func (h *oidcHarness) inTeam(t *testing.T, username, team string) bool { + t.Helper() + tm := h.teamNamed(t, team) + for _, id := range h.userNamed(t, username).Teams { + if id == tm.ID { + return true + } + } + return false +} + +func teamNames(m meBody) []string { + names := make([]string, 0, len(m.Teams)) + for _, tm := range m.Teams { + names = append(names, tm.Name) + } + return names +} + +func TestOIDCCallbackCSRFStateMismatch(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + q := cb.Query() + q.Set("state", "another-state-value") + cb.RawQuery = q.Encode() + failures := oidcLoginCount("failure") + + rec := h.callback(t, cb, tx) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) + assert.Equal(t, failures+1, oidcLoginCount("failure")) +} + +// A login CSRF: the attacker's code and state are delivered to the victim's +// browser, which holds its own transaction cookie. +func TestOIDCCallbackForeignLogin(t *testing.T) { + h := newOIDCHarness(t, nil) + _, victim := h.start(t, "") + attackerAuthURL, _ := h.start(t, "") + attackerCB := h.idp.Authorize(t, attackerAuthURL.String()) + + rec := h.callback(t, attackerCB, victim) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) + _, err := h.f.users.GetByUsername(context.Background(), "alice") + assert.ErrorIs(t, err, store.ErrNotFound) +} + +func TestOIDCCallbackWithoutTransaction(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, _ := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + failures, successes := oidcLoginCount("failure"), oidcLoginCount("success") + + rec := h.callback(t, cb, nil) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) + assert.Equal(t, failures, oidcLoginCount("failure")) + assert.Equal(t, successes, oidcLoginCount("success")) +} + +func TestOIDCCallbackReplayWithoutCookie(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + require.Equal(t, http.StatusSeeOther, h.callback(t, cb, tx).Code) + + rec := h.callback(t, cb, nil) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 1, h.idp.TokenRequests()) +} + +func TestOIDCCallbackReplayWithCapturedCookie(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + first := h.callback(t, cb, tx) + require.Equal(t, http.StatusSeeOther, first.Code) + require.NotNil(t, cookieNamed(first, auth.SessionCookieName)) + + // The browser dropped the cleared cookie, an attacker kept a copy. + rec := h.callback(t, cb, tx) + requireRefusedRedirect(t, rec, "/login?error=oidc_failed") + assert.Equal(t, 2, h.idp.TokenRequests()) +} + +func TestOIDCCallbackIDTokenRejections(t *testing.T) { + now := time.Now() + cases := []struct { + name string + setup func(*ssotest.IdP) + }{ + {"nonce mismatch", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { c["nonce"] = "forged" }) + }}, + {"wrong audience", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { c["aud"] = "another-client" }) + }}, + {"expired id_token", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { + c["exp"] = now.Add(-time.Hour).Unix() + c["iat"] = now.Add(-2 * time.Hour).Unix() + }) + }}, + {"wrong issuer", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { c["iss"] = "https://evil.example" }) + }}, + {"alg none", func(i *ssotest.IdP) { i.SetSigning(ssotest.SignNone) }}, + {"unknown key", func(i *ssotest.IdP) { i.SetSigning(ssotest.SignForeignKey) }}, + {"algorithm confusion HS256 with the public key", func(i *ssotest.IdP) { i.SetSigning(ssotest.SignHS256PublicKey) }}, + {"missing sub", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { delete(c, "sub") }) + }}, + {"empty sub", func(i *ssotest.IdP) { + i.SetTokenMutator(func(c map[string]any) { c["sub"] = "" }) + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + h := newOIDCHarness(t, nil) + tc.setup(h.idp) + failures := oidcLoginCount("failure") + + rec := h.login(t, "") + requireRefusedRedirect(t, rec, "/login?error=oidc_failed") + assert.Equal(t, failures+1, oidcLoginCount("failure")) + assert.Equal(t, 1, h.idp.TokenRequests()) + users, err := h.f.users.List(context.Background()) + require.NoError(t, err) + assert.Len(t, users, 1, "only the bootstrap admin exists, no user is created") + }) + } +} + +func TestOIDCCallbackTamperedCookie(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + v := []byte(tx.Value) + mid := len(v) / 2 + if v[mid] == 'A' { + v[mid] = 'B' + } else { + v[mid] = 'A' + } + + rec := h.callback(t, cb, &http.Cookie{Name: tx.Name, Value: string(v)}) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) +} + +func TestOIDCCallbackCookieFromOtherSecret(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + plain, err := h.codec.Decode(tx.Value) + require.NoError(t, err) + other, err := sso.NewTransactionCodec(bytes.Repeat([]byte{8}, 32)) + require.NoError(t, err) + forged, err := other.Encode(plain) + require.NoError(t, err) + + rec := h.callback(t, cb, sso.TransactionCookie(forged, false)) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) +} + +func TestOIDCCallbackExpiredTransaction(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.codec.Now = func() time.Time { return time.Now().Add(11 * time.Minute) } + + rec := h.callback(t, cb, tx) + requireRefusedRedirect(t, rec, "/login?error=oidc_state") + assert.Equal(t, 0, h.idp.TokenRequests()) + assert.Contains(t, h.logs.String(), "transaction_expired") +} + +func TestOIDCCallbackIdPError(t *testing.T) { + h := newOIDCHarness(t, nil) + h.idp.SetAuthorizeError("access_denied", "") + authURL, tx := h.start(t, "") + failures := oidcLoginCount("failure") + + rec := h.callback(t, h.idp.Authorize(t, authURL.String()), tx) + requireRefusedRedirect(t, rec, "/login?error=oidc_denied") + assert.NotContains(t, rec.Header().Get("Location"), "script") + assert.NotContains(t, rec.Body.String(), "script") + assert.Equal(t, failures+1, oidcLoginCount("failure")) + assert.Equal(t, 0, h.idp.TokenRequests()) +} + +// Every failure path of the callback expires the transaction cookie and issues no session. +func TestOIDCTransactionCookieClearedOnEveryFailure(t *testing.T) { + type outcome struct { + rec *httptest.ResponseRecorder + } + cases := []struct { + name string + run func(t *testing.T) outcome + }{ + {"idp error", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + h.idp.SetAuthorizeError("access_denied", "no") + authURL, tx := h.start(t, "") + return outcome{h.callback(t, h.idp.Authorize(t, authURL.String()), tx)} + }}, + {"bad state", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + q := cb.Query() + q.Set("state", "x") + cb.RawQuery = q.Encode() + return outcome{h.callback(t, cb, tx)} + }}, + {"expired transaction", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.codec.Now = func() time.Time { return time.Now().Add(time.Hour) } + return outcome{h.callback(t, cb, tx)} + }}, + {"exchange failure", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.oidc.provider = stubProvider{err: sso.ErrExchange} + return outcome{h.callback(t, cb, tx)} + }}, + {"id_token failure", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + h.idp.SetSigning(ssotest.SignNone) + return outcome{h.login(t, "")} + }}, + {"provisioning disabled", func(t *testing.T) outcome { + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC.UserProvisioning = false }) + return outcome{h.login(t, "")} + }}, + {"disabled user", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + u := h.userNamed(t, "alice") + u.Disabled = true + require.NoError(t, h.f.users.Update(context.Background(), u)) + return outcome{h.login(t, "")} + }}, + {"missing groups claim", func(t *testing.T) outcome { + h := newOIDCHarness(t, nil) + h.createTeam(t, "Platform", "platform-eng") + h.setGroups([]string{"platform-eng"}) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"preferred_username": "alice"}}) + return outcome{h.login(t, "")} + }}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + rec := tc.run(t).rec + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + requireTransactionCleared(t, rec) + }) + } +} + +func TestOIDCOpenRedirect(t *testing.T) { + cases := []struct{ redirect, want string }{ + {"//evil.example", "/"}, + {"/\\evil.example", "/"}, + {"https://evil.example", "/"}, + {"/\r\nX:1", "/"}, + {"/locks?tab=1", "/locks?tab=1"}, + } + for _, tc := range cases { + t.Run(tc.redirect, func(t *testing.T) { + h := newOIDCHarness(t, nil) + path := auth.OIDCLoginPath + "?" + url.Values{"redirect": {tc.redirect}}.Encode() + rec := h.get(path, nil) + require.Equal(t, http.StatusFound, rec.Code) + authURL, err := url.Parse(rec.Header().Get("Location")) + require.NoError(t, err) + tx := cookieNamed(rec, sso.TransactionCookieName) + require.NotNil(t, tx) + + cbRec := h.callback(t, h.idp.Authorize(t, authURL.String()), tx) + require.Equal(t, http.StatusSeeOther, cbRec.Code) + assert.Equal(t, tc.want, cbRec.Header().Get("Location")) + assert.Empty(t, cbRec.Header().Get("X")) + }) + } +} + +func TestOIDCUsernameCollision(t *testing.T) { + h := newOIDCHarness(t, nil) + local := &store.User{Username: "alice", Source: store.UserSourceLocal, PasswordHash: "local-hash"} + require.NoError(t, h.f.users.Create(context.Background(), local)) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + me := h.me(t, cookieNamed(rec, auth.SessionCookieName)) + assert.Equal(t, "alice-2", me.Username) + assert.Equal(t, "oidc", me.Source) + + after := h.userNamed(t, "alice") + assert.Equal(t, store.UserSourceLocal, after.Source) + assert.Equal(t, "local-hash", after.PasswordHash) + assert.Empty(t, after.OIDCSubject) + assert.Empty(t, after.OIDCIssuer) +} + +func TestOIDCCannotTakeOverAdmin(t *testing.T) { + h := newOIDCHarness(t, nil) + before := h.userNamed(t, "admin") + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "admin", "groups": []string{}, + }}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + me := h.me(t, cookieNamed(rec, auth.SessionCookieName)) + assert.Equal(t, "admin-2", me.Username) + assert.False(t, me.IsAdmin) + + after := h.userNamed(t, "admin") + assert.Equal(t, before.ID, after.ID) + assert.Equal(t, before.PasswordHash, after.PasswordHash) + assert.Equal(t, store.UserSourceLocal, after.Source) + assert.Empty(t, after.OIDCSubject) + assert.Equal(t, before.Teams, after.Teams) +} + +func TestOIDCDisabledUser(t *testing.T) { + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + u := h.userNamed(t, "alice") + u.Disabled = true + require.NoError(t, h.f.users.Update(context.Background(), u)) + failures := oidcLoginCount("failure") + + rec := h.login(t, "") + requireRefusedPage(t, rec, "disabled") + assert.Equal(t, failures+1, oidcLoginCount("failure")) +} + +func TestOIDCProvisioningDisabled(t *testing.T) { + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC.UserProvisioning = false }) + + rec := h.login(t, "") + requireRefusedPage(t, rec, "not registered") + _, err := h.f.users.GetByUsername(context.Background(), "alice") + assert.ErrorIs(t, err, store.ErrNotFound) + + require.NoError(t, h.f.users.Create(context.Background(), &store.User{ + Username: "alice", Source: store.UserSourceOIDC, OIDCIssuer: h.idp.URL, OIDCSubject: "user-1", + })) + rec = h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.NotNil(t, cookieNamed(rec, auth.SessionCookieName)) +} + +func TestOIDCNoUsableUsername(t *testing.T) { + h := newOIDCHarness(t, nil) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"groups": []string{}}}) + + rec := h.login(t, "") + requireRefusedRedirect(t, rec, "/login?error=oidc_failed") +} + +func TestOIDCGroupsSingleString(t *testing.T) { + h := newOIDCHarness(t, nil) + h.createTeam(t, "Platform", "platform-eng") + h.setGroups("platform-eng") + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.Equal(t, []string{"Platform"}, teamNames(h.me(t, cookieNamed(rec, auth.SessionCookieName)))) +} + +func TestOIDCGroupsCaseSensitive(t *testing.T) { + h := newOIDCHarness(t, nil) + h.createTeam(t, "Platform", "platform-eng") + h.setGroups([]string{"Platform-Eng"}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.Empty(t, teamNames(h.me(t, cookieNamed(rec, auth.SessionCookieName)))) + assert.False(t, h.inTeam(t, "alice", "Platform")) +} + +// A groups claim that disappears never strips the mapped teams of a user who +// holds one: the login is refused before any write (fail closed). A user +// without mapped membership is not affected, see TestOIDCGroupsClaimAbsent*. +func TestOIDCGroupsAbsentKeepsMappedTeams(t *testing.T) { + h := newOIDCHarness(t, nil) + h.createTeam(t, "Platform", "platform-eng") + h.setGroups([]string{"platform-eng"}) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + require.True(t, h.inTeam(t, "alice", "Platform")) + + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"preferred_username": "alice"}}) + rec := h.login(t, "") + requireRefusedRedirect(t, rec, "/login?error=oidc_failed") + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + assert.True(t, h.inTeam(t, "alice", "Platform"), "membership untouched") + assert.Contains(t, h.logs.String(), "groups_claim_missing") +} + +func TestOIDCTeamSyncDisabled(t *testing.T) { + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC.TeamSync = false }) + h.createTeam(t, "Platform", "platform-eng") + h.setGroups([]string{"platform-eng"}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.Empty(t, teamNames(h.me(t, cookieNamed(rec, auth.SessionCookieName)))) +} + +func TestOIDCManualTeamKept(t *testing.T) { + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + manual := h.createTeam(t, "Manual") + h.createTeam(t, "Ops", "ops") + u := h.userNamed(t, "alice") + u.Teams = append(u.Teams, manual.ID) + require.NoError(t, h.f.users.Update(context.Background(), u)) + h.setGroups([]string{"ops"}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.ElementsMatch(t, []string{"Manual", "Ops"}, teamNames(h.me(t, cookieNamed(rec, auth.SessionCookieName)))) +} + +func (h *oidcHarness) mapAdministrators(t *testing.T, group string) { + t.Helper() + admins := h.teamNamed(t, store.AdministratorsTeamName) + admins.OIDCGroups = []string{group} + require.NoError(t, h.f.teams.Update(context.Background(), admins)) +} + +func TestOIDCAdministratorsMapping(t *testing.T) { + h := newOIDCHarness(t, nil) + h.mapAdministrators(t, "tracker-admins") + + h.setGroups([]string{"tracker-admins"}) + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.True(t, h.me(t, cookieNamed(rec, auth.SessionCookieName)).IsAdmin) + + h.setGroups([]string{}) + rec = h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.False(t, h.me(t, cookieNamed(rec, auth.SessionCookieName)).IsAdmin) + + assert.True(t, h.inTeam(t, "admin", store.AdministratorsTeamName), "the local admin keeps Administrators") +} + +func TestOIDCLastAdminKept(t *testing.T) { + h := newOIDCHarness(t, nil) + h.mapAdministrators(t, "tracker-admins") + h.setGroups([]string{"tracker-admins"}) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + + admin := h.userNamed(t, "admin") + admin.Disabled = true + require.NoError(t, h.f.users.Update(context.Background(), admin)) + + h.setGroups([]string{}) + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.True(t, h.me(t, cookieNamed(rec, auth.SessionCookieName)).IsAdmin) + assert.Contains(t, h.logs.String(), "last enabled administrator") +} + +func TestOIDCLogsCarryNoSecrets(t *testing.T) { + h := newOIDCHarness(t, nil) + secrets := []string{h.idp.ClientSecret} + collect := func(cb *url.URL, tx *http.Cookie) { + plain, err := h.codec.Decode(tx.Value) + require.NoError(t, err) + secrets = append(secrets, cb.Query().Get("code"), cb.Query().Get("state"), tx.Value, + plain.Verifier, plain.State, plain.Nonce, h.idp.LastIDToken()) + } + + // Successful login. + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + require.Equal(t, http.StatusSeeOther, h.callback(t, cb, tx).Code) + collect(cb, tx) + + // Forged nonce. + h.idp.SetTokenMutator(func(c map[string]any) { c["nonce"] = "forged-nonce-value" }) + authURL, tx = h.start(t, "") + cb = h.idp.Authorize(t, authURL.String()) + require.Equal(t, "/login?error=oidc_failed", h.callback(t, cb, tx).Header().Get("Location")) + collect(cb, tx) + h.idp.SetTokenMutator(nil) + + // Replay of the forged flow: the code is already consumed. + require.Equal(t, "/login?error=oidc_failed", h.callback(t, cb, tx).Header().Get("Location")) + + logs := h.logs.String() + require.NotEmpty(t, logs) + for _, s := range secrets { + require.NotEmpty(t, s) + assert.False(t, strings.Contains(logs, s), "log leaks a secret") + } +} + +func TestOIDCMetrics(t *testing.T) { + h := newOIDCHarness(t, nil) + success, failure := oidcLoginCount("success"), oidcLoginCount("failure") + localSuccess := oidcLocalCount(authz.LoginSuccess) + localFailure := oidcLocalCount(authz.LoginFailure) + + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + h.idp.SetTokenMutator(func(c map[string]any) { c["nonce"] = "forged" }) + assert.Equal(t, "/login?error=oidc_failed", h.login(t, "").Header().Get("Location")) + + assert.Equal(t, success+1, oidcLoginCount("success")) + assert.Equal(t, failure+1, oidcLoginCount("failure")) + assert.Equal(t, localSuccess, oidcLocalCount(authz.LoginSuccess)) + assert.Equal(t, localFailure, oidcLocalCount(authz.LoginFailure)) +} + +// oidcLocalCount reads tracker_auth_logins_total{method="local"} for one result. +func oidcLocalCount(result string) float64 { + return testutil.ToFloat64(authz.AuthLogins.WithLabelValues(authz.LoginMethodLocal, result)) +} diff --git a/server/auth_oidc_test.go b/server/auth_oidc_test.go new file mode 100644 index 0000000..f50ad09 --- /dev/null +++ b/server/auth_oidc_test.go @@ -0,0 +1,410 @@ +package server + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/bananaops/tracker/internal/auth" + "github.com/bananaops/tracker/internal/auth/sso" + "github.com/bananaops/tracker/internal/auth/sso/ssotest" + store "github.com/bananaops/tracker/internal/stores" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.mongodb.org/mongo-driver/bson/primitive" + "golang.org/x/oauth2" +) + +func TestOIDCRoutesAbsentWhenDisabled(t *testing.T) { + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC = auth.OIDCConfig{} }) + + assert.Equal(t, http.StatusNotFound, h.get(auth.OIDCLoginPath, nil).Code) + assert.Equal(t, http.StatusNotFound, h.get(auth.OIDCCallbackPath, nil).Code) + rec := h.get("/api/v1alpha1/auth/config", nil) + require.Equal(t, http.StatusOK, rec.Code) + assert.NotContains(t, rec.Body.String(), `"oidcEnabled":true`) +} + +func TestOIDCLoginRedirectsToProvider(t *testing.T) { + h := newOIDCHarness(t, nil) + rec := h.startRaw("/locks") + + require.Equal(t, http.StatusFound, rec.Code) + assert.Contains(t, rec.Header().Get("Cache-Control"), "no-store") + loc, err := url.Parse(rec.Header().Get("Location")) + require.NoError(t, err) + assert.Equal(t, h.idp.URL+"/authorize", loc.Scheme+"://"+loc.Host+loc.Path) + assert.Equal(t, "S256", loc.Query().Get("code_challenge_method")) + + c := cookieNamed(rec, sso.TransactionCookieName) + require.NotNil(t, c) + assert.True(t, c.HttpOnly) + assert.Equal(t, http.SameSiteLaxMode, c.SameSite) + assert.Equal(t, auth.OIDCCallbackPath, c.Path) + assert.Equal(t, 600, c.MaxAge) + assert.False(t, c.Secure) + tx, err := h.codec.Decode(c.Value) + require.NoError(t, err) + assert.Equal(t, "/locks", tx.Redirect) + assert.Equal(t, tx.State, loc.Query().Get("state")) + assert.Equal(t, tx.Nonce, loc.Query().Get("nonce")) +} + +func TestOIDCLoginSecureCookie(t *testing.T) { + h := newOIDCHarness(t, func(c *auth.Config) { c.CookieSecure = true }) + _, tx := h.start(t, "") + assert.True(t, tx.Secure) +} + +func TestOIDCLoginEndToEnd(t *testing.T) { + h := newOIDCHarness(t, nil) + before := oidcLoginCount("success") + + rec := h.login(t, "/locks") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.Equal(t, "/locks", rec.Header().Get("Location")) + assert.Equal(t, "no-referrer", rec.Header().Get("Referrer-Policy")) + cleared := cookieNamed(rec, sso.TransactionCookieName) + require.NotNil(t, cleared) + assert.Equal(t, -1, cleared.MaxAge) + session := cookieNamed(rec, auth.SessionCookieName) + require.NotNil(t, session) + _, err := h.f.sessions.Verify(session.Value) + require.NoError(t, err) + + me := h.me(t, session) + assert.True(t, me.Authenticated) + assert.Equal(t, "alice", me.Username) + assert.Equal(t, "oidc", me.Source) + assert.False(t, me.MustChangePassword) + + u, err := h.f.users.GetByUsername(context.Background(), "alice") + require.NoError(t, err) + assert.Equal(t, store.UserSourceOIDC, u.Source) + assert.Equal(t, h.idp.URL, u.OIDCIssuer) + assert.Equal(t, "user-1", u.OIDCSubject) + assert.NotNil(t, u.LastLoginAt) + assert.Equal(t, before+1, oidcLoginCount("success")) +} + +func TestOIDCSecondLoginSameUser(t *testing.T) { + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + + users, err := h.f.users.List(context.Background()) + require.NoError(t, err) + n := 0 + for _, u := range users { + if u.Username == "alice" { + n++ + } + } + assert.Equal(t, 1, n) +} + +func TestOIDCTeamMappingEndToEnd(t *testing.T) { + h := newOIDCHarness(t, nil) + team := &store.Team{Name: "Platform", Permissions: []string{"event:read"}, OIDCGroups: []string{"platform-eng"}} + require.NoError(t, h.f.teams.Create(context.Background(), team)) + + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "alice", "groups": []string{"platform-eng"}, + }}) + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + me := h.me(t, cookieNamed(rec, auth.SessionCookieName)) + require.Len(t, me.Teams, 1) + assert.Equal(t, "Platform", me.Teams[0].Name) + + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "alice", "groups": []string{}, + }}) + rec = h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.Empty(t, h.me(t, cookieNamed(rec, auth.SessionCookieName)).Teams) +} + +func TestOIDCGroupsClaimMissingFailsClosed(t *testing.T) { + ctx := context.Background() + h := newOIDCHarness(t, nil) + + // Existing OIDC user, provisioned while no team was mapped. + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + before, err := h.f.users.GetByUsername(ctx, "alice") + require.NoError(t, err) + + platform := h.createTeam(t, "Platform", "platform-eng") + require.NoError(t, h.f.users.SyncTeams(ctx, before.ID, []primitive.ObjectID{platform.ID}, nil)) + before, err = h.f.users.GetByUsername(ctx, "alice") + require.NoError(t, err) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "alice", "email": "changed@example.com", "name": "Changed Name", + }}) + failures := oidcLoginCount("failure") + + rec := h.login(t, "") + assert.Equal(t, http.StatusSeeOther, rec.Code) + assert.Equal(t, "/login?error=oidc_failed", rec.Header().Get("Location")) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + assert.Equal(t, failures+1, oidcLoginCount("failure")) + assert.Contains(t, h.logs.String(), `"level":"ERROR"`) + assert.Contains(t, h.logs.String(), `"claim":"groups"`) + assert.Contains(t, h.logs.String(), `"username":"alice"`) + + after, err := h.f.users.GetByUsername(ctx, "alice") + require.NoError(t, err) + assert.Equal(t, before.Email, after.Email) + assert.Equal(t, before.DisplayName, after.DisplayName) + assert.Equal(t, before.LastLoginAt.UTC(), after.LastLoginAt.UTC()) + assert.Equal(t, before.Teams, after.Teams) +} + +func TestOIDCDiscoveryErrorBodyNeverLogged(t *testing.T) { + broken := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte("LEAKED-DISCOVERY-BODY")) + })) + t.Cleanup(broken.Close) + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC.Issuer = broken.URL }) + + rec := h.startRaw("") + assert.Equal(t, "/login?error=oidc_unavailable", rec.Header().Get("Location")) + assert.Contains(t, h.logs.String(), "provider_unavailable") + assert.NotContains(t, h.logs.String(), "LEAKED-DISCOVERY-BODY") +} + +func TestOIDCVerificationErrorNeverLogsDetail(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.oidc.provider = stubProvider{err: fmt.Errorf("%w: failed to fetch keys: LEAKED-JWKS-BODY", sso.ErrIDToken)} + + rec := h.callback(t, cb, tx) + assert.Equal(t, "/login?error=oidc_failed", rec.Header().Get("Location")) + assert.Contains(t, h.logs.String(), "id_token_verification_failed") + assert.NotContains(t, h.logs.String(), "LEAKED-JWKS-BODY") +} + +func TestOIDCUserCannotChangePassword(t *testing.T) { + h := newOIDCHarness(t, nil) + session := cookieNamed(h.login(t, ""), auth.SessionCookieName) + require.NotNil(t, session) + + rec := post(h.handler, "/api/v1alpha1/auth/password", `{"currentPassword":"x","newPassword":"another-long-password-1"}`, session) + assert.Equal(t, http.StatusBadRequest, rec.Code) + assert.Contains(t, rec.Body.String(), "managed by the identity provider") +} + +func TestLocalLoginStillWorksWhenProviderDown(t *testing.T) { + h := newOIDCHarness(t, nil) + h.idp.Close() + before := oidcLoginCount("failure") + + rec := h.startRaw("") + assert.Equal(t, http.StatusSeeOther, rec.Code) + assert.Equal(t, "/login?error=oidc_unavailable", rec.Header().Get("Location")) + assert.Nil(t, cookieNamed(rec, sso.TransactionCookieName)) + + rec = post(h.handler, "/api/v1alpha1/auth/login", `{"username":"admin","password":"admin-password-123"}`, nil) + assert.Equal(t, http.StatusNoContent, rec.Code) + assert.Contains(t, h.get("/api/v1alpha1/auth/config", nil).Body.String(), `"oidcEnabled":true`) + assert.Equal(t, before+1, oidcLoginCount("failure")) +} + +func TestLoginPageRedirectIsSanitized(t *testing.T) { + h := newOIDCHarness(t, nil) + _, c := h.start(t, "//evil.example") + tx, err := h.codec.Decode(c.Value) + require.NoError(t, err) + assert.Equal(t, "/", tx.Redirect) +} + +func TestOIDCCallbackRedirectReSanitized(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, orig := h.start(t, "") + tx, err := h.codec.Decode(orig.Value) + require.NoError(t, err) + // A hostile redirect that would somehow be sealed in the cookie is + // neutralised again on the way out. + tx.Redirect = "//evil.example" + v, err := h.codec.Encode(tx) + require.NoError(t, err) + + rec := h.callback(t, h.idp.Authorize(t, authURL.String()), sso.TransactionCookie(v, false)) + require.Equal(t, http.StatusSeeOther, rec.Code) + assert.Equal(t, "/", rec.Header().Get("Location")) +} + +func TestOIDCCallbackRefusals(t *testing.T) { + t.Run("missing transaction cookie", func(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, _ := h.start(t, "") + rec := h.callback(t, h.idp.Authorize(t, authURL.String()), nil) + assert.Equal(t, "/login?error=oidc_state", rec.Header().Get("Location")) + assert.Equal(t, 0, h.idp.TokenRequests()) + }) + t.Run("state mismatch", func(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + q := cb.Query() + q.Set("state", "forged") + cb.RawQuery = q.Encode() + before := oidcLoginCount("failure") + rec := h.callback(t, cb, tx) + assert.Equal(t, "/login?error=oidc_state", rec.Header().Get("Location")) + assert.Equal(t, 0, h.idp.TokenRequests()) + assert.Equal(t, before+1, oidcLoginCount("failure")) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + }) + t.Run("idp error", func(t *testing.T) { + h := newOIDCHarness(t, nil) + h.idp.SetAuthorizeError("access_denied", "nope") + authURL, tx := h.start(t, "") + rec := h.callback(t, h.idp.Authorize(t, authURL.String()), tx) + assert.Equal(t, "/login?error=oidc_denied", rec.Header().Get("Location")) + assert.Contains(t, h.logs.String(), "access_denied") + }) + t.Run("provisioning disabled", func(t *testing.T) { + h := newOIDCHarness(t, func(c *auth.Config) { c.OIDC.UserProvisioning = false }) + rec := h.login(t, "") + assert.Equal(t, http.StatusForbidden, rec.Code) + assert.Contains(t, rec.Body.String(), "not registered in Tracker") + assert.Equal(t, "default-src 'none'", rec.Header().Get("Content-Security-Policy")) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + }) + t.Run("disabled user", func(t *testing.T) { + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + u, err := h.f.users.GetByUsername(context.Background(), "alice") + require.NoError(t, err) + u.Disabled = true + require.NoError(t, h.f.users.Update(context.Background(), u)) + rec := h.login(t, "") + assert.Equal(t, http.StatusForbidden, rec.Code) + assert.Contains(t, rec.Body.String(), "account is disabled") + }) +} + +// stubProvider fails the exchange with a canned error. +type stubProvider struct{ err error } + +func (stubProvider) AuthCodeURL(context.Context, string, string, string) (string, error) { + return "http://idp.test/authorize", nil +} + +func (s stubProvider) Exchange(context.Context, string, string, string) (sso.Claims, error) { + return sso.Claims{}, s.err +} + +func TestOIDCExchangeErrorNeverLogsRawBody(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.oidc.provider = stubProvider{err: fmt.Errorf("%w: %w", sso.ErrExchange, &oauth2.RetrieveError{ + Response: &http.Response{StatusCode: http.StatusBadGateway}, + Body: []byte("LEAKED-RAW-BODY access_token=abc"), + })} + + rec := h.callback(t, cb, tx) + assert.Equal(t, "/login?error=oidc_failed", rec.Header().Get("Location")) + assert.NotContains(t, h.logs.String(), "LEAKED-RAW-BODY") + assert.NotContains(t, h.logs.String(), "access_token") + assert.Contains(t, h.logs.String(), `"status":502`) + assert.Contains(t, h.logs.String(), "exchange_failed") +} + +func TestOIDCExchangeUnavailableRedirects(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + h.oidc.provider = stubProvider{err: fmt.Errorf("%w: boom", sso.ErrUnavailable)} + rec := h.callback(t, cb, tx) + assert.Equal(t, "/login?error=oidc_unavailable", rec.Header().Get("Location")) +} + +func TestOIDCNoSecretsInLogs(t *testing.T) { + h := newOIDCHarness(t, nil) + authURL, tx := h.start(t, "") + cb := h.idp.Authorize(t, authURL.String()) + rec := h.callback(t, cb, tx) + require.Equal(t, http.StatusSeeOther, rec.Code) + logs := h.logs.String() + for _, secret := range []string{h.idp.ClientSecret, h.idp.LastIDToken(), cb.Query().Get("code"), cb.Query().Get("state"), tx.Value} { + require.NotEmpty(t, secret) + assert.False(t, strings.Contains(logs, secret), "log leaks a secret") + } +} + +// A provider that omits the groups claim for a user without groups must not +// lock that user out: only users holding a mapped membership are refused. +func TestOIDCGroupsClaimAbsentNewSubject(t *testing.T) { + ctx := context.Background() + h := newOIDCHarness(t, nil) + h.createTeam(t, "Platform", "platform-eng") + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"preferred_username": "alice"}}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.NotNil(t, cookieNamed(rec, auth.SessionCookieName)) + u, err := h.f.users.GetByUsername(ctx, "alice") + require.NoError(t, err) + assert.Empty(t, u.Teams) + assert.Contains(t, h.logs.String(), `"level":"WARN"`) + assert.Contains(t, h.logs.String(), `"reason":"groups_claim_missing_accepted"`) + assert.Contains(t, h.logs.String(), `"claim":"groups"`) + assert.Contains(t, h.logs.String(), `"username":"alice"`) +} + +func TestOIDCGroupsClaimAbsentUserWithoutMappedMembership(t *testing.T) { + ctx := context.Background() + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + manual := h.createTeam(t, "Manual") + h.createTeam(t, "Platform", "platform-eng") + u := h.userNamed(t, "alice") + require.NoError(t, h.f.users.SyncTeams(ctx, u.ID, []primitive.ObjectID{manual.ID}, nil)) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"preferred_username": "alice"}}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.NotNil(t, cookieNamed(rec, auth.SessionCookieName)) + assert.Equal(t, []string{"Manual"}, teamNames(h.me(t, cookieNamed(rec, auth.SessionCookieName)))) + assert.Contains(t, h.logs.String(), `"reason":"groups_claim_missing_accepted"`) +} + +// Without any mapped team an absent claim is normal: no warning. +func TestOIDCGroupsClaimAbsentNoMappedTeamNoWarning(t *testing.T) { + h := newOIDCHarness(t, nil) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{"preferred_username": "alice"}}) + + rec := h.login(t, "") + require.Equal(t, http.StatusSeeOther, rec.Code, rec.Body.String()) + assert.NotContains(t, h.logs.String(), "groups_claim_missing_accepted") +} + +// A groups claim of an unexpected type is treated as absent: it never removes +// mapped teams and it is reported. +func TestOIDCGroupsClaimUnexpectedTypeTreatedAsAbsent(t *testing.T) { + ctx := context.Background() + h := newOIDCHarness(t, nil) + require.Equal(t, http.StatusSeeOther, h.login(t, "").Code) + platform := h.createTeam(t, "Platform", "platform-eng") + u := h.userNamed(t, "alice") + require.NoError(t, h.f.users.SyncTeams(ctx, u.ID, []primitive.ObjectID{platform.ID}, nil)) + h.idp.SetUser(ssotest.User{Subject: "user-1", Claims: map[string]any{ + "preferred_username": "alice", "groups": map[string]any{"platform-eng": true}, + }}) + + rec := h.login(t, "") + assert.Equal(t, "/login?error=oidc_failed", rec.Header().Get("Location")) + assert.Nil(t, cookieNamed(rec, auth.SessionCookieName)) + assert.Contains(t, h.logs.String(), `"reason":"groups_claim_unexpected_type"`) + assert.Contains(t, h.logs.String(), `"reason":"groups_claim_missing"`) + assert.Equal(t, []primitive.ObjectID{platform.ID}, h.userNamed(t, "alice").Teams) +} diff --git a/server/auth_oidc_testing_test.go b/server/auth_oidc_testing_test.go new file mode 100644 index 0000000..069cec4 --- /dev/null +++ b/server/auth_oidc_testing_test.go @@ -0,0 +1,161 @@ +package server + +import ( + "bytes" + "context" + "encoding/json" + "log/slog" + "net/http" + "net/http/httptest" + "net/url" + "testing" + + authv1 "github.com/bananaops/tracker/generated/proto/auth/v1alpha1" + "github.com/bananaops/tracker/internal/auth" + "github.com/bananaops/tracker/internal/auth/authz" + "github.com/bananaops/tracker/internal/auth/sso" + "github.com/bananaops/tracker/internal/auth/sso/ssotest" + "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/require" +) + +const oidcTestPublicURL = "http://tracker.test" + +type oidcHarness struct { + f *authFixture + idp *ssotest.IdP + cfg auth.Config + codec *sso.TransactionCodec + oidc *OIDCHTTP + handler http.Handler + logs *bytes.Buffer +} + +// newOIDCHarness wires the gateway (AuthService), the cookie endpoints and the +// OIDC routes on a real mux, behind the real auth middleware. +func newOIDCHarness(t *testing.T, mutate func(*auth.Config)) *oidcHarness { + t.Helper() + f := newAuthFixture(t) + idp := ssotest.New(t) + + cfg := f.cfg + cfg.PublicURL = oidcTestPublicURL + cfg.OIDC = auth.OIDCConfig{ + Issuer: idp.URL, + ClientID: idp.ClientID, + ClientSecret: idp.ClientSecret, + Scopes: []string{"openid", "profile", "email", "groups"}, + GroupsClaim: "groups", + UsernameClaim: "preferred_username", + UserProvisioning: true, + TeamSync: true, + ButtonLabel: "Single Sign-On", + } + if mutate != nil { + mutate(&cfg) + } + + codec, err := sso.NewTransactionCodec(bytes.Repeat([]byte{7}, 32)) + require.NoError(t, err) + + h := &oidcHarness{f: f, idp: idp, cfg: cfg, codec: codec, logs: &bytes.Buffer{}} + mux := runtime.NewServeMux() + require.NoError(t, authv1.RegisterAuthServiceHandlerServer(context.Background(), mux, NewAuth(f.users, f.teams, f.keys, cfg))) + NewAuthHTTP(f.users, f.sessions, cfg).Register(mux) + if cfg.OIDC.Enabled() { + provider := sso.NewOIDCProvider(cfg.OIDC, cfg.OIDCRedirectURL()) + h.oidc = NewOIDCHTTP(f.users, f.teams, f.sessions, provider, codec, cfg) + h.oidc.logger = slog.New(slog.NewJSONHandler(h.logs, nil)) + h.oidc.Register(mux) + } + h.handler = auth.HTTPMiddleware(f.resolver, cfg)(mux) + return h +} + +func (h *oidcHarness) get(path string, headers map[string]string, cookies ...*http.Cookie) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, path, nil) + for k, v := range headers { + req.Header.Set(k, v) + } + for _, c := range cookies { + if c != nil { + req.AddCookie(c) + } + } + rec := httptest.NewRecorder() + h.handler.ServeHTTP(rec, req) + return rec +} + +// startRaw calls the login route. +func (h *oidcHarness) startRaw(redirect string) *httptest.ResponseRecorder { + path := auth.OIDCLoginPath + if redirect != "" { + path += "?redirect=" + url.QueryEscape(redirect) + } + return h.get(path, nil) +} + +// start calls the login route and returns the IdP authorization URL and the transaction cookie. +func (h *oidcHarness) start(t *testing.T, redirect string) (*url.URL, *http.Cookie) { + t.Helper() + rec := h.startRaw(redirect) + require.Equal(t, http.StatusFound, rec.Code, rec.Body.String()) + loc, err := url.Parse(rec.Header().Get("Location")) + require.NoError(t, err) + tx := cookieNamed(rec, sso.TransactionCookieName) + require.NotNil(t, tx) + return loc, tx +} + +// callback replays the IdP redirect on the Tracker callback, as a browser coming back cross-site. +func (h *oidcHarness) callback(t *testing.T, cb *url.URL, tx *http.Cookie) *httptest.ResponseRecorder { + t.Helper() + return h.get(cb.RequestURI(), map[string]string{"Sec-Fetch-Site": "cross-site"}, tx) +} + +// login runs start, the IdP authorization and the callback. +func (h *oidcHarness) login(t *testing.T, redirect string) *httptest.ResponseRecorder { + t.Helper() + authURL, tx := h.start(t, redirect) + return h.callback(t, h.idp.Authorize(t, authURL.String()), tx) +} + +type meBody struct { + Authenticated bool `json:"authenticated"` + Username string `json:"username"` + Source string `json:"source"` + Kind string `json:"kind"` + MustChangePassword bool `json:"mustChangePassword"` + IsAdmin bool `json:"isAdmin"` + Teams []struct { + ID string `json:"id"` + Name string `json:"name"` + } `json:"teams"` +} + +// me calls GET /api/v1alpha1/auth/me with a same-origin session cookie. +func (h *oidcHarness) me(t *testing.T, session *http.Cookie) meBody { + t.Helper() + rec := h.get("/api/v1alpha1/auth/me", map[string]string{"Sec-Fetch-Site": "same-origin"}, session) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + var b meBody + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &b)) + return b +} + +// cookieNamed returns the named Set-Cookie of the response, nil when absent. +func cookieNamed(rec *httptest.ResponseRecorder, name string) *http.Cookie { + for _, c := range rec.Result().Cookies() { + if c.Name == name { + return c + } + } + return nil +} + +// oidcLoginCount reads tracker_auth_logins_total{method="oidc"} for one result. +func oidcLoginCount(result string) float64 { + return testutil.ToFloat64(authz.AuthLogins.WithLabelValues(authz.LoginMethodOIDC, result)) +} diff --git a/server/auth_test.go b/server/auth_test.go index c99a18f..4894ff7 100644 --- a/server/auth_test.go +++ b/server/auth_test.go @@ -28,6 +28,18 @@ func TestGetAuthConfigIsPublic(t *testing.T) { assert.False(t, resp.OidcEnabled) assert.Equal(t, []string{"event:read"}, resp.AnonymousPermissions) assert.True(t, resp.DemoMode) + assert.Equal(t, "", resp.OidcButtonLabel) +} + +func TestGetAuthConfigReportsOIDC(t *testing.T) { + f := newAuthFixture(t) + f.cfg.OIDC = auth.OIDCConfig{Issuer: "https://idp.example.com", ButtonLabel: "Sign in with Okta"} + svc := newAuthService(f) + + resp, err := svc.GetAuthConfig(rpcCtx(auth.Anonymous(nil), "GetAuthConfig"), &authv1.GetAuthConfigRequest{}) + require.NoError(t, err) + assert.True(t, resp.OidcEnabled) + assert.Equal(t, "Sign in with Okta", resp.OidcButtonLabel) } func TestMe(t *testing.T) {