package com.osglab.account.features.oobe import com.nimbusds.jose.JWSAlgorithm import com.nimbusds.jose.JWSHeader import com.nimbusds.jose.crypto.MACSigner import com.nimbusds.jose.crypto.MACVerifier import com.nimbusds.jwt.JWTClaimsSet import com.nimbusds.jwt.SignedJWT import com.osglab.account.features.gateway.models.GatewayCapability import com.osglab.account.features.gateway.models.GatewayPrincipal import com.osglab.account.features.gateway.models.GatewaySubjectType import com.osglab.account.features.integrity.AppAttestService import java.nio.charset.StandardCharsets import java.security.MessageDigest import java.time.Clock import java.time.Duration import java.util.Base64 import java.util.Date import java.util.UUID import javax.crypto.Mac import javax.crypto.spec.SecretKeySpec class OobeGrantService( private val repository: OobeRepository, private val appAttest: AppAttestService, private val settings: OobeTokenSettings, private val clock: Clock = Clock.systemUTC(), ) { suspend fun create(request: CreateOobeGrantRequest): OobeGrantTokens { val installationId = canonicalInstallationId(request.installationId) val challenge = decodeChallenge(request.challenge) val canonicalPayload = OobeContract.canonicalAssertionPayload( challenge = challenge, keyId = request.keyId, installationId = installationId, ) appAttest.verifyBoundAssertion( challengeId = request.challengeId, challenge = challenge, keyId = request.keyId, assertionObject = request.assertion, expectedClientDataHash = sha256(canonicalPayload), ) val now = clock.instant() val subject = repository.findOrCreateSubject( keyId = request.keyId, installationHash = sha256Hex(installationId.toByteArray(StandardCharsets.UTF_8)), subjectId = UUID.randomUUID().toString(), now = now, ) val grantId = UUID.randomUUID().toString() val tokenId = UUID.randomUUID().toString() val familyId = UUID.randomUUID().toString() val grantExpiresAt = now.plus(GRANT_LIFETIME) val refreshToken = refreshToken(grantId, familyId, tokenId) val stored = repository.createGrant( NewOobeGrant( grant = OobeGrant(grantId, subject.id, grantExpiresAt), refreshTokenId = tokenId, refreshFamilyId = familyId, refreshTokenHash = tokenHash(refreshToken), refreshExpiresAt = grantExpiresAt, ), now, ) return issue(stored) } suspend fun refresh(refreshToken: String, idempotencyKey: String): OobeGrantTokens { require(IDEMPOTENCY_KEY.matches(idempotencyKey)) { "Idempotency key is invalid" } if (refreshToken.length !in 32..MAX_REFRESH_TOKEN_CHARS) { throw OobeRefreshTokenInvalidException() } parseRefreshToken(refreshToken) val now = clock.instant() val tokenId = UUID.randomUUID().toString() val result = repository.rotateRefresh( currentTokenHash = tokenHash(refreshToken), rotationIdempotencyKey = idempotencyKey, newTokenId = tokenId, newTokenHash = tokenHash(replaceTokenId(refreshToken, tokenId)), newExpiresAt = now.plus(GRANT_LIFETIME), now = now, ) return when (result) { is OobeRefreshRotationResult.Rotated -> issue(result.refresh) OobeRefreshRotationResult.Invalid -> throw OobeRefreshTokenInvalidException() OobeRefreshRotationResult.ReuseDetected -> throw OobeRefreshTokenReuseException() } } suspend fun authenticate(serialized: String): GatewayPrincipal? { val principal = verifyAccessToken(serialized) ?: return null return repository.findActiveGrant( grantId = requireNotNull(principal.grantId), subjectId = principal.userId, now = clock.instant(), )?.let { principal } } private fun issue(refresh: StoredOobeRefresh): OobeGrantTokens { val now = clock.instant() val accessExpiresAt = minOf(now.plus(ACCESS_LIFETIME), refresh.grant.expiresAt) require(accessExpiresAt.isAfter(now)) { "OOBE gateway grant has expired" } require(refresh.expiresAt.isAfter(now)) { "OOBE refresh token has expired" } val claims = JWTClaimsSet.Builder() .issuer(settings.issuer) .audience(settings.audience) .subject("$SUBJECT_PREFIX${refresh.grant.subjectId}") .jwtID(UUID.randomUUID().toString()) .issueTime(Date.from(now)) .notBeforeTime(Date.from(now.minusSeconds(CLOCK_SKEW_SECONDS))) .expirationTime(Date.from(accessExpiresAt)) .claim(CLAIM_TYPE, ACCESS_TOKEN_TYPE) .claim(CLAIM_GRANT_ID, refresh.grant.id) .claim(CLAIM_SCOPES, OobeContract.scopes.map { it.name.lowercase() }.sorted()) .claim(CLAIM_FEATURES, OobeContract.features.map { it.name.lowercase() }.sorted()) .build() val jwt = SignedJWT(JWSHeader(JWSAlgorithm.HS256), claims) jwt.sign(MACSigner(settings.accessTokenHmacSecret)) return OobeGrantTokens( grantId = refresh.grant.id, scopes = OobeContract.scopes, features = OobeContract.features, accessToken = jwt.serialize(), accessExpiresAt = accessExpiresAt.toString(), refreshToken = refreshToken(refresh.grant.id, refresh.familyId, refresh.tokenId), refreshExpiresAt = refresh.expiresAt.toString(), ) } private fun verifyAccessToken(serialized: String): GatewayPrincipal? = runCatching { val jwt = SignedJWT.parse(serialized) require(jwt.header.algorithm == JWSAlgorithm.HS256) require(jwt.verify(MACVerifier(settings.accessTokenHmacSecret))) val claims = jwt.jwtClaimsSet val now = clock.instant() require(claims.issuer == settings.issuer) require(settings.audience in claims.audience) require(claims.getStringClaim(CLAIM_TYPE) == ACCESS_TOKEN_TYPE) require(claims.expirationTime?.toInstant()?.isAfter(now) == true) require(claims.notBeforeTime?.toInstant()?.isBefore(now.plusSeconds(CLOCK_SKEW_SECONDS)) != false) require(claims.issueTime?.toInstant()?.isAfter(now.plusSeconds(CLOCK_SKEW_SECONDS)) != true) require(claims.getStringListClaim(CLAIM_SCOPES).map(String::uppercase) .map(GatewayCapability::valueOf).toSet() == OobeContract.scopes) require(claims.getStringListClaim(CLAIM_FEATURES).map(String::uppercase).toSet() == OobeContract.features.map { it.name }.toSet()) val subject = claims.subject require(subject.startsWith(SUBJECT_PREFIX)) val subjectId = UUID.fromString(subject.removePrefix(SUBJECT_PREFIX)).toString() GatewayPrincipal( userId = subjectId, grantId = UUID.fromString(claims.getStringClaim(CLAIM_GRANT_ID)).toString(), scopes = OobeContract.scopes, subjectType = GatewaySubjectType.OOBE, ) }.getOrNull() private fun refreshToken(grantId: String, familyId: String, tokenId: String): String { val publicPart = "$grantId.$familyId.$tokenId" val mac = Mac.getInstance(HMAC_ALGORITHM) mac.init(SecretKeySpec(settings.refreshTokenHmacSecret, HMAC_ALGORITHM)) val secret = Base64.getUrlEncoder().withoutPadding() .encodeToString(mac.doFinal("$REFRESH_CONTEXT:$publicPart".toByteArray(StandardCharsets.US_ASCII))) return "$REFRESH_PREFIX$publicPart.$secret" } private fun parseRefreshToken(value: String) { if (!value.startsWith(REFRESH_PREFIX)) throw OobeRefreshTokenInvalidException() val parts = value.removePrefix(REFRESH_PREFIX).split('.') if (parts.size != 4) throw OobeRefreshTokenInvalidException() val grantId = canonicalUuid(parts[0]) val familyId = canonicalUuid(parts[1]) val tokenId = canonicalUuid(parts[2]) val expected = refreshToken(grantId, familyId, tokenId) if (!MessageDigest.isEqual( expected.toByteArray(StandardCharsets.US_ASCII), value.toByteArray(StandardCharsets.US_ASCII), ) ) { throw OobeRefreshTokenInvalidException() } } private fun replaceTokenId(value: String, newTokenId: String): String { val parts = value.removePrefix(REFRESH_PREFIX).split('.') return refreshToken(parts[0], parts[1], newTokenId) } private fun canonicalInstallationId(value: String): String = runCatching { UUID.fromString(value).toString() } .getOrElse { throw IllegalArgumentException("installationId must be a UUID") } private fun canonicalUuid(value: String): String = runCatching { UUID.fromString(value).toString() } .getOrElse { throw OobeRefreshTokenInvalidException() } private fun decodeChallenge(value: String): ByteArray = runCatching { Base64.getUrlDecoder().decode(value) } .getOrElse { throw IllegalArgumentException("challenge must be Base64URL") } .also { require(it.size == CHALLENGE_BYTES) { "challenge size is invalid" } } private fun tokenHash(value: String): String = sha256Hex(value.toByteArray(StandardCharsets.US_ASCII)) private companion object { val GRANT_LIFETIME: Duration = Duration.ofMinutes(30) val ACCESS_LIFETIME: Duration = Duration.ofMinutes(5) const val HMAC_ALGORITHM = "HmacSHA256" const val REFRESH_CONTEXT = "oobe-refresh" const val CLAIM_TYPE = "typ" const val CLAIM_GRANT_ID = "gid" const val CLAIM_SCOPES = "scp" const val CLAIM_FEATURES = "features" const val ACCESS_TOKEN_TYPE = "oobe_gateway_access" const val SUBJECT_PREFIX = "oobe:" const val REFRESH_PREFIX = "oobert_" const val CLOCK_SKEW_SECONDS = 30L const val CHALLENGE_BYTES = 32 const val MAX_REFRESH_TOKEN_CHARS = 512 val IDEMPOTENCY_KEY = Regex("[A-Za-z0-9._:-]{8,128}") } } class OobeRefreshTokenInvalidException : RuntimeException("OOBE refresh token is invalid") class OobeRefreshTokenReuseException : RuntimeException("OOBE refresh token reuse was detected") data class OobeTokenSettings( val issuer: String, val audience: String, val accessTokenHmacSecret: ByteArray, val refreshTokenHmacSecret: ByteArray, ) private fun sha256(value: ByteArray): ByteArray = MessageDigest.getInstance("SHA-256").digest(value) private fun sha256Hex(value: ByteArray): String = sha256(value).joinToString("") { "%02x".format(it.toInt() and 0xff) }