diff --git a/source/api/src/main/kotlin/com/clerk/api/Clerk.kt b/source/api/src/main/kotlin/com/clerk/api/Clerk.kt index a57637671..184071185 100644 --- a/source/api/src/main/kotlin/com/clerk/api/Clerk.kt +++ b/source/api/src/main/kotlin/com/clerk/api/Clerk.kt @@ -31,6 +31,7 @@ import com.clerk.api.network.serialization.ClerkResult import com.clerk.api.organizations.Organization import com.clerk.api.organizations.OrganizationMembership import com.clerk.api.session.Session +import com.clerk.api.session.SessionTokenFetcher import com.clerk.api.session.SessionTokensCache import com.clerk.api.sharedsession.SharedSessionSyncCoordinator import com.clerk.api.sharedsession.SharedSessionSyncProvider @@ -751,6 +752,7 @@ object Clerk { StorageHelper.deleteValue(StorageKey.DEVICE_TOKEN) StorageHelper.deleteValue(StorageKey.SHARED_SESSION_SYNC_SNAPSHOT) clearSessionAndUserState() + SessionTokenFetcher.shared.reset() SessionTokensCache.clear() SSOService.cancelPendingAuthentication() ExternalAccountService.cancelPendingExternalAccountConnection() diff --git a/source/api/src/main/kotlin/com/clerk/api/configuration/ConfigurationManager.kt b/source/api/src/main/kotlin/com/clerk/api/configuration/ConfigurationManager.kt index 37fe8873c..ec23cb550 100644 --- a/source/api/src/main/kotlin/com/clerk/api/configuration/ConfigurationManager.kt +++ b/source/api/src/main/kotlin/com/clerk/api/configuration/ConfigurationManager.kt @@ -21,6 +21,7 @@ import com.clerk.api.network.model.error.ClerkErrorResponse import com.clerk.api.network.serialization.ClerkResult import com.clerk.api.network.serialization.fold import com.clerk.api.session.GetTokenOptions +import com.clerk.api.session.SessionTokenFetcher import com.clerk.api.session.SessionTokensCache import com.clerk.api.session.fetchToken import com.clerk.api.sso.SSOService @@ -348,6 +349,7 @@ internal class ConfigurationManager { StorageHelper.deleteValue(StorageKey.DEVICE_TOKEN) Clerk.updateClient(Client()) Clerk.clearSessionAndUserState() + SessionTokenFetcher.shared.reset() SessionTokensCache.clear() val result = diff --git a/source/api/src/main/kotlin/com/clerk/api/network/api/SessionApi.kt b/source/api/src/main/kotlin/com/clerk/api/network/api/SessionApi.kt index d4efdb484..6cb979d10 100644 --- a/source/api/src/main/kotlin/com/clerk/api/network/api/SessionApi.kt +++ b/source/api/src/main/kotlin/com/clerk/api/network/api/SessionApi.kt @@ -9,6 +9,7 @@ import com.clerk.api.network.serialization.ClerkResult import com.clerk.api.session.Session import com.clerk.api.session.SessionVerification import retrofit2.http.DELETE +import retrofit2.http.Field import retrofit2.http.FieldMap import retrofit2.http.FormUrlEncoded import retrofit2.http.GET @@ -76,8 +77,12 @@ internal interface SessionApi { * failure */ @POST(ApiPaths.Client.Sessions.TOKENS) + @FormUrlEncoded suspend fun tokens( - @Path(ApiParams.ID) sessionId: String + @Path(ApiParams.ID) sessionId: String, + @Field("organization_id") organizationId: String = "", + @Field("token") token: String? = null, + @Field("force_origin") forceOrigin: String? = null, ): ClerkResult /** diff --git a/source/api/src/main/kotlin/com/clerk/api/network/model/environment/AuthConfig.kt b/source/api/src/main/kotlin/com/clerk/api/network/model/environment/AuthConfig.kt index 365e50a8d..04c9bff71 100644 --- a/source/api/src/main/kotlin/com/clerk/api/network/model/environment/AuthConfig.kt +++ b/source/api/src/main/kotlin/com/clerk/api/network/model/environment/AuthConfig.kt @@ -11,6 +11,7 @@ import kotlinx.serialization.Serializable * * @property singleSessionMode Whether the application is configured for single session mode. When * true, only one active session is allowed per user at a time. + * @property sessionMinter Whether session token minting at the edge is enabled. */ @Serializable internal data class AuthConfig( @@ -18,5 +19,8 @@ internal data class AuthConfig( * Whether the application is configured for single session mode. When true, only one active * session is allowed per user at a time. */ - @SerialName("single_session_mode") val singleSessionMode: Boolean + @SerialName("single_session_mode") val singleSessionMode: Boolean, + + /** Whether session token minting at the edge is enabled. */ + @SerialName("session_minter") val sessionMinter: Boolean = false, ) diff --git a/source/api/src/main/kotlin/com/clerk/api/session/Session.kt b/source/api/src/main/kotlin/com/clerk/api/session/Session.kt index 9aa56a75d..41609e779 100644 --- a/source/api/src/main/kotlin/com/clerk/api/session/Session.kt +++ b/source/api/src/main/kotlin/com/clerk/api/session/Session.kt @@ -213,7 +213,7 @@ suspend fun Session.delete(): ClerkResult { suspend fun Session.fetchToken( options: GetTokenOptions = GetTokenOptions() ): ClerkResult { - val token = SessionTokenFetcher().getToken(this, options) + val token = SessionTokenFetcher.shared.getToken(this, options) return if (token != null) { ClerkResult.success(token) } else { diff --git a/source/api/src/main/kotlin/com/clerk/api/session/SessionTokenFetcher.kt b/source/api/src/main/kotlin/com/clerk/api/session/SessionTokenFetcher.kt index 462881d8b..01ca575b7 100644 --- a/source/api/src/main/kotlin/com/clerk/api/session/SessionTokenFetcher.kt +++ b/source/api/src/main/kotlin/com/clerk/api/session/SessionTokenFetcher.kt @@ -26,8 +26,10 @@ import kotlinx.coroutines.Deferred * @param jwtManager The JWT manager used for token parsing and validation */ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManagerImpl()) { - private companion object { - val sessionInvalidationErrorCodes = + internal companion object { + internal val shared: SessionTokenFetcher by lazy { SessionTokenFetcher() } + + private val sessionInvalidationErrorCodes = setOf( "session_revoked", "session_expired", @@ -40,9 +42,23 @@ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManag ) } + private data class FetchContext( + val session: Session, + val cacheKey: String, + val sessionMinterEnabled: Boolean, + ) + /** Map of cache keys to deferred token fetch tasks for request deduplication */ private val tokenTasks = ConcurrentHashMap>() + /** + * Releases deduplicated waiters and removes requests registered by the previous Clerk runtime. + */ + internal fun reset() { + tokenTasks.values.forEach { it.cancel() } + tokenTasks.clear() + } + /** * Retrieves a token for the specified session with the given options. * @@ -76,19 +92,24 @@ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManag session: Session, options: GetTokenOptions, ): TokenResource? { - val cacheKey = session.tokenCacheKey(options.template) + val context = makeFetchContext(session, options.template) ClerkLog.d( - "Fetching token for session ${session.id} with options: $options and cache key: $cacheKey" + "Fetching token for session ${context.session.id} with options: $options and cache key: " + + context.cacheKey ) - return tokenTasks[cacheKey]?.await() + if (options.skipCache) { + return fetchToken(context, options) + } + + return tokenTasks[context.cacheKey]?.await() ?: run { val deferred = CompletableDeferred() - val existingTask = tokenTasks.putIfAbsent(cacheKey, deferred) + val existingTask = tokenTasks.putIfAbsent(context.cacheKey, deferred) existingTask?.await() ?: try { - fetchToken(session, options).also { deferred.complete(it) } + fetchToken(context, options).also { deferred.complete(it) } } catch (e: CancellationException) { deferred.cancel(e) throw e @@ -96,11 +117,21 @@ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManag deferred.completeExceptionally(t) throw t } finally { - tokenTasks.remove(cacheKey, deferred) + tokenTasks.remove(context.cacheKey, deferred) } } } + private fun makeFetchContext(session: Session, template: String?): FetchContext { + val currentSession = + Clerk.clientFlow.value?.sessions?.firstOrNull { it.id == session.id } ?: session + return FetchContext( + session = currentSession, + cacheKey = currentSession.tokenCacheKey(template), + sessionMinterEnabled = Clerk.environment?.authConfig?.sessionMinter == true, + ) + } + /** * Internal method to fetch a token from cache or network. * @@ -112,8 +143,17 @@ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManag * @param options Options controlling the fetch behavior * @return The token resource, or null if the fetch failed */ - private suspend fun fetchToken(session: Session, options: GetTokenOptions): TokenResource? { - val cacheKey = session.tokenCacheKey(options.template) + private suspend fun fetchToken(context: FetchContext, options: GetTokenOptions): TokenResource? { + val session = context.session + val cacheKey = context.cacheKey + + if (options.template == null) { + session.lastActiveToken + ?.takeIf { + TokenFreshness.matches(it, session.id, session.lastActiveOrganizationId) + } + ?.let { SessionTokensCache.hydrate(cacheKey, it) } + } // Check cache first (unless skipped) if (!options.skipCache) { @@ -134,12 +174,25 @@ internal class SessionTokenFetcher(private val jwtManager: JWTManager = JWTManag if (options.template != null) { ClerkApi.session.tokens(session.id, options.template) } else { - ClerkApi.session.tokens(session.id) + val cachedToken = SessionTokensCache.getToken(cacheKey) + val previousToken = + cachedToken?.let { + TokenFreshness.pickFreshest( + existing = session.lastActiveToken, + incoming = it, + ) + } ?: session.lastActiveToken + ClerkApi.session.tokens( + sessionId = session.id, + organizationId = session.lastActiveOrganizationId.orEmpty(), + token = previousToken?.jwt.takeIf { context.sessionMinterEnabled }, + forceOrigin = "true".takeIf { context.sessionMinterEnabled && options.skipCache }, + ) } when (tokensRequest) { is ClerkResult.Success -> { - SessionTokensCache.setToken(cacheKey, tokensRequest.value) + SessionTokensCache.storeIfFresher(cacheKey, tokensRequest.value) tokensRequest.value } is ClerkResult.Failure -> { @@ -228,10 +281,11 @@ data class GetTokenOptions( /** * Extension function to generate a cache key for session tokens. * - * This function creates a unique cache key based on the session ID and optional template name. This - * ensures that tokens for different templates are cached separately. + * This function creates a unique cache key based on the session ID and either its active + * organization or the optional template name. * * @param template Optional template name to include in the cache key * @return A unique cache key string for the session and template combination */ -internal fun Session.tokenCacheKey(template: String?): String = template?.let { "$id-$it" } ?: id +internal fun Session.tokenCacheKey(template: String?): String = + template?.let { "$id-template-$it" } ?: "$id-organization-${lastActiveOrganizationId.orEmpty()}" diff --git a/source/api/src/main/kotlin/com/clerk/api/session/SessionTokensCache.kt b/source/api/src/main/kotlin/com/clerk/api/session/SessionTokensCache.kt index 7c2b9314a..88a2dd8b9 100644 --- a/source/api/src/main/kotlin/com/clerk/api/session/SessionTokensCache.kt +++ b/source/api/src/main/kotlin/com/clerk/api/session/SessionTokensCache.kt @@ -6,6 +6,11 @@ import java.util.concurrent.ConcurrentHashMap internal object SessionTokensCache { private val cache = ConcurrentHashMap() + internal data class StoreResult( + val canonicalToken: TokenResource, + val didChangeCanonicalToken: Boolean, + ) + /** Returns a session token for the given cache key. */ internal fun getToken(cacheKey: String): TokenResource? = cache[cacheKey] @@ -14,6 +19,35 @@ internal object SessionTokensCache { cache[cacheKey] = token } + /** Reconciles a session snapshot without replacing an equally fresh canonical token. */ + internal fun hydrate(cacheKey: String, token: TokenResource) { + cache.compute(cacheKey) { _, existing -> + TokenFreshness.pickFreshest( + existing = existing, + incoming = token, + tieBreaker = TokenFreshness.TieBreaker.EXISTING, + ) + } + } + + /** Atomically stores [token] unless the cache already contains a fresher token. */ + internal fun storeIfFresher( + cacheKey: String, + token: TokenResource, + nowMillis: Long = System.currentTimeMillis(), + ): StoreResult { + var didChangeCanonicalToken = false + val canonicalToken = + checkNotNull( + cache.compute(cacheKey) { _, existing -> + TokenFreshness.pickFreshest(existing, token, nowMillis).also { canonical -> + didChangeCanonicalToken = existing?.jwt != canonical.jwt + } + } + ) + return StoreResult(canonicalToken, didChangeCanonicalToken) + } + /** Removes a session token for the given cache key. */ internal fun removeToken(cacheKey: String): TokenResource? = cache.remove(cacheKey) diff --git a/source/api/src/main/kotlin/com/clerk/api/session/TokenFreshness.kt b/source/api/src/main/kotlin/com/clerk/api/session/TokenFreshness.kt new file mode 100644 index 000000000..b48f2a279 --- /dev/null +++ b/source/api/src/main/kotlin/com/clerk/api/session/TokenFreshness.kt @@ -0,0 +1,128 @@ +package com.clerk.api.session + +import com.auth0.android.jwt.JWT +import com.clerk.api.network.model.token.TokenResource + +/** Chooses the canonical token when session minter responses arrive out of order. */ +internal object TokenFreshness { + private data class DecodedToken(val resource: TokenResource, val jwt: JWT) + + internal enum class TieBreaker { + EXISTING, + INCOMING, + } + + internal fun pickFreshest( + existing: TokenResource?, + incoming: TokenResource, + nowMillis: Long = System.currentTimeMillis(), + tieBreaker: TieBreaker = TieBreaker.INCOMING, + ): TokenResource { + existing ?: return incoming + + val existingJwt = decode(existing.jwt) + val incomingJwt = decode(incoming.jwt) + return when { + existingJwt != null && incomingJwt != null -> + pickFreshestDecoded( + existing = DecodedToken(existing, existingJwt), + incoming = DecodedToken(incoming, incomingJwt), + nowMillis = nowMillis, + tieBreaker = tieBreaker, + ) + existingJwt != null -> existing + else -> incoming + } + } + + internal fun matches( + token: TokenResource, + sessionId: String, + organizationId: String?, + ): Boolean { + val jwt = decode(token.jwt) + val tokenSessionId = jwt?.getClaim("sid")?.asString() + return if (jwt == null || tokenSessionId == null) { + true + } else { + tokenSessionId == sessionId && jwt.organizationId().orEmpty() == organizationId.orEmpty() + } + } + + private fun pickFreshestDecoded( + existing: DecodedToken, + incoming: DecodedToken, + nowMillis: Long, + tieBreaker: TieBreaker, + ): TokenResource = + if (!haveMatchingContext(existing.jwt, incoming.jwt)) { + incoming.resource + } else { + pickByExpiration(existing, incoming, nowMillis) + ?: pickByOriginIssuedAt(existing, incoming, tieBreaker) + } + + private fun haveMatchingContext(existing: JWT, incoming: JWT): Boolean = + existing.getClaim("sid").asString() == incoming.getClaim("sid").asString() && + existing.organizationId().orEmpty() == incoming.organizationId().orEmpty() + + private fun pickByExpiration( + existing: DecodedToken, + incoming: DecodedToken, + nowMillis: Long, + ): TokenResource? { + val existingIsExpired = existing.jwt.expiresAt?.time?.let { it <= nowMillis } + val incomingIsExpired = incoming.jwt.expiresAt?.time?.let { it <= nowMillis } + return when { + existingIsExpired == true && incomingIsExpired == false -> incoming.resource + existingIsExpired == false && incomingIsExpired == true -> existing.resource + else -> null + } + } + + private fun pickByOriginIssuedAt( + existing: DecodedToken, + incoming: DecodedToken, + tieBreaker: TieBreaker, + ): TokenResource { + val existingOriginIssuedAt = existing.jwt.originIssuedAt() + val incomingOriginIssuedAt = incoming.jwt.originIssuedAt() + return when { + existingOriginIssuedAt == null && incomingOriginIssuedAt == null -> + pickByIssuedAt(existing, incoming, tieBreaker) + existingOriginIssuedAt != null && incomingOriginIssuedAt == null -> existing.resource + existingOriginIssuedAt == null -> incoming.resource + existingOriginIssuedAt > checkNotNull(incomingOriginIssuedAt) -> existing.resource + incomingOriginIssuedAt > existingOriginIssuedAt -> incoming.resource + else -> pickByIssuedAt(existing, incoming, tieBreaker) + } + } + + private fun pickByIssuedAt( + existing: DecodedToken, + incoming: DecodedToken, + tieBreaker: TieBreaker, + ): TokenResource { + val existingIssuedAt = existing.jwt.issuedAt?.time ?: 0 + val incomingIssuedAt = incoming.jwt.issuedAt?.time ?: 0 + return when { + existingIssuedAt > incomingIssuedAt -> existing.resource + incomingIssuedAt > existingIssuedAt -> incoming.resource + tieBreaker == TieBreaker.EXISTING -> existing.resource + else -> incoming.resource + } + } + + private fun decode(token: String): JWT? = + try { + JWT(token) + } catch (_: Exception) { + null + } + + private fun JWT.originIssuedAt(): Long? = header["oiat"]?.toLongOrNull() + + private fun JWT.organizationId(): String? = + getClaim("org_id").asString() + ?: runCatching { getClaim("o").asObject(Map::class.java)?.get("id") as? String }.getOrNull() +} diff --git a/source/api/src/test/java/com/clerk/api/network/model/environment/AuthConfigTest.kt b/source/api/src/test/java/com/clerk/api/network/model/environment/AuthConfigTest.kt new file mode 100644 index 000000000..438c74cb9 --- /dev/null +++ b/source/api/src/test/java/com/clerk/api/network/model/environment/AuthConfigTest.kt @@ -0,0 +1,23 @@ +package com.clerk.api.network.model.environment + +import kotlinx.serialization.json.Json +import org.junit.Assert.assertFalse +import org.junit.Assert.assertTrue +import org.junit.Test + +class AuthConfigTest { + @Test + fun `decodes session minter flag`() { + val authConfig = + Json.decodeFromString("""{"single_session_mode":false,"session_minter":true}""") + + assertTrue(authConfig.sessionMinter) + } + + @Test + fun `defaults session minter to false when omitted`() { + val authConfig = Json.decodeFromString("""{"single_session_mode":false}""") + + assertFalse(authConfig.sessionMinter) + } +} diff --git a/source/api/src/test/java/com/clerk/api/session/SessionTokenFetcherTest.kt b/source/api/src/test/java/com/clerk/api/session/SessionTokenFetcherTest.kt index 7cdf31de1..5dd6b38e1 100644 --- a/source/api/src/test/java/com/clerk/api/session/SessionTokenFetcherTest.kt +++ b/source/api/src/test/java/com/clerk/api/session/SessionTokenFetcherTest.kt @@ -4,6 +4,8 @@ import com.auth0.android.jwt.JWT import com.clerk.api.Clerk import com.clerk.api.network.ClerkApi import com.clerk.api.network.api.SessionApi +import com.clerk.api.network.model.environment.AuthConfig +import com.clerk.api.network.model.environment.Environment import com.clerk.api.network.model.error.ClerkErrorResponse import com.clerk.api.network.model.error.Error import com.clerk.api.network.model.token.TokenResource @@ -17,6 +19,8 @@ import io.mockk.slot import io.mockk.unmockkAll import io.mockk.verify import java.util.Date +import java.util.concurrent.atomic.AtomicInteger +import kotlinx.coroutines.CompletableDeferred import kotlinx.coroutines.async import kotlinx.coroutines.delay import kotlinx.coroutines.test.runTest @@ -67,6 +71,7 @@ class SessionTokenFetcherTest { every { Clerk.clearSessionAndUserState() } returns Unit // Mock SessionTokensCache + SessionTokensCache.clear() mockkObject(SessionTokensCache) } @@ -78,7 +83,7 @@ class SessionTokenFetcherTest { @Test fun `getToken returns cached token if valid and cache not skipped`() = runTest { // Given - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" val futureTime = Date(System.currentTimeMillis() + 120000) // 2 minutes from now every { mockTokenResource.jwt } returns "valid.jwt.token" @@ -97,13 +102,14 @@ class SessionTokenFetcherTest { @Test fun `getToken fetches from network if cache is empty`() = runTest { // Given - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" val setTokenSlot = slot() coEvery { SessionTokensCache.getToken(cacheKey) } returns null coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, capture(setTokenSlot)) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, capture(setTokenSlot), any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession) @@ -112,14 +118,14 @@ class SessionTokenFetcherTest { assertEquals(mockTokenResource, result) coVerify { SessionTokensCache.getToken(cacheKey) } coVerify { mockClerkApiService.tokens("session_123") } - coVerify { SessionTokensCache.setToken(cacheKey, mockTokenResource) } + coVerify { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } assertEquals(mockTokenResource, setTokenSlot.captured) } @Test fun `getToken fetches from network if cached token is expired`() = runTest { // Given - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" val pastTime = Date(System.currentTimeMillis() - 60000) // 1 minute ago val freshToken = mockk(relaxed = true) @@ -127,7 +133,8 @@ class SessionTokenFetcherTest { every { mockJWT.expiresAt } returns pastTime coEvery { SessionTokensCache.getToken(cacheKey) } returns mockTokenResource coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(freshToken) - coEvery { SessionTokensCache.setToken(cacheKey, freshToken) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, freshToken, any()) } returns + SessionTokensCache.StoreResult(freshToken, true) // When val result = sessionTokenFetcher.getToken(mockSession) @@ -136,20 +143,21 @@ class SessionTokenFetcherTest { assertEquals(freshToken, result) coVerify { SessionTokensCache.getToken(cacheKey) } coVerify { mockClerkApiService.tokens("session_123") } - coVerify { SessionTokensCache.setToken(cacheKey, freshToken) } + coVerify { SessionTokensCache.storeIfFresher(cacheKey, freshToken, any()) } } @Test fun `getToken uses template in API call when provided`() = runTest { // Given val template = "custom_template" - val cacheKey = "session_123-custom_template" + val cacheKey = "session_123-template-custom_template" val options = GetTokenOptions(template = template) coEvery { SessionTokensCache.getToken(cacheKey) } returns null coEvery { mockClerkApiService.tokens("session_123", template) } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession, options) @@ -157,27 +165,69 @@ class SessionTokenFetcherTest { // Then assertEquals(mockTokenResource, result) coVerify { mockClerkApiService.tokens("session_123", template) } - coVerify { SessionTokensCache.setToken(cacheKey, mockTokenResource) } + coVerify { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } } @Test - fun `getToken skips cache when skipCache is true`() = runTest { + fun `getToken bypasses cached result when skipCache is true`() = runTest { // Given val options = GetTokenOptions(skipCache = true) - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" + coEvery { SessionTokensCache.getToken(cacheKey) } returns null coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession, options) // Then assertEquals(mockTokenResource, result) - coVerify(exactly = 0) { SessionTokensCache.getToken(any()) } + coVerify { SessionTokensCache.getToken(cacheKey) } coVerify { mockClerkApiService.tokens("session_123") } - coVerify { SessionTokensCache.setToken(cacheKey, mockTokenResource) } + coVerify { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } + } + + @Test + fun `session minter passes previous token and forces origin for forced refresh`() = runTest { + // Given + val cacheKey = "session_123-organization-org_123" + val previousToken = TokenResource("previous.token.value") + val environment = mockk() + every { environment.authConfig } returns + AuthConfig(singleSessionMode = false, sessionMinter = true) + every { Clerk.environment } returns environment + every { mockSession.lastActiveOrganizationId } returns "org_123" + every { mockSession.lastActiveToken } returns previousToken + every { SessionTokensCache.hydrate(cacheKey, previousToken) } returns Unit + every { SessionTokensCache.getToken(cacheKey) } returns null + coEvery { + mockClerkApiService.tokens( + sessionId = "session_123", + organizationId = "org_123", + token = previousToken.jwt, + forceOrigin = "true", + ) + } returns ClerkResult.success(mockTokenResource) + every { + SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) + } returns SessionTokensCache.StoreResult(mockTokenResource, true) + + // When + val result = sessionTokenFetcher.getToken(mockSession, GetTokenOptions(skipCache = true)) + + // Then + assertEquals(mockTokenResource, result) + coVerify { + mockClerkApiService.tokens( + sessionId = "session_123", + organizationId = "org_123", + token = previousToken.jwt, + forceOrigin = "true", + ) + } } @Test @@ -201,7 +251,7 @@ class SessionTokenFetcherTest { // Then assertNull(result) coVerify { mockClerkApiService.tokens("session_123") } - coVerify(exactly = 0) { SessionTokensCache.setToken(any(), any()) } + coVerify(exactly = 0) { SessionTokensCache.storeIfFresher(any(), any(), any()) } verify(exactly = 0) { Clerk.clearSessionAndUserState() } } @@ -266,7 +316,7 @@ class SessionTokenFetcherTest { // Given val customBuffer = 120L // 2 minutes val options = GetTokenOptions(expirationBuffer = customBuffer) - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" // Token expires in 90 seconds (less than 2-minute buffer) val soonExpiredTime = Date(System.currentTimeMillis() + 90000) @@ -275,7 +325,8 @@ class SessionTokenFetcherTest { coEvery { SessionTokensCache.getToken(cacheKey) } returns mockTokenResource coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession, options) @@ -289,14 +340,15 @@ class SessionTokenFetcherTest { @Test fun `getToken handles JWT parsing exception gracefully`() = runTest { // Given - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" every { mockTokenResource.jwt } returns "invalid.jwt.token" every { mockJWTManager.createFromString(any()) } throws RuntimeException("Invalid JWT") coEvery { SessionTokensCache.getToken(cacheKey) } returns mockTokenResource coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession) @@ -310,7 +362,7 @@ class SessionTokenFetcherTest { @Test fun `getToken handles concurrent requests properly`() = runTest { // Given - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" coEvery { SessionTokensCache.getToken(cacheKey) } returns null coEvery { mockClerkApiService.tokens("session_123") } coAnswers @@ -318,7 +370,8 @@ class SessionTokenFetcherTest { delay(100) // Simulate network delay ClerkResult.success(mockTokenResource) } - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When - Launch multiple concurrent requests val deferred1 = async { sessionTokenFetcher.getToken(mockSession) } @@ -338,6 +391,49 @@ class SessionTokenFetcherTest { coVerify(exactly = 1) { mockClerkApiService.tokens("session_123") } } + @Test + fun `forced refreshes are not deduplicated`() = runTest { + // Given + val cacheKey = "session_123-organization-" + val firstCallStarted = CompletableDeferred() + val releaseFirstCall = CompletableDeferred() + val callCount = AtomicInteger() + val secondToken = mockk(relaxed = true) + coEvery { SessionTokensCache.getToken(cacheKey) } returns null + coEvery { mockClerkApiService.tokens("session_123") } coAnswers + { + if (callCount.getAndIncrement() == 0) { + firstCallStarted.complete(Unit) + releaseFirstCall.await() + ClerkResult.success(mockTokenResource) + } else { + ClerkResult.success(secondToken) + } + } + coEvery { SessionTokensCache.storeIfFresher(cacheKey, any(), any()) } answers + { + val token = secondArg() + SessionTokensCache.StoreResult(token, true) + } + + // When + val first = async { + sessionTokenFetcher.getToken(mockSession, GetTokenOptions(skipCache = true)) + } + firstCallStarted.await() + val second = async { + sessionTokenFetcher.getToken(mockSession, GetTokenOptions(skipCache = true)) + } + val secondResult = second.await() + releaseFirstCall.complete(Unit) + val firstResult = first.await() + + // Then + assertEquals(mockTokenResource, firstResult) + assertEquals(secondToken, secondResult) + coVerify(exactly = 2) { mockClerkApiService.tokens("session_123") } + } + @Test fun `getToken handles API exception gracefully`() = runTest { // Given @@ -350,7 +446,7 @@ class SessionTokenFetcherTest { // Then assertNull(result) coVerify { mockClerkApiService.tokens("session_123") } - coVerify(exactly = 0) { SessionTokensCache.setToken(any(), any()) } + coVerify(exactly = 0) { SessionTokensCache.storeIfFresher(any(), any(), any()) } } @Test @@ -362,7 +458,7 @@ class SessionTokenFetcherTest { val cacheKey = mockSession.tokenCacheKey(null) // Then - assertEquals("session_456", cacheKey) + assertEquals("session_456-organization-", cacheKey) } @Test @@ -375,7 +471,7 @@ class SessionTokenFetcherTest { val cacheKey = mockSession.tokenCacheKey(template) // Then - assertEquals("session_456-admin_template", cacheKey) + assertEquals("session_456-template-admin_template", cacheKey) } @Test @@ -386,13 +482,17 @@ class SessionTokenFetcherTest { every { session1.id } returns "session_1" every { session2.id } returns "session_2" - coEvery { SessionTokensCache.getToken("session_1") } returns null - coEvery { SessionTokensCache.getToken("session_2") } returns null + coEvery { SessionTokensCache.getToken("session_1-organization-") } returns null + coEvery { SessionTokensCache.getToken("session_2-organization-") } returns null coEvery { mockClerkApiService.tokens("session_1") } returns ClerkResult.success(mockTokenResource) coEvery { mockClerkApiService.tokens("session_2") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(any(), any()) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(any(), any(), any()) } answers + { + val token = secondArg() + SessionTokensCache.StoreResult(token, true) + } // When sessionTokenFetcher.getToken(session1) @@ -401,8 +501,12 @@ class SessionTokenFetcherTest { // Then coVerify { mockClerkApiService.tokens("session_1") } coVerify { mockClerkApiService.tokens("session_2") } - coVerify { SessionTokensCache.setToken("session_1", mockTokenResource) } - coVerify { SessionTokensCache.setToken("session_2", mockTokenResource) } + coVerify { + SessionTokensCache.storeIfFresher("session_1-organization-", mockTokenResource, any()) + } + coVerify { + SessionTokensCache.storeIfFresher("session_2-organization-", mockTokenResource, any()) + } } @Test @@ -422,13 +526,14 @@ class SessionTokenFetcherTest { @Test fun `getToken proceeds normally for active session`() = runTest { // Given - a session with ACTIVE status - val cacheKey = "session_123" + val cacheKey = "session_123-organization-" every { mockSession.status } returns Session.SessionStatus.ACTIVE coEvery { SessionTokensCache.getToken(cacheKey) } returns null coEvery { mockClerkApiService.tokens("session_123") } returns ClerkResult.success(mockTokenResource) - coEvery { SessionTokensCache.setToken(cacheKey, mockTokenResource) } returns Unit + coEvery { SessionTokensCache.storeIfFresher(cacheKey, mockTokenResource, any()) } returns + SessionTokensCache.StoreResult(mockTokenResource, true) // When val result = sessionTokenFetcher.getToken(mockSession) diff --git a/source/api/src/test/java/com/clerk/api/session/TokenFreshnessTest.kt b/source/api/src/test/java/com/clerk/api/session/TokenFreshnessTest.kt new file mode 100644 index 000000000..d35488f24 --- /dev/null +++ b/source/api/src/test/java/com/clerk/api/session/TokenFreshnessTest.kt @@ -0,0 +1,150 @@ +package com.clerk.api.session + +import com.clerk.api.network.model.token.TokenResource +import java.nio.charset.StandardCharsets +import java.util.Base64 +import kotlinx.coroutines.CompletableDeferred +import kotlinx.coroutines.async +import kotlinx.coroutines.test.runTest +import org.junit.After +import org.junit.Assert.assertEquals +import org.junit.Assert.assertFalse +import org.junit.Assert.assertTrue +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner + +@RunWith(RobolectricTestRunner::class) +class TokenFreshnessTest { + @After + fun tearDown() { + SessionTokensCache.clear() + } + + @Test + fun `keeps token with higher origin issued at`() { + val existing = token(originIssuedAt = 200, issuedAt = 200) + val incoming = token(originIssuedAt = 100, issuedAt = 300) + + val result = TokenFreshness.pickFreshest(existing, incoming) + + assertEquals(existing, result) + } + + @Test + fun `uses issued at to break equal origin issued at`() { + val existing = token(originIssuedAt = 200, issuedAt = 300) + val incoming = token(originIssuedAt = 200, issuedAt = 400) + + val result = TokenFreshness.pickFreshest(existing, incoming) + + assertEquals(incoming, result) + } + + @Test + fun `replaces an expired existing token`() { + val existing = token(originIssuedAt = 300, issuedAt = 300, expiresAt = 900) + val incoming = token(originIssuedAt = 100, issuedAt = 100, expiresAt = 2_000) + + val result = TokenFreshness.pickFreshest(existing, incoming, nowMillis = 1_000_000) + + assertEquals(incoming, result) + } + + @Test + fun `accepts incoming token when organization changes`() { + val existing = token(organizationId = "org_one", originIssuedAt = 300, issuedAt = 300) + val incoming = token(organizationId = "org_two", originIssuedAt = 100, issuedAt = 100) + + val result = TokenFreshness.pickFreshest(existing, incoming) + + assertEquals(incoming, result) + } + + @Test + fun `keeps decodable existing token when incoming cannot be decoded`() { + val existing = token(originIssuedAt = 100, issuedAt = 100) + val incoming = TokenResource("malformed") + + val result = TokenFreshness.pickFreshest(existing, incoming) + + assertEquals(existing, result) + } + + @Test + fun `cache cannot be rolled back by a stale response`() { + val cacheKey = "session-organization-" + val stale = token(originIssuedAt = 100, issuedAt = 100, signature = "stale") + val fresh = token(originIssuedAt = 200, issuedAt = 200, signature = "fresh") + + val firstStore = SessionTokensCache.storeIfFresher(cacheKey, fresh) + val staleStore = SessionTokensCache.storeIfFresher(cacheKey, stale) + val duplicateStore = SessionTokensCache.storeIfFresher(cacheKey, fresh) + + assertTrue(firstStore.didChangeCanonicalToken) + assertFalse(staleStore.didChangeCanonicalToken) + assertFalse(duplicateStore.didChangeCanonicalToken) + assertEquals(fresh, SessionTokensCache.getToken(cacheKey)) + } + + @Test + fun `out of order cache writes retain the freshest response`() = runTest { + val cacheKey = "session-organization-" + val stale = token(originIssuedAt = 100, issuedAt = 100, signature = "stale") + val fresh = token(originIssuedAt = 200, issuedAt = 200, signature = "fresh") + val releaseStaleResponse = CompletableDeferred() + val staleResponse = async { + releaseStaleResponse.await() + SessionTokensCache.storeIfFresher(cacheKey, stale) + } + + SessionTokensCache.storeIfFresher(cacheKey, fresh) + releaseStaleResponse.complete(Unit) + staleResponse.await() + + assertEquals(fresh, SessionTokensCache.getToken(cacheKey)) + } + + @Test + fun `hydration preserves canonical token on a timestamp tie`() { + val cacheKey = "session-organization-" + val canonical = token(originIssuedAt = 100, issuedAt = 100, signature = "canonical") + val snapshot = token(originIssuedAt = 100, issuedAt = 100, signature = "snapshot") + SessionTokensCache.setToken(cacheKey, canonical) + + SessionTokensCache.hydrate(cacheKey, snapshot) + + assertEquals(canonical, SessionTokensCache.getToken(cacheKey)) + } + + private fun token( + sessionId: String = "session", + organizationId: String? = null, + originIssuedAt: Long?, + issuedAt: Long, + expiresAt: Long = 4_000_000_000, + signature: String = "signature", + ): TokenResource { + val headerClaims = + buildList { + add("\"alg\":\"none\"") + add("\"typ\":\"JWT\"") + originIssuedAt?.let { add("\"oiat\":$it") } + } + .joinToString(",") + val payloadClaims = + buildList { + add("\"sid\":\"$sessionId\"") + add("\"iat\":$issuedAt") + add("\"exp\":$expiresAt") + organizationId?.let { add("\"org_id\":\"$it\"") } + } + .joinToString(",") + return TokenResource("${encode("{$headerClaims}")}.${encode("{$payloadClaims}")}.$signature") + } + + private fun encode(value: String): String = + Base64.getUrlEncoder() + .withoutPadding() + .encodeToString(value.toByteArray(StandardCharsets.UTF_8)) +}