diff --git a/docs/auth.md b/docs/auth.md index ec864ef7e9b..04588ae3913 100644 --- a/docs/auth.md +++ b/docs/auth.md @@ -1145,6 +1145,17 @@ The returned OAuth client is a normal Cozy OAuth client: The created OAuth client is bound to the upstream OIDC session so it can be revoked by OIDC backchannel logout. +Repeated exchanges reuse an existing client when the OIDC provider/context, +instance, session (`sid`), and software ID match, including across allowed +origins. The client ID and secret remain the same, while each response contains +tokens with the scope validated for that request. Revoking this client revokes +all tokens issued through it. Different sessions and applications keep separate +clients. + +Reuse relies on the existing OIDC session bindings. If a binding is lost or +expires (after 31 days without renewal in Redis), another client can be created. +Existing duplicate clients are not automatically deleted by token exchange. + ### POST /auth/session_code This endpoint can be used by the flagship application in order to create a diff --git a/web/auth/auth_test.go b/web/auth/auth_test.go index 91068c85a73..ab12d8346cc 100644 --- a/web/auth/auth_test.go +++ b/web/auth/auth_test.go @@ -15,6 +15,8 @@ import ( "net/http" "net/http/httptest" "net/url" + "strings" + "sync/atomic" "testing" "time" @@ -31,21 +33,39 @@ import ( "github.com/cozy/cozy-stack/pkg/consts" "github.com/cozy/cozy-stack/pkg/couchdb" "github.com/cozy/cozy-stack/pkg/crypto" + "github.com/cozy/cozy-stack/pkg/limits" + "github.com/cozy/cozy-stack/pkg/lock" "github.com/cozy/cozy-stack/pkg/metadata" "github.com/cozy/cozy-stack/tests/testutils" "github.com/cozy/cozy-stack/web" "github.com/cozy/cozy-stack/web/apps" + "github.com/cozy/cozy-stack/web/auth" "github.com/cozy/cozy-stack/web/errors" "github.com/cozy/cozy-stack/web/middlewares" "github.com/gavv/httpexpect/v2" "github.com/golang-jwt/jwt/v5" "github.com/labstack/echo/v4" + "github.com/redis/go-redis/v9" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) const domain = "cozy.example.net" +func TestLockOAuthClientFailure(t *testing.T) { + config.UseTestFile(t) + conf := config.GetConfig() + previousLock := conf.Lock + t.Cleanup(func() { conf.Lock = previousLock }) + redisClient := redis.NewClient(&redis.Options{}) + require.NoError(t, redisClient.Close()) + conf.Lock = lock.New(redisClient) + + unlock, err := auth.LockOAuthClient(&instance.Instance{Domain: domain}, "client") + require.ErrorIs(t, err, redis.ErrClosed) + require.Nil(t, unlock) +} + func TestAuth(t *testing.T) { if testing.Short() { t.Skip("an instance is required for this test: test skipped due to the use of --short flag") @@ -2126,6 +2146,11 @@ func TestTokenExchange(t *testing.T) { config.UseTestFile(t) conf := config.GetConfig() + couchTransport := &tokenExchangeCouchDBTransport{RoundTripper: conf.CouchDB.Client.Transport} + if couchTransport.RoundTripper == nil { + couchTransport.RoundTripper = http.DefaultTransport + } + conf.CouchDB.Client.Transport = couchTransport conf.Assets = "../../assets" _ = web.LoadSupportedLocales() if _, err := couchdb.CheckStatus(context.Background()); err != nil { @@ -2221,6 +2246,30 @@ func TestTokenExchange(t *testing.T) { delete(appConfig, "instance_claim") }) } + exchange := func(t *testing.T, sid, exchangeType, scope, origin string, status int) *httpexpect.Object { + t.Helper() + audience := clientID + if exchangeType == "app" { + audience = appTokenAudience + } + idToken := makeTokenExchangeSignedJWT(t, privateKey, kid, map[string]interface{}{ + "iss": issuer, "aud": []string{audience}, "sub": "mail-user", "sid": sid, + "iat": time.Now().Unix(), "exp": time.Now().Add(time.Hour).Unix(), + "org_id": testInstance.OrgID, "org_domain": testInstance.OrgDomain, "org_role": "owner", + }) + response := httpexpect.WithConfig(httpexpect.Config{ + BaseURL: ts.URL, Reporter: httpexpect.NewAssertReporter(t), + }).POST("/auth/token_exchange"). + WithHost(testInstance.Domain). + WithHeader("Accept", "application/json"). + WithHeader("Origin", origin). + WithJSON(map[string]string{"id_token": idToken, "exchange_type": exchangeType, "scope": scope}). + Expect().Status(status) + if status != http.StatusOK { + return nil + } + return response.JSON().Object() + } t.Run("RequiresMandatoryParameters", func(t *testing.T) { e.POST("/auth/token_exchange"). @@ -3149,6 +3198,227 @@ func TestTokenExchange(t *testing.T) { Expect(). Status(http.StatusUnauthorized) }) + + t.Run("ReusesClientWithCurrentScopeAndOriginalCredentials", func(t *testing.T) { + const sid = "reuse-admin-sid" + first := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + id := first.Value("client_id").String().Raw() + client, err := oauth.FindClient(testInstance, id) + require.NoError(t, err) + client.LastRefreshedAt = time.Now().Add(-time.Hour) + require.NoError(t, couchdb.UpdateDoc(testInstance, client)) + + second := exchange(t, sid, "admin", "io.cozy.contacts", "https://admin.example.com", http.StatusOK) + for _, field := range []string{"client_id", "client_secret", "registration_access_token"} { + second.ValueEqual(field, first.Value(field).Raw()) + } + assertValidToken(t, testInstance, first.Value("access_token").String().Raw(), consts.AccessTokenAudience, id, "io.cozy.files") + assertValidToken(t, testInstance, second.Value("access_token").String().Raw(), consts.AccessTokenAudience, id, "io.cozy.contacts") + assertValidToken(t, testInstance, second.Value("refresh_token").String().Raw(), consts.RefreshTokenAudience, id, "io.cozy.contacts") + client, err = oauth.FindClient(testInstance, id) + require.NoError(t, err) + refreshed, err := time.Parse(time.RFC3339Nano, client.LastRefreshedAt.(string)) + require.NoError(t, err) + require.WithinDuration(t, time.Now(), refreshed, time.Minute) + refs, err := oidcbinding.ListOAuthClients(contextName, sid) + require.NoError(t, err) + require.Len(t, refs, 1) + + e.POST("/auth/access_token").WithHost(testInstance.Domain). + WithForm(map[string]string{ + "grant_type": "refresh_token", "client_id": id, + "client_secret": first.Value("client_secret").String().Raw(), + "refresh_token": first.Value("refresh_token").String().Raw(), + }).Expect().Status(http.StatusOK).JSON().Object().ValueEqual("scope", "io.cozy.files") + e.DELETE("/auth/register/"+id).WithHost(testInstance.Domain). + WithHeader("Authorization", "Bearer "+first.Value("registration_access_token").String().Raw()). + Expect().Status(http.StatusNoContent) + e.POST("/auth/access_token").WithHost(testInstance.Domain). + WithForm(map[string]string{ + "grant_type": "refresh_token", "client_id": id, + "client_secret": second.Value("client_secret").String().Raw(), + "refresh_token": second.Value("refresh_token").String().Raw(), + }).Expect().Status(http.StatusBadRequest) + }) + + t.Run("ReusesAppClientAndPreservesSessionIsolation", func(t *testing.T) { + const sid = "reuse-app-sid" + origin := "https://mail." + testInstance.Domain + first := exchange(t, sid, "app", "", origin, http.StatusOK) + second := exchange(t, sid, "app", "", origin, http.StatusOK) + id := first.Value("client_id").String().Raw() + second.ValueEqual("client_id", id) + second.ValueEqual("client_secret", first.Value("client_secret").Raw()) + otherSession := exchange(t, "other-app-sid", "app", "", origin, http.StatusOK) + require.NotEqual(t, id, otherSession.Value("client_id").String().Raw()) + admin := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + require.NotEqual(t, id, admin.Value("client_id").String().Raw()) + otherOrigin := exchange(t, sid, "admin", "io.cozy.files", "https://workspace.sales.example.com", http.StatusOK) + for _, field := range []string{"client_id", "client_secret", "registration_access_token"} { + otherOrigin.ValueEqual(field, admin.Value(field).Raw()) + } + client, err := oauth.FindClient(testInstance, admin.Value("client_id").String().Raw()) + require.NoError(t, err) + require.Equal(t, []string{"https://admin.example.com"}, client.RedirectURIs) + + deleted, err := oauth.DeleteByOIDCSession(contextName, sid) + require.NoError(t, err) + require.Equal(t, 2, deleted) + _, err = oauth.FindClient(testInstance, id) + require.True(t, couchdb.IsNotFoundError(err)) + _, err = oauth.FindClient(testInstance, otherSession.Value("client_id").String().Raw()) + require.NoError(t, err) + }) + + t.Run("ConcurrentExchangesCreateOneClient", func(t *testing.T) { + const sid = "concurrent-exchange-sid" + origins := []string{"https://admin.example.com", "https://workspace.sales.example.com"} + ids := make([]string, 6) + t.Run("Requests", func(t *testing.T) { + for i := range ids { + t.Run(fmt.Sprint(i), func(t *testing.T) { + t.Parallel() + response := exchange(t, sid, "admin", "io.cozy.files", origins[i%len(origins)], http.StatusOK) + ids[i] = response.Value("client_id").String().Raw() + }) + } + }) + for _, id := range ids { + require.NotEmpty(t, id) + require.Equal(t, ids[0], id) + } + refs, err := oidcbinding.ListOAuthClients(contextName, sid) + require.NoError(t, err) + require.Len(t, refs, 1) + }) + + t.Run("DoesNotReuseUnrelatedOrStaleClients", func(t *testing.T) { + for _, mismatch := range []string{"provider", "instance", "session", "software", "pending", "deleted", "missing-binding"} { + t.Run(mismatch, func(t *testing.T) { + sid := "mismatch-" + mismatch + first := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + id := first.Value("client_id").String().Raw() + client, err := oauth.FindClient(testInstance, id) + require.NoError(t, err) + switch mismatch { + case "provider", "instance", "missing-binding": + require.NoError(t, oidcbinding.UnbindOAuthClient(contextName, testInstance.Domain, sid, id)) + if mismatch == "provider" { + require.NoError(t, oidcbinding.BindOAuthClient("other-provider", testInstance.Domain, sid, id)) + } else if mismatch == "instance" { + require.NoError(t, oidcbinding.BindOAuthClient(contextName, "other.example.com", sid, id)) + } + case "deleted": + require.NoError(t, couchdb.DeleteDoc(testInstance, client)) + default: + switch mismatch { + case "session": + client.OIDCSessionID = "other-session" + case "software": + client.SoftwareID = "other-software" + case "pending": + client.Pending = true + } + require.NoError(t, couchdb.UpdateDoc(testInstance, client)) + } + second := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + require.NotEqual(t, id, second.Value("client_id").String().Raw()) + if mismatch != "deleted" { + _, err = oauth.FindClient(testInstance, id) + require.NoError(t, err) + } else { + refs, err := oidcbinding.ListOAuthClients(contextName, sid) + require.NoError(t, err) + require.Len(t, refs, 1) + } + }) + } + }) + + t.Run("ReusesOneExistingDuplicateWithoutDeletingOthers", func(t *testing.T) { + const sid = "existing-duplicate-sid" + first := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + duplicate, err := oauth.FindClient(testInstance, first.Value("client_id").String().Raw()) + require.NoError(t, err) + require.Nil(t, duplicate.Create(testInstance, oauth.NotPending)) + require.NoError(t, oidcbinding.BindOAuthClient(contextName, testInstance.Domain, sid, duplicate.ClientID)) + refs, err := oidcbinding.ListOAuthClients(contextName, sid) + require.NoError(t, err) + require.Len(t, refs, 2) + for range 2 { + exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK). + ValueEqual("client_id", refs[0].OAuthClientID) + } + for _, ref := range refs { + _, err := oauth.FindClient(testInstance, ref.OAuthClientID) + require.NoError(t, err) + } + }) + + t.Run("RegeneratesMissingRegistrationTokenWithoutRegisteringAgain", func(t *testing.T) { + const sid = "missing-registration-token-sid" + first := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + id := first.Value("client_id").String().Raw() + client, err := oauth.FindClient(testInstance, id) + require.NoError(t, err) + client.RegistrationToken = "" + require.NoError(t, couchdb.UpdateDoc(testInstance, client)) + oldLimit := limits.GetMaximumLimit(limits.OAuthClientType) + limits.SetMaximumLimit(limits.OAuthClientType, 0) + t.Cleanup(func() { limits.SetMaximumLimit(limits.OAuthClientType, oldLimit) }) + second := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + second.ValueEqual("client_id", id) + assertValidToken(t, testInstance, second.Value("registration_access_token").String().Raw(), consts.RegistrationTokenAudience, id, "") + }) + + t.Run("StorageFailuresDoNotDeleteReusedClientsOrCreateDuplicates", func(t *testing.T) { + const sid = "reuse-storage-failure-sid" + first := exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK) + id := first.Value("client_id").String().Raw() + before, _, err := oauth.GetAll(testInstance, 1000, "") + require.NoError(t, err) + for _, method := range []string{http.MethodGet, http.MethodPut} { + t.Run(method, func(t *testing.T) { + couchTransport.failMethod.Store(method) + t.Cleanup(func() { couchTransport.failMethod.Store("") }) + exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusServiceUnavailable) + }) + client, err := oauth.FindClient(testInstance, id) + require.NoError(t, err) + require.Equal(t, first.Value("client_secret").String().Raw(), client.ClientSecret) + refs, err := oidcbinding.ListOAuthClients(contextName, sid) + require.NoError(t, err) + require.Len(t, refs, 1) + after, _, err := oauth.GetAll(testInstance, 1000, "") + require.NoError(t, err) + require.Len(t, after, len(before)) + } + exchange(t, sid, "admin", "io.cozy.files", "https://admin.example.com", http.StatusOK).ValueEqual("client_id", id) + }) + + t.Run("FailedBindingDeletesOnlyNewClient", func(t *testing.T) { + before, _, err := oauth.GetAll(testInstance, 1000, "") + require.NoError(t, err) + couchTransport.failMethod.Store(http.MethodPut) + t.Cleanup(func() { couchTransport.failMethod.Store("") }) + exchange(t, "new-client-binding-failure", "admin", "io.cozy.files", "https://admin.example.com", http.StatusServiceUnavailable) + couchTransport.failMethod.Store("") + after, _, err := oauth.GetAll(testInstance, 1000, "") + require.NoError(t, err) + require.Len(t, after, len(before)) + }) +} + +type tokenExchangeCouchDBTransport struct { + http.RoundTripper + failMethod atomic.Value +} + +func (transport *tokenExchangeCouchDBTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if req.Method == transport.failMethod.Load() && strings.Contains(req.URL.Path, couchdb.EscapeCouchdbName(consts.OAuthClients)+"/") { + return nil, fmt.Errorf("simulated OAuth client storage failure") + } + return transport.RoundTripper.RoundTrip(req) } func getLoginCSRFToken(e *httpexpect.Expect) string { diff --git a/web/auth/oauth.go b/web/auth/oauth.go index b27a74ecb4f..08d99caa0db 100644 --- a/web/auth/oauth.go +++ b/web/auth/oauth.go @@ -918,10 +918,12 @@ type AccessTokenReponse struct { Refresh string `json:"refresh_token,omitempty"` } -func LockOAuthClient(inst *instance.Instance, clientID string) func() { +func LockOAuthClient(inst *instance.Instance, clientID string) (func(), error) { mu := config.Lock().ReadWrite(inst, "oauth/"+clientID) - _ = mu.Lock() - return mu.Unlock + if err := mu.Lock(); err != nil { + return nil, err + } + return mu.Unlock, nil } func accessToken(c echo.Context) error { @@ -946,7 +948,11 @@ func accessToken(c echo.Context) error { "error": "the client_secret parameter is mandatory", }) } - defer LockOAuthClient(instance, clientID)() + unlock, err := LockOAuthClient(instance, clientID) + if err != nil { + return err + } + defer unlock() client, err := oauth.FindClient(instance, clientID) if err != nil { diff --git a/web/auth/register.go b/web/auth/register.go index b3a38de6f0c..8af8ecb42af 100644 --- a/web/auth/register.go +++ b/web/auth/register.go @@ -53,7 +53,11 @@ func updateClient(c echo.Context) error { } clientID := c.Param("client-id") - defer LockOAuthClient(instance, clientID)() + unlock, err := LockOAuthClient(instance, clientID) + if err != nil { + return err + } + defer unlock() oldClient, err := oauth.FindClient(instance, clientID) if err != nil { @@ -80,7 +84,11 @@ func updateClient(c echo.Context) error { func deleteClient(c echo.Context) error { instance := middlewares.GetInstance(c) clientID := c.Param("client-id") - defer LockOAuthClient(instance, clientID)() + unlock, err := LockOAuthClient(instance, clientID) + if err != nil { + return err + } + defer unlock() client, err := oauth.FindClient(instance, clientID) if err != nil { diff --git a/web/auth/token_exchange.go b/web/auth/token_exchange.go index 6ffa590a50f..b81e89922a5 100644 --- a/web/auth/token_exchange.go +++ b/web/auth/token_exchange.go @@ -81,25 +81,68 @@ func executeTokenExchange(c echo.Context, inst *instance.Instance, req tokenExch return nil, echo.NewHTTPError(http.StatusBadRequest, err.Error()) } - var client *oauth.Client + var params *tokenExchangeOAuthClientParams var scope string if req.ExchangeType == tokenExchangeTypeApp { if validated.AppConfig == nil { return nil, echo.NewHTTPError(http.StatusBadRequest, "invalid token audience") } - client, scope, err = executeTokenExchangeApp(c, inst, req, *validated.AppConfig) + params, scope, err = tokenExchangeAppClientParams(inst, req, *validated.AppConfig) } else { - client, scope, err = executeTokenExchangeAdmin(c, inst, req) + params, scope, err = tokenExchangeAdminClientParams(req) } if err != nil { return nil, err } - defer LockOAuthClient(inst, client.ClientID)() + // Two locks guard the exchange. This session-scoped lock serialises all + // exchanges for the same OIDC session, so concurrent requests reuse a single + // client instead of each creating their own. The per-client LockOAuthClient + // below then serialises this exchange against client refresh/revoke/update. + // The lock getter already namespaces by instance (DBPrefix), so the session + // id is the only discriminator needed here. + sessionID, _ := tokenExchangeClaimString(validated.Claims, "sid") + mu := config.Lock().ReadWrite(inst, fmt.Sprintf("token-exchange/%q", sessionID)) + if err := mu.Lock(); err != nil { + return nil, err + } + defer mu.Unlock() + + client, err := findTokenExchangeOAuthClient(inst, sessionID, params.SoftwareID) + if err != nil { + return nil, err + } + created := client == nil + if created { + client, err = createTokenExchangeOAuthClient(c, inst, *params) + if err != nil { + return nil, err + } + } + + unlock, err := LockOAuthClient(inst, client.ClientID) + if err != nil { + return nil, err + } + defer unlock() + + if !created { + // Re-read the reused client under its own lock to pick up the current + // revision and reject it if it changed since we looked it up. + client, err = oauth.FindClient(inst, client.ClientID) + if err != nil { + return nil, err + } + if !isTokenExchangeOAuthClient(client, sessionID, params.SoftwareID) { + return nil, echo.NewHTTPError(http.StatusConflict, "OAuth client changed during token exchange") + } + } if err := bindTokenExchangeOIDCSession(inst, client, validated.Claims); err != nil { - if delErr := client.Delete(inst); delErr != nil { - inst.Logger().WithNamespace("oidc").Warnf("Cannot delete orphaned OAuth client %s: %s", client.CouchID, delErr.Description) + if created { + if delErr := client.Delete(inst); delErr != nil { + inst.Logger().WithNamespace("oidc").Warnf("Cannot delete orphaned OAuth client %s: %s", client.CouchID, delErr.Description) + } } return nil, err } @@ -107,22 +150,18 @@ func executeTokenExchange(c echo.Context, inst *instance.Instance, req tokenExch return buildTokenExchangeResponse(inst, client, scope) } -func executeTokenExchangeAdmin(c echo.Context, inst *instance.Instance, req tokenExchangeRequest) (*oauth.Client, string, error) { +func tokenExchangeAdminClientParams(req tokenExchangeRequest) (*tokenExchangeOAuthClientParams, string, error) { if err := validateTokenExchangeScope(req.Scope); err != nil { return nil, "", echo.NewHTTPError(http.StatusBadRequest, err.Error()) } - client, err := createTokenExchangeOAuthClient(c, inst, tokenExchangeOAuthClientParams{ + return &tokenExchangeOAuthClientParams{ ClientName: tokenExchangeOAuthClientName, SoftwareID: tokenExchangeOAuthClientSoftwareID, - }) - if err != nil { - return nil, "", err - } - return client, req.Scope, nil + }, req.Scope, nil } -func executeTokenExchangeApp(c echo.Context, inst *instance.Instance, req tokenExchangeRequest, appConfig config.OIDCAppTokenExchangeAppConfig) (*oauth.Client, string, error) { +func tokenExchangeAppClientParams(inst *instance.Instance, req tokenExchangeRequest, appConfig config.OIDCAppTokenExchangeAppConfig) (*tokenExchangeOAuthClientParams, string, error) { if req.Scope != "" { return nil, "", echo.NewHTTPError(http.StatusBadRequest, "scope is not allowed for app token exchange") } @@ -138,17 +177,12 @@ func executeTokenExchangeApp(c echo.Context, inst *instance.Instance, req tokenE if err := tokenExchangeAssertManifestTrusted(manifest, slug); err != nil { return nil, "", err } - client, err := createTokenExchangeOAuthClient(c, inst, tokenExchangeOAuthClientParams{ + return &tokenExchangeOAuthClientParams{ AppSlug: slug, ClientName: tokenExchangeLinkedAppClientName(manifest, slug), SoftwareID: appConfig.SoftwareID, SoftwareIDPrevalidated: true, - }) - if err != nil { - return nil, "", err - } - - return client, oauth.BuildLinkedAppScope(slug), nil + }, oauth.BuildLinkedAppScope(slug), nil } // tokenExchangeAssertManifestTrusted confirms the locally installed manifest @@ -349,6 +383,37 @@ func tokenExchangeCheckAppInstance(conf *oidcprovider.Config, inst *instance.Ins return nil } +func findTokenExchangeOAuthClient(inst *instance.Instance, sessionID, softwareID string) (*oauth.Client, error) { + refs, err := oidcbinding.ListOAuthClients(inst.ContextName, sessionID) + if err != nil { + return nil, err + } + // Scan this session's bindings; index by application if sessions hold many clients. + for _, ref := range refs { + if ref.Domain != inst.Domain { + continue + } + client, err := oauth.FindClient(inst, ref.OAuthClientID) + if couchdb.IsNotFoundError(err) { + if err := oidcbinding.UnbindOAuthClient(inst.ContextName, inst.Domain, sessionID, ref.OAuthClientID); err != nil { + return nil, err + } + continue + } + if err != nil { + return nil, err + } + if isTokenExchangeOAuthClient(client, sessionID, softwareID) { + return client, nil + } + } + return nil, nil +} + +func isTokenExchangeOAuthClient(client *oauth.Client, sessionID, softwareID string) bool { + return !client.Pending && client.OIDCSessionID == sessionID && client.SoftwareID == softwareID +} + func createTokenExchangeOAuthClient(c echo.Context, inst *instance.Instance, params tokenExchangeOAuthClientParams) (*oauth.Client, error) { redirectURI, err := tokenExchangeRedirectURI(c, inst, params.AppSlug) if err != nil { @@ -417,6 +482,13 @@ func bindTokenExchangeOIDCSession(inst *instance.Instance, client *oauth.Client, } func buildTokenExchangeResponse(inst *instance.Instance, client *oauth.Client, scope string) (*tokenExchangeResponse, error) { + if client.RegistrationToken == "" { + registrationToken, err := client.CreateJWT(inst, consts.RegistrationTokenAudience, "") + if err != nil { + return nil, echo.NewHTTPError(http.StatusInternalServerError, "Can't generate registration token") + } + client.RegistrationToken = registrationToken + } out := &tokenExchangeResponse{ AccessTokenReponse: AccessTokenReponse{ Type: "bearer", @@ -510,8 +582,7 @@ func tokenExchangeClaimString(claims jwt.MapClaims, key string) (string, bool) { // // For app exchanges (appSlug != ""), the redirect URI is always the app's // own subdomain on this instance: callers cannot influence it via the -// Origin header because executeTokenExchangeApp has already enforced that -// the Origin (if any) is exactly that subdomain. +// Origin header; the request handler has already checked that Origin is allowed. // // For admin exchanges, the Origin header is honoured when it points // somewhere other than the instance itself, which lets the admin panel host diff --git a/web/settings/clients.go b/web/settings/clients.go index adf5dfa981a..c07ddcd2352 100644 --- a/web/settings/clients.go +++ b/web/settings/clients.go @@ -86,7 +86,11 @@ func (h *HTTPHandler) revokeClient(c echo.Context) error { } clientID := c.Param("id") - defer auth.LockOAuthClient(instance, clientID)() + unlock, err := auth.LockOAuthClient(instance, clientID) + if err != nil { + return err + } + defer unlock() client, err := oauth.FindClient(instance, clientID) if err != nil { @@ -112,7 +116,11 @@ func (h *HTTPHandler) synchronized(c echo.Context) error { return err } - defer auth.LockOAuthClient(instance, claims.Subject)() + unlock, err := auth.LockOAuthClient(instance, claims.Subject) + if err != nil { + return err + } + defer unlock() client, err := oauth.FindClient(instance, claims.Subject) if err != nil {