Make session refresh retries idempotent
CI / verify (push) Has been cancelled
CI / publish (push) Has been cancelled

Preserve the successor session for legitimate refresh retries so transient failures no longer revoke the user's session family.
This commit is contained in:
Rocky
2026-08-25 13:13:41 +08:00
parent 36a926f12f
commit 4c9e5feec0
14 changed files with 409 additions and 65 deletions
+1
View File
@@ -23,6 +23,7 @@ JWT_AUDIENCE=osgkeyboard-ios
JWT_SECRET=replace-with-at-least-32-random-bytes JWT_SECRET=replace-with-at-least-32-random-bytes
ACCESS_TOKEN_MINUTES=15 ACCESS_TOKEN_MINUTES=15
REFRESH_TOKEN_DAYS=30 REFRESH_TOKEN_DAYS=30
LEGACY_REFRESH_REPLAY_SECONDS=30
GATEWAY_GRANT_DAYS=30 GATEWAY_GRANT_DAYS=30
FIELD_ENCRYPTION_KEY=replace-with-exactly-32-random-bytes-as-base64 FIELD_ENCRYPTION_KEY=replace-with-exactly-32-random-bytes-as-base64
IDENTITY_HMAC_KEY=replace-with-a-distinct-32-random-bytes-as-base64 IDENTITY_HMAC_KEY=replace-with-a-distinct-32-random-bytes-as-base64
+1
View File
@@ -26,6 +26,7 @@ services:
JWT_SECRET: ${JWT_SECRET:?set a random JWT secret} JWT_SECRET: ${JWT_SECRET:?set a random JWT secret}
ACCESS_TOKEN_MINUTES: ${ACCESS_TOKEN_MINUTES:-15} ACCESS_TOKEN_MINUTES: ${ACCESS_TOKEN_MINUTES:-15}
REFRESH_TOKEN_DAYS: ${REFRESH_TOKEN_DAYS:-30} REFRESH_TOKEN_DAYS: ${REFRESH_TOKEN_DAYS:-30}
LEGACY_REFRESH_REPLAY_SECONDS: ${LEGACY_REFRESH_REPLAY_SECONDS:-30}
GATEWAY_GRANT_DAYS: ${GATEWAY_GRANT_DAYS:-30} GATEWAY_GRANT_DAYS: ${GATEWAY_GRANT_DAYS:-30}
FIELD_ENCRYPTION_KEY: ${FIELD_ENCRYPTION_KEY:?set a 32-byte Base64 key} FIELD_ENCRYPTION_KEY: ${FIELD_ENCRYPTION_KEY:?set a 32-byte Base64 key}
IDENTITY_HMAC_KEY: ${IDENTITY_HMAC_KEY:?set a distinct Base64 key} IDENTITY_HMAC_KEY: ${IDENTITY_HMAC_KEY:?set a distinct Base64 key}
+7
View File
@@ -60,6 +60,13 @@ paths:
required: [refreshToken] required: [refreshToken]
properties: properties:
refreshToken: { type: string, minLength: 32 } refreshToken: { type: string, minLength: 32 }
refreshOperationId:
type: string
format: uuid
description: |
Stable ID for one logical refresh attempt. Retrying with the same
ID returns the same successor while that session remains current
and unexpired.
responses: responses:
"200": "200":
description: Rotated session description: Rotated session
@@ -56,6 +56,10 @@ data class AppConfig(
hmacSecret = config.secret("app.session.secret", production).toByteArray(), hmacSecret = config.secret("app.session.secret", production).toByteArray(),
accessMinutes = config.positiveLong("app.session.accessMinutes"), accessMinutes = config.positiveLong("app.session.accessMinutes"),
refreshDays = config.positiveLong("app.session.refreshDays"), refreshDays = config.positiveLong("app.session.refreshDays"),
legacyRefreshReplaySeconds = config.positiveLong(
"app.session.legacyRefreshReplaySeconds",
30,
),
gatewayGrantDays = config.positiveLong("app.session.gatewayGrantDays", 30), gatewayGrantDays = config.positiveLong("app.session.gatewayGrantDays", 30),
) )
val encryption = EncryptionConfig( val encryption = EncryptionConfig(
@@ -244,6 +248,9 @@ data class AppConfig(
require(session.refreshDays in 1..365) { require(session.refreshDays in 1..365) {
"app.session.refreshDays must be between 1 and 365" "app.session.refreshDays must be between 1 and 365"
} }
require(session.legacyRefreshReplaySeconds in 5..120) {
"app.session.legacyRefreshReplaySeconds must be between 5 and 120"
}
require(!production || providers.volcengine.credentialsAvailable) { require(!production || providers.volcengine.credentialsAvailable) {
"Production Volcengine credentials are missing" "Production Volcengine credentials are missing"
} }
@@ -380,6 +387,7 @@ data class SessionConfig(
val hmacSecret: ByteArray, val hmacSecret: ByteArray,
val accessMinutes: Long, val accessMinutes: Long,
val refreshDays: Long, val refreshDays: Long,
val legacyRefreshReplaySeconds: Long = 30,
val gatewayGrantDays: Long = 30, val gatewayGrantDays: Long = 30,
) )
@@ -6,7 +6,9 @@ import org.jetbrains.exposed.v1.core.Table
import org.jetbrains.exposed.v1.core.and import org.jetbrains.exposed.v1.core.and
import org.jetbrains.exposed.v1.core.eq import org.jetbrains.exposed.v1.core.eq
import org.jetbrains.exposed.v1.core.greater import org.jetbrains.exposed.v1.core.greater
import org.jetbrains.exposed.v1.core.isNotNull
import org.jetbrains.exposed.v1.core.isNull import org.jetbrains.exposed.v1.core.isNull
import org.jetbrains.exposed.v1.core.lessEq
import org.jetbrains.exposed.v1.javatime.timestamp import org.jetbrains.exposed.v1.javatime.timestamp
import org.jetbrains.exposed.v1.jdbc.insert import org.jetbrains.exposed.v1.jdbc.insert
import org.jetbrains.exposed.v1.jdbc.insertIgnore import org.jetbrains.exposed.v1.jdbc.insertIgnore
@@ -47,6 +49,10 @@ internal object SessionsTable : Table("sessions") {
val familyId = varchar("family_id", 36).index() val familyId = varchar("family_id", 36).index()
val refreshTokenHash = varchar("refresh_token_hash", 64).uniqueIndex() val refreshTokenHash = varchar("refresh_token_hash", 64).uniqueIndex()
val replacedById = varchar("replaced_by_id", 36).nullable() val replacedById = varchar("replaced_by_id", 36).nullable()
val refreshOperationId = varchar("refresh_operation_id", 36).nullable()
val encryptedReplacementRefreshToken =
varchar("encrypted_replacement_refresh_token", 255).nullable()
val refreshReplayUntil = timestamp("refresh_replay_until").nullable().index()
val createdAt = timestamp("created_at") val createdAt = timestamp("created_at")
val expiresAt = timestamp("expires_at") val expiresAt = timestamp("expires_at")
val revokedAt = timestamp("revoked_at").nullable() val revokedAt = timestamp("revoked_at").nullable()
@@ -66,6 +72,16 @@ data class CreatedSession(
val familyId: UUID, val familyId: UUID,
) )
data class RefreshRotationAttempt(
val currentTokenHash: String,
val newTokenHash: String,
val encryptedNewToken: String,
val newExpiresAt: Instant,
val operationId: UUID?,
val replayUntil: Instant,
val now: Instant,
)
sealed interface RefreshRotationResult { sealed interface RefreshRotationResult {
data class Rotated( data class Rotated(
val accountId: UUID, val accountId: UUID,
@@ -73,12 +89,21 @@ sealed interface RefreshRotationResult {
val familyId: UUID, val familyId: UUID,
) : RefreshRotationResult ) : RefreshRotationResult
data class Replayed(
val accountId: UUID,
val sessionId: UUID,
val familyId: UUID,
val encryptedRefreshToken: String,
val refreshTokenExpiresAt: Instant,
) : RefreshRotationResult
data object Invalid : RefreshRotationResult data object Invalid : RefreshRotationResult
data object ReuseDetected : RefreshRotationResult data object ReuseDetected : RefreshRotationResult
} }
internal enum class RefreshRotationDecision { internal enum class RefreshRotationDecision {
ROTATE, ROTATE,
REPLAY_ROTATION,
REVOKE_EXPIRED, REVOKE_EXPIRED,
REVOKE_REUSED_FAMILY, REVOKE_REUSED_FAMILY,
} }
@@ -91,9 +116,11 @@ internal object RefreshRotationPolicy {
fun decide( fun decide(
revoked: Boolean, revoked: Boolean,
replaced: Boolean, replaced: Boolean,
replayable: Boolean,
expiresAt: Instant, expiresAt: Instant,
now: Instant, now: Instant,
): RefreshRotationDecision = when { ): RefreshRotationDecision = when {
replayable -> RefreshRotationDecision.REPLAY_ROTATION
revoked || replaced -> RefreshRotationDecision.REVOKE_REUSED_FAMILY revoked || replaced -> RefreshRotationDecision.REVOKE_REUSED_FAMILY
!expiresAt.isAfter(now) -> RefreshRotationDecision.REVOKE_EXPIRED !expiresAt.isAfter(now) -> RefreshRotationDecision.REVOKE_EXPIRED
else -> RefreshRotationDecision.ROTATE else -> RefreshRotationDecision.ROTATE
@@ -114,12 +141,7 @@ interface AuthRepository {
now: Instant, now: Instant,
): CreatedSession ): CreatedSession
suspend fun rotateRefreshToken( suspend fun rotateRefreshToken(attempt: RefreshRotationAttempt): RefreshRotationResult
currentTokenHash: String,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult
suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean
suspend fun isSessionActive(accountId: UUID, sessionId: UUID, now: Instant): Boolean suspend fun isSessionActive(accountId: UUID, sessionId: UUID, now: Instant): Boolean
@@ -215,37 +237,78 @@ class ExposedAuthRepository(
} }
override suspend fun rotateRefreshToken( override suspend fun rotateRefreshToken(
currentTokenHash: String, attempt: RefreshRotationAttempt,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult = databaseFactory.query { ): RefreshRotationResult = databaseFactory.query {
SessionsTable.update({
SessionsTable.refreshReplayUntil.isNotNull() and
(SessionsTable.refreshReplayUntil lessEq attempt.now)
}) {
it[refreshOperationId] = null
it[encryptedReplacementRefreshToken] = null
it[refreshReplayUntil] = null
}
val current = SessionsTable.selectAll() val current = SessionsTable.selectAll()
.where { SessionsTable.refreshTokenHash eq currentTokenHash } .where { SessionsTable.refreshTokenHash eq attempt.currentTokenHash }
.forUpdate() .forUpdate()
.singleOrNull() .singleOrNull()
?: return@query RefreshRotationResult.Invalid ?: return@query RefreshRotationResult.Invalid
val familyId = current[SessionsTable.familyId] val familyId = current[SessionsTable.familyId]
val replacement = current[SessionsTable.replacedById]?.takeIf {
current[SessionsTable.refreshOperationId] == attempt.operationId?.toString() &&
current[SessionsTable.refreshReplayUntil]?.isAfter(attempt.now) == true &&
current[SessionsTable.encryptedReplacementRefreshToken] != null
}?.let { replacementId ->
SessionsTable.selectAll()
.where { SessionsTable.id eq replacementId }
.forUpdate()
.singleOrNull()
?.takeIf {
it[SessionsTable.accountId] == current[SessionsTable.accountId] &&
it[SessionsTable.familyId] == familyId &&
it[SessionsTable.revokedAt] == null &&
it[SessionsTable.replacedById] == null &&
it[SessionsTable.expiresAt].isAfter(attempt.now)
}
}
when ( when (
RefreshRotationPolicy.decide( RefreshRotationPolicy.decide(
revoked = current[SessionsTable.revokedAt] != null, revoked = current[SessionsTable.revokedAt] != null,
replaced = current[SessionsTable.replacedById] != null, replaced = current[SessionsTable.replacedById] != null,
replayable = replacement != null,
expiresAt = current[SessionsTable.expiresAt], expiresAt = current[SessionsTable.expiresAt],
now = now, now = attempt.now,
) )
) { ) {
RefreshRotationDecision.REPLAY_ROTATION -> {
val replayed = requireNotNull(replacement)
return@query RefreshRotationResult.Replayed(
accountId = UUID.fromString(replayed[SessionsTable.accountId]),
sessionId = UUID.fromString(replayed[SessionsTable.id]),
familyId = UUID.fromString(replayed[SessionsTable.familyId]),
encryptedRefreshToken = requireNotNull(
current[SessionsTable.encryptedReplacementRefreshToken],
),
refreshTokenExpiresAt = replayed[SessionsTable.expiresAt],
)
}
RefreshRotationDecision.REVOKE_REUSED_FAMILY -> { RefreshRotationDecision.REVOKE_REUSED_FAMILY -> {
SessionsTable.update({ SessionsTable.familyId eq familyId }) { SessionsTable.update({ SessionsTable.familyId eq familyId }) {
it[SessionsTable.revokedAt] = now it[SessionsTable.revokedAt] = attempt.now
it[refreshOperationId] = null
it[encryptedReplacementRefreshToken] = null
it[refreshReplayUntil] = null
} }
SessionsTable.update({ SessionsTable.id eq current[SessionsTable.id] }) { SessionsTable.update({ SessionsTable.id eq current[SessionsTable.id] }) {
it[SessionsTable.reuseDetectedAt] = now it[SessionsTable.reuseDetectedAt] = attempt.now
} }
return@query RefreshRotationResult.ReuseDetected return@query RefreshRotationResult.ReuseDetected
} }
RefreshRotationDecision.REVOKE_EXPIRED -> { RefreshRotationDecision.REVOKE_EXPIRED -> {
SessionsTable.update({ SessionsTable.familyId eq familyId }) { SessionsTable.update({ SessionsTable.familyId eq familyId }) {
it[SessionsTable.revokedAt] = now it[SessionsTable.revokedAt] = attempt.now
it[refreshOperationId] = null
it[encryptedReplacementRefreshToken] = null
it[refreshReplayUntil] = null
} }
return@query RefreshRotationResult.Invalid return@query RefreshRotationResult.Invalid
} }
@@ -257,13 +320,16 @@ class ExposedAuthRepository(
it[SessionsTable.id] = newSessionId.toString() it[SessionsTable.id] = newSessionId.toString()
it[SessionsTable.accountId] = current[SessionsTable.accountId] it[SessionsTable.accountId] = current[SessionsTable.accountId]
it[SessionsTable.familyId] = familyId it[SessionsTable.familyId] = familyId
it[SessionsTable.refreshTokenHash] = newTokenHash it[SessionsTable.refreshTokenHash] = attempt.newTokenHash
it[SessionsTable.createdAt] = now it[SessionsTable.createdAt] = attempt.now
it[SessionsTable.expiresAt] = newExpiresAt it[SessionsTable.expiresAt] = attempt.newExpiresAt
} }
SessionsTable.update({ SessionsTable.id eq current[SessionsTable.id] }) { SessionsTable.update({ SessionsTable.id eq current[SessionsTable.id] }) {
it[SessionsTable.replacedById] = newSessionId.toString() it[SessionsTable.replacedById] = newSessionId.toString()
it[SessionsTable.revokedAt] = now it[SessionsTable.revokedAt] = attempt.now
it[SessionsTable.refreshOperationId] = attempt.operationId?.toString()
it[SessionsTable.encryptedReplacementRefreshToken] = attempt.encryptedNewToken
it[SessionsTable.refreshReplayUntil] = attempt.replayUntil
} }
RefreshRotationResult.Rotated( RefreshRotationResult.Rotated(
accountId = UUID.fromString(current[SessionsTable.accountId]), accountId = UUID.fromString(current[SessionsTable.accountId]),
@@ -290,6 +356,9 @@ class ExposedAuthRepository(
(SessionsTable.familyId eq session[SessionsTable.familyId]) (SessionsTable.familyId eq session[SessionsTable.familyId])
}) { }) {
it[SessionsTable.revokedAt] = now it[SessionsTable.revokedAt] = now
it[refreshOperationId] = null
it[encryptedReplacementRefreshToken] = null
it[refreshReplayUntil] = null
} > 0 } > 0
} }
@@ -1,6 +1,7 @@
package com.osglab.account.features.auth package com.osglab.account.features.auth
import com.osglab.account.common.api.ApiResponse import com.osglab.account.common.api.ApiResponse
import com.osglab.account.common.errors.InvalidRequestException
import com.osglab.account.common.errors.UnauthorizedException import com.osglab.account.common.errors.UnauthorizedException
import com.osglab.account.common.security.AccountPrincipal import com.osglab.account.common.security.AccountPrincipal
import com.osglab.account.common.security.SESSION_AUTH_NAME import com.osglab.account.common.security.SESSION_AUTH_NAME
@@ -15,6 +16,7 @@ import io.ktor.server.routing.Route
import io.ktor.server.routing.post import io.ktor.server.routing.post
import io.ktor.server.routing.route import io.ktor.server.routing.route
import kotlinx.serialization.Serializable import kotlinx.serialization.Serializable
import java.util.UUID
@Serializable @Serializable
data class AppleSignInRequest( data class AppleSignInRequest(
@@ -44,8 +46,12 @@ data class AppAttestRequest(
} }
@Serializable @Serializable
data class RefreshSessionRequest(val refreshToken: String) { data class RefreshSessionRequest(
override fun toString(): String = "RefreshSessionRequest(refreshToken=[REDACTED])" val refreshToken: String,
val refreshOperationId: String? = null,
) {
override fun toString(): String =
"RefreshSessionRequest(refreshToken=[REDACTED], refreshOperationId=$refreshOperationId)"
} }
@Serializable @Serializable
@@ -92,8 +98,11 @@ class AuthRoutes(
} }
post("/refresh") { post("/refresh") {
val request = call.receive<RefreshSessionRequest>() val request = call.receive<RefreshSessionRequest>()
val operationId = request.refreshOperationId?.let(::parseRefreshOperationId)
call.respond( call.respond(
ApiResponse(data = sessionService.refresh(request.refreshToken).toResponse()), ApiResponse(
data = sessionService.refresh(request.refreshToken, operationId).toResponse(),
),
) )
} }
authenticate(SESSION_AUTH_NAME) { authenticate(SESSION_AUTH_NAME) {
@@ -112,6 +121,10 @@ class AuthRoutes(
fun Route.authRoutes(sessionService: SessionService) = fun Route.authRoutes(sessionService: SessionService) =
AuthRoutes(sessionService).register(this) AuthRoutes(sessionService).register(this)
private fun parseRefreshOperationId(value: String): UUID =
runCatching { UUID.fromString(value) }
.getOrElse { throw InvalidRequestException("refreshOperationId must be a UUID") }
private fun SessionTokens.toResponse(): SessionTokenResponse = SessionTokenResponse( private fun SessionTokens.toResponse(): SessionTokenResponse = SessionTokenResponse(
accountId = accountId.toString(), accountId = accountId.toString(),
accessToken = accessToken, accessToken = accessToken,
@@ -96,30 +96,53 @@ class SessionService(
return createSession(account.id, now) return createSession(account.id, now)
} }
suspend fun refresh(refreshToken: String): SessionTokens { suspend fun refresh(refreshToken: String, operationId: UUID? = null): SessionTokens {
requireValue(refreshToken, "refreshToken", MAX_REFRESH_TOKEN_LENGTH) requireValue(refreshToken, "refreshToken", MAX_REFRESH_TOKEN_LENGTH)
val now = clock.instant() val now = clock.instant()
val currentTokenHash = TokenHash.sha256(refreshToken)
val replacement = tokenGenerator.newRefreshToken() val replacement = tokenGenerator.newRefreshToken()
val replacementExpiresAt = now.plus(Duration.ofDays(sessionConfig.refreshDays)) val replacementExpiresAt = now.plus(Duration.ofDays(sessionConfig.refreshDays))
val replayUntil = if (operationId == null) {
now.plusSeconds(sessionConfig.legacyRefreshReplaySeconds)
} else {
// A stable operation ID lets a crashed client recover until the successor expires.
replacementExpiresAt
}
return when ( return when (
val result = repository.rotateRefreshToken( val result = repository.rotateRefreshToken(
currentTokenHash = TokenHash.sha256(refreshToken), RefreshRotationAttempt(
currentTokenHash = currentTokenHash,
newTokenHash = TokenHash.sha256(replacement), newTokenHash = TokenHash.sha256(replacement),
encryptedNewToken = fieldEncryptor.encrypt(
replacement,
refreshReplayContext(currentTokenHash),
),
newExpiresAt = replacementExpiresAt, newExpiresAt = replacementExpiresAt,
operationId = operationId,
replayUntil = replayUntil,
now = now, now = now,
),
) )
) { ) {
RefreshRotationResult.Invalid -> throw UnauthorizedException("Refresh token is invalid or expired") RefreshRotationResult.Invalid -> throw UnauthorizedException("Refresh token is invalid or expired")
RefreshRotationResult.ReuseDetected -> throw TokenReuseException() RefreshRotationResult.ReuseDetected -> throw TokenReuseException()
is RefreshRotationResult.Rotated -> { is RefreshRotationResult.Rotated -> issueSessionTokens(
val access = sessionJwt.issue(result.accountId, result.sessionId)
SessionTokens(
accountId = result.accountId, accountId = result.accountId,
accessToken = access.value, sessionId = result.sessionId,
accessTokenExpiresAt = access.expiresAt,
refreshToken = replacement, refreshToken = replacement,
refreshTokenExpiresAt = replacementExpiresAt, refreshTokenExpiresAt = replacementExpiresAt,
) )
is RefreshRotationResult.Replayed -> {
val replayedRefreshToken = fieldEncryptor.decrypt(
result.encryptedRefreshToken,
refreshReplayContext(currentTokenHash),
)
issueSessionTokens(
accountId = result.accountId,
sessionId = result.sessionId,
refreshToken = replayedRefreshToken,
refreshTokenExpiresAt = result.refreshTokenExpiresAt,
)
} }
} }
} }
@@ -147,6 +170,22 @@ class SessionService(
) )
} }
private fun issueSessionTokens(
accountId: UUID,
sessionId: UUID,
refreshToken: String,
refreshTokenExpiresAt: Instant,
): SessionTokens {
val access = sessionJwt.issue(accountId, sessionId)
return SessionTokens(
accountId = accountId,
accessToken = access.value,
accessTokenExpiresAt = access.expiresAt,
refreshToken = refreshToken,
refreshTokenExpiresAt = refreshTokenExpiresAt,
)
}
private suspend fun verifyIdentityToken(token: String, nonce: String): AppleIdentity = private suspend fun verifyIdentityToken(token: String, nonce: String): AppleIdentity =
try { try {
appleIdentityVerifier.verify(token, nonce) appleIdentityVerifier.verify(token, nonce)
@@ -185,3 +224,4 @@ class SessionService(
fun appleRefreshContext(accountId: UUID): String = "apple-refresh-token:$accountId" fun appleRefreshContext(accountId: UUID): String = "apple-refresh-token:$accountId"
fun appleSubjectContext(identityFingerprint: String): String = "apple-subject:$identityFingerprint" fun appleSubjectContext(identityFingerprint: String): String = "apple-subject:$identityFingerprint"
fun refreshReplayContext(currentTokenHash: String): String = "session-refresh-replay:$currentTokenHash"
+1
View File
@@ -24,6 +24,7 @@ app:
secret: "$JWT_SECRET" secret: "$JWT_SECRET"
accessMinutes: "$ACCESS_TOKEN_MINUTES:15" accessMinutes: "$ACCESS_TOKEN_MINUTES:15"
refreshDays: "$REFRESH_TOKEN_DAYS:30" refreshDays: "$REFRESH_TOKEN_DAYS:30"
legacyRefreshReplaySeconds: "$LEGACY_REFRESH_REPLAY_SECONDS:30"
gatewayGrantDays: "$GATEWAY_GRANT_DAYS:30" gatewayGrantDays: "$GATEWAY_GRANT_DAYS:30"
encryption: encryption:
keyBase64: "$FIELD_ENCRYPTION_KEY" keyBase64: "$FIELD_ENCRYPTION_KEY"
@@ -0,0 +1,21 @@
ALTER TABLE sessions
ADD COLUMN refresh_operation_id VARCHAR(36) CHARACTER SET ascii COLLATE ascii_bin NULL
AFTER replaced_by_id,
ADD COLUMN encrypted_replacement_refresh_token VARCHAR(255)
CHARACTER SET ascii COLLATE ascii_bin NULL
AFTER refresh_operation_id,
ADD COLUMN refresh_replay_until DATETIME(6) NULL
AFTER encrypted_replacement_refresh_token,
ADD INDEX ix_sessions_refresh_replay_expiry (refresh_replay_until),
ADD CONSTRAINT chk_sessions_refresh_replay_payload CHECK (
(
encrypted_replacement_refresh_token IS NULL
AND refresh_replay_until IS NULL
AND refresh_operation_id IS NULL
)
OR
(
encrypted_replacement_refresh_token IS NOT NULL
AND refresh_replay_until IS NOT NULL
)
);
@@ -17,9 +17,20 @@ class AppConfigTest : FunSpec({
config.credits.signupTrial shouldBe 1_000 config.credits.signupTrial shouldBe 1_000
config.credits.referralInviter shouldBe 1_000 config.credits.referralInviter shouldBe 1_000
config.credits.referralInvitee shouldBe 1_000 config.credits.referralInvitee shouldBe 1_000
config.session.legacyRefreshReplaySeconds shouldBe 30
config.admin.mtlsRequired shouldBe true config.admin.mtlsRequired shouldBe true
} }
test("refresh replay window is bounded") {
val config = validConfig("test").apply {
put("app.session.legacyRefreshReplaySeconds", "121")
}
shouldThrow<IllegalArgumentException> {
AppConfig.from(config)
}.message.orEmpty() shouldContain "legacyRefreshReplaySeconds"
}
test("production rejects placeholder secrets") { test("production rejects placeholder secrets") {
val config = validProductionConfig().apply { val config = validProductionConfig().apply {
put("app.session.secret", "replace-with-secret") put("app.session.secret", "replace-with-secret")
@@ -20,6 +20,24 @@ class DeploymentConsistencyTest : FunSpec({
documentedPaths shouldBe EXPECTED_PUBLIC_PATHS documentedPaths shouldBe EXPECTED_PUBLIC_PATHS
} }
test("session refresh idempotency stays aligned across API, schema, and deployment") {
val openApi = root.read("docs/openapi.yaml")
val migration = root.read(
"src/main/resources/db/migration/V29__idempotent_session_refresh.sql",
)
openApi shouldContain "refreshOperationId"
migration shouldContain "encrypted_replacement_refresh_token"
migration shouldContain "refresh_replay_until"
listOf(
root.read("src/main/resources/application.yaml"),
root.read(".env.example"),
root.read("compose.yaml"),
).forEach { configuration ->
configuration shouldContain "LEGACY_REFRESH_REPLAY_SECONDS"
}
}
test("OpenAPI defines admin pagination and response contracts") { test("OpenAPI defines admin pagination and response contracts") {
val openApi = root.read("docs/openapi.yaml") val openApi = root.read("docs/openapi.yaml")
val sessionSchema = openApi val sessionSchema = openApi
@@ -65,10 +65,7 @@ private class MutableSessionStateRepository : AuthRepository {
): CreatedSession = error("Not used") ): CreatedSession = error("Not used")
override suspend fun rotateRefreshToken( override suspend fun rotateRefreshToken(
currentTokenHash: String, attempt: RefreshRotationAttempt,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult = error("Not used") ): RefreshRotationResult = error("Not used")
override suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean = override suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean =
@@ -137,8 +137,8 @@ class SessionServiceTest : FunSpec({
} }
} }
test("concurrent refresh accepts once and revokes the family on replay") { test("concurrent retries return the same successor without revoking the family") {
val repository = ConcurrentRotationRepository() val repository = IdempotentRotationRepository()
val sessionConfig = sessionConfig() val sessionConfig = sessionConfig()
val service = SessionService( val service = SessionService(
repository = repository, repository = repository,
@@ -162,8 +162,35 @@ class SessionServiceTest : FunSpec({
}.awaitAll() }.awaitAll()
} }
results.count { it.isSuccess } shouldBe 1 results.count { it.isSuccess } shouldBe 2
results.count { it.exceptionOrNull() is TokenReuseException } shouldBe 1 results.map { it.getOrThrow().refreshToken }.distinct().size shouldBe 1
repository.familyRevoked shouldBe false
}
test("retrying one refresh operation returns the original successor token") {
val repository = IdempotentRotationRepository()
val service = sessionService(repository)
val operationId = UUID.randomUUID()
val first = service.refresh("response-lost-token", operationId)
val replay = service.refresh("response-lost-token", operationId)
replay.accountId shouldBe first.accountId
replay.refreshToken shouldBe first.refreshToken
replay.refreshTokenExpiresAt shouldBe first.refreshTokenExpiresAt
repository.replayUntil shouldBe first.refreshTokenExpiresAt
repository.familyRevoked shouldBe false
}
test("replaying a consumed token for a different operation revokes the family") {
val repository = IdempotentRotationRepository()
val service = sessionService(repository)
service.refresh("stolen-refresh-token", UUID.randomUUID())
shouldThrow<TokenReuseException> {
service.refresh("stolen-refresh-token", UUID.randomUUID())
}
repository.familyRevoked shouldBe true repository.familyRevoked shouldBe true
} }
@@ -173,29 +200,40 @@ class SessionServiceTest : FunSpec({
RefreshRotationPolicy.decide( RefreshRotationPolicy.decide(
revoked = false, revoked = false,
replaced = false, replaced = false,
replayable = false,
expiresAt = now.plusSeconds(1), expiresAt = now.plusSeconds(1),
now = now, now = now,
) shouldBe RefreshRotationDecision.ROTATE ) shouldBe RefreshRotationDecision.ROTATE
RefreshRotationPolicy.decide( RefreshRotationPolicy.decide(
revoked = false, revoked = false,
replaced = false, replaced = false,
replayable = false,
expiresAt = now, expiresAt = now,
now = now, now = now,
) shouldBe RefreshRotationDecision.REVOKE_EXPIRED ) shouldBe RefreshRotationDecision.REVOKE_EXPIRED
} }
test("refresh rotation policy treats any consumed token as family reuse") { test("refresh rotation policy replays only an eligible consumed token") {
val now = Instant.parse("2026-08-16T00:00:00Z") val now = Instant.parse("2026-08-16T00:00:00Z")
RefreshRotationPolicy.decide(
revoked = true,
replaced = true,
replayable = true,
expiresAt = now.plusSeconds(60),
now = now,
) shouldBe RefreshRotationDecision.REPLAY_ROTATION
RefreshRotationPolicy.decide( RefreshRotationPolicy.decide(
revoked = true, revoked = true,
replaced = false, replaced = false,
replayable = false,
expiresAt = now.plusSeconds(60), expiresAt = now.plusSeconds(60),
now = now, now = now,
) shouldBe RefreshRotationDecision.REVOKE_REUSED_FAMILY ) shouldBe RefreshRotationDecision.REVOKE_REUSED_FAMILY
RefreshRotationPolicy.decide( RefreshRotationPolicy.decide(
revoked = false, revoked = false,
replaced = true, replaced = true,
replayable = false,
expiresAt = now.plusSeconds(60), expiresAt = now.plusSeconds(60),
now = now, now = now,
) shouldBe RefreshRotationDecision.REVOKE_REUSED_FAMILY ) shouldBe RefreshRotationDecision.REVOKE_REUSED_FAMILY
@@ -240,10 +278,7 @@ private class SuccessfulAuthRepository(
} }
override suspend fun rotateRefreshToken( override suspend fun rotateRefreshToken(
currentTokenHash: String, attempt: RefreshRotationAttempt,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult = error("Not used") ): RefreshRotationResult = error("Not used")
override suspend fun revokeSessionFamily( override suspend fun revokeSessionFamily(
@@ -283,10 +318,7 @@ private data object ReuseDetectingRepository : AuthRepository {
): CreatedSession = error("Not used") ): CreatedSession = error("Not used")
override suspend fun rotateRefreshToken( override suspend fun rotateRefreshToken(
currentTokenHash: String, attempt: RefreshRotationAttempt,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult = RefreshRotationResult.ReuseDetected ): RefreshRotationResult = RefreshRotationResult.ReuseDetected
override suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean = override suspend fun revokeSessionFamily(accountId: UUID, sessionId: UUID, now: Instant): Boolean =
@@ -301,28 +333,46 @@ private data object ReuseDetectingRepository : AuthRepository {
override suspend fun restrictAccountForAntiAbuse(accountId: UUID, now: Instant) = Unit override suspend fun restrictAccountForAntiAbuse(accountId: UUID, now: Instant) = Unit
} }
private class ConcurrentRotationRepository : AuthRepository { private class IdempotentRotationRepository : AuthRepository {
private val mutex = Mutex() private val mutex = Mutex()
private var consumed = false private val accountId = UUID.randomUUID()
private val sessionId = UUID.randomUUID()
private val familyId = UUID.randomUUID()
private var rotation: StoredRotation? = null
var familyRevoked = false var familyRevoked = false
private set private set
val replayUntil: Instant?
get() = rotation?.replayUntil
override suspend fun rotateRefreshToken( override suspend fun rotateRefreshToken(
currentTokenHash: String, attempt: RefreshRotationAttempt,
newTokenHash: String,
newExpiresAt: Instant,
now: Instant,
): RefreshRotationResult = mutex.withLock { ): RefreshRotationResult = mutex.withLock {
if (consumed) { val stored = rotation
if (stored == null) {
rotation = StoredRotation(
currentTokenHash = attempt.currentTokenHash,
encryptedRefreshToken = attempt.encryptedNewToken,
refreshTokenExpiresAt = attempt.newExpiresAt,
operationId = attempt.operationId,
replayUntil = attempt.replayUntil,
)
RefreshRotationResult.Rotated(accountId, sessionId, familyId)
} else if (
!familyRevoked &&
stored.currentTokenHash == attempt.currentTokenHash &&
stored.operationId == attempt.operationId &&
attempt.now.isBefore(stored.replayUntil)
) {
RefreshRotationResult.Replayed(
accountId = accountId,
sessionId = sessionId,
familyId = familyId,
encryptedRefreshToken = stored.encryptedRefreshToken,
refreshTokenExpiresAt = stored.refreshTokenExpiresAt,
)
} else {
familyRevoked = true familyRevoked = true
RefreshRotationResult.ReuseDetected RefreshRotationResult.ReuseDetected
} else {
consumed = true
RefreshRotationResult.Rotated(
accountId = UUID.randomUUID(),
sessionId = UUID.randomUUID(),
familyId = UUID.randomUUID(),
)
} }
} }
@@ -360,6 +410,33 @@ private class ConcurrentRotationRepository : AuthRepository {
override suspend fun restrictAccountForAntiAbuse(accountId: UUID, now: Instant) = Unit override suspend fun restrictAccountForAntiAbuse(accountId: UUID, now: Instant) = Unit
} }
private data class StoredRotation(
val currentTokenHash: String,
val encryptedRefreshToken: String,
val refreshTokenExpiresAt: Instant,
val operationId: UUID?,
val replayUntil: Instant,
)
private fun sessionService(repository: AuthRepository): SessionService {
val config = sessionConfig()
return SessionService(
repository = repository,
appleIdentityVerifier = AppleIdentityTokenVerifier(
appleConfig(),
object : AppleJwksProvider {
override suspend fun rsaKey(keyId: String): RSAKey? = null
},
),
appleTokenClient = UnavailableAppleTokenClient(),
integrityService = monitorOnlyIntegrityService(),
sessionJwt = SessionJwt(config),
fieldEncryptor = FieldEncryptor(ByteArray(32) { 4 }),
identityFingerprint = IdentityFingerprint(ByteArray(32) { 6 }),
sessionConfig = config,
)
}
private fun appleConfig() = AppleConfig( private fun appleConfig() = AppleConfig(
teamId = null, teamId = null,
keyId = null, keyId = null,
@@ -1,13 +1,18 @@
package com.osglab.account.integration package com.osglab.account.integration
import com.osglab.account.common.security.FieldEncryptor
import com.osglab.account.common.security.IdentityFingerprint import com.osglab.account.common.security.IdentityFingerprint
import com.osglab.account.common.security.SessionJwt import com.osglab.account.common.security.SessionJwt
import com.osglab.account.common.security.TokenHash
import com.osglab.account.config.DatabaseConfig import com.osglab.account.config.DatabaseConfig
import com.osglab.account.config.DatabaseFactory import com.osglab.account.config.DatabaseFactory
import com.osglab.account.config.SessionConfig import com.osglab.account.config.SessionConfig
import com.osglab.account.features.account.ExposedAccountRepository import com.osglab.account.features.account.ExposedAccountRepository
import com.osglab.account.features.auth.ExposedAuthRepository import com.osglab.account.features.auth.ExposedAuthRepository
import com.osglab.account.features.auth.RefreshRotationAttempt
import com.osglab.account.features.auth.RefreshRotationResult
import com.osglab.account.features.auth.SessionAccessAuthenticator import com.osglab.account.features.auth.SessionAccessAuthenticator
import com.osglab.account.features.auth.refreshReplayContext
import com.osglab.account.features.credits.domain.CreditConflict import com.osglab.account.features.credits.domain.CreditConflict
import com.osglab.account.features.credits.domain.UsageMeasurement import com.osglab.account.features.credits.domain.UsageMeasurement
import com.osglab.account.features.credits.repositories.ExposedBillingTransactionRunner import com.osglab.account.features.credits.repositories.ExposedBillingTransactionRunner
@@ -25,6 +30,8 @@ import io.kotest.matchers.ints.shouldBeExactly
import io.kotest.matchers.longs.shouldBeExactly import io.kotest.matchers.longs.shouldBeExactly
import io.kotest.matchers.nulls.shouldBeNull import io.kotest.matchers.nulls.shouldBeNull
import io.kotest.matchers.nulls.shouldNotBeNull import io.kotest.matchers.nulls.shouldNotBeNull
import io.kotest.matchers.shouldBe
import io.kotest.matchers.types.shouldBeInstanceOf
import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.async import kotlinx.coroutines.async
import kotlinx.coroutines.awaitAll import kotlinx.coroutines.awaitAll
@@ -107,6 +114,79 @@ class MySqlSecurityIntegrationTest : FunSpec({
val sessionJwt = SessionJwt(sessionConfig) val sessionJwt = SessionJwt(sessionConfig)
val authenticator = SessionAccessAuthenticator(sessionJwt, authRepository) val authenticator = SessionAccessAuthenticator(sessionJwt, authRepository)
val refreshAccount = UUID.randomUUID()
connection().use {
insertAccount(
it,
refreshAccount,
"refresh-apple-sub",
identity.ofAppleSubject("refresh-apple-sub"),
)
}
val refreshNow = Instant.now()
val originalRefreshToken = "integration-original-refresh-token"
val originalRefreshTokenHash = TokenHash.sha256(originalRefreshToken)
val replacementRefreshToken = "integration-replacement-refresh-token"
val operationId = UUID.randomUUID()
val refreshEncryptor = FieldEncryptor(ByteArray(32) { 9 })
authRepository.createSession(
accountId = refreshAccount,
refreshTokenHash = originalRefreshTokenHash,
expiresAt = refreshNow.plus(Duration.ofDays(30)),
now = refreshNow,
)
val firstRotation = authRepository.rotateRefreshToken(
RefreshRotationAttempt(
currentTokenHash = originalRefreshTokenHash,
newTokenHash = TokenHash.sha256(replacementRefreshToken),
encryptedNewToken = refreshEncryptor.encrypt(
replacementRefreshToken,
refreshReplayContext(originalRefreshTokenHash),
),
newExpiresAt = refreshNow.plus(Duration.ofDays(30)),
operationId = operationId,
replayUntil = refreshNow.plusSeconds(30),
now = refreshNow,
),
).shouldBeInstanceOf<RefreshRotationResult.Rotated>()
val replayedRotation = authRepository.rotateRefreshToken(
RefreshRotationAttempt(
currentTokenHash = originalRefreshTokenHash,
newTokenHash = TokenHash.sha256("discarded-retry-token"),
encryptedNewToken = "discarded-retry-ciphertext",
newExpiresAt = refreshNow.plus(Duration.ofDays(30)),
operationId = operationId,
replayUntil = refreshNow.plusSeconds(31),
now = refreshNow.plusSeconds(1),
),
).shouldBeInstanceOf<RefreshRotationResult.Replayed>()
replayedRotation.sessionId shouldBe firstRotation.sessionId
refreshEncryptor.decrypt(
replayedRotation.encryptedRefreshToken,
refreshReplayContext(originalRefreshTokenHash),
) shouldBe replacementRefreshToken
authRepository.isSessionActive(
refreshAccount,
firstRotation.sessionId,
refreshNow.plusSeconds(1),
) shouldBe true
authRepository.rotateRefreshToken(
RefreshRotationAttempt(
currentTokenHash = originalRefreshTokenHash,
newTokenHash = TokenHash.sha256("attacker-replacement-token"),
encryptedNewToken = "attacker-ciphertext",
newExpiresAt = refreshNow.plus(Duration.ofDays(30)),
operationId = UUID.randomUUID(),
replayUntil = refreshNow.plusSeconds(32),
now = refreshNow.plusSeconds(2),
),
) shouldBe RefreshRotationResult.ReuseDetected
authRepository.isSessionActive(
refreshAccount,
firstRotation.sessionId,
refreshNow.plusSeconds(2),
) shouldBe false
val deletedUser = UUID.randomUUID() val deletedUser = UUID.randomUUID()
val deletedFamily = UUID.randomUUID() val deletedFamily = UUID.randomUUID()
val deletedSession = UUID.randomUUID() val deletedSession = UUID.randomUUID()