Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -22,9 +22,11 @@ internal class JwtProviderImpl @Inject constructor(
private val tokenMutex = Mutex()

private var _cachedAccessToken: AccessToken? = null

override val cachedAccessToken: AccessToken
get() {
val accessToken = _cachedAccessToken ?: throw CannotUseAccessTokenException()
val accessToken =
_cachedAccessToken ?: throw CannotUseAccessTokenException()

if (accessToken.isExpired()) {
throw CannotUseAccessTokenException()
Expand All @@ -33,25 +35,29 @@ internal class JwtProviderImpl @Inject constructor(
return accessToken
}

private val _isCachedAccessTokenAvailable: MutableStateFlow<Boolean> =
private val _isCachedAccessTokenAvailable =
MutableStateFlow(checkIsAccessTokenAvailable())

override val isCachedAccessTokenAvailable: StateFlow<Boolean> =
_isCachedAccessTokenAvailable.asStateFlow()

private var _cachedRefreshToken: RefreshToken? = null

override val cachedRefreshToken: RefreshToken
get() {
if (_cachedRefreshToken == null) {
throw CannotUseRefreshTokenException()
}
if (_cachedRefreshToken!!.isExpired()) {
val refreshToken =
_cachedRefreshToken ?: throw CannotUseRefreshTokenException()

if (refreshToken.isExpired()) {
throw CannotUseRefreshTokenException()
}
return _cachedRefreshToken!!

return refreshToken
}

private val _isCachedRefreshTokenAvailable: MutableStateFlow<Boolean> =
private val _isCachedRefreshTokenAvailable =
MutableStateFlow(checkIsRefreshTokenAvailable())

override val isCachedRefreshTokenAvailable: StateFlow<Boolean> =
_isCachedRefreshTokenAvailable.asStateFlow()

Expand All @@ -63,12 +69,13 @@ internal class JwtProviderImpl @Inject constructor(
runCatching {
jwtDataStoreDataSource.loadTokens()
}.onSuccess { tokens ->
this@JwtProviderImpl._cachedAccessToken = tokens.accessToken
this@JwtProviderImpl._cachedRefreshToken = tokens.refreshToken
_cachedAccessToken = tokens.accessToken
_cachedRefreshToken = tokens.refreshToken
}.onFailure { exception ->
Log.e("JwtProvider", "Failed to persist tokens", exception)
Log.e("JwtProvider", "Failed to load tokens", exception)
}
this.refreshTokenAbility()

refreshTokenAbility()
}

override fun updateTokens(tokens: Tokens) {
Expand All @@ -87,46 +94,45 @@ internal class JwtProviderImpl @Inject constructor(
}
}

override suspend fun resolveSession(): Boolean {
override suspend fun resolveSession(): Boolean =
tokenMutex.withLock {
val accessToken = _cachedAccessToken

if (accessToken != null && !accessToken.isExpired()) {
refreshTokenAbility()
return true
return@withLock true
}

reissueTokensLocked()
val reissued = reissueTokensLocked()
refreshTokenAbility()
return checkIsAccessTokenAvailable() || checkIsRefreshTokenAvailable()

reissued && checkIsAccessTokenAvailable()
}
}

override suspend fun refreshSession(): Boolean {
override suspend fun refreshSession(): Boolean =
tokenMutex.withLock {
reissueTokensLocked()
val reissued = reissueTokensLocked()
refreshTokenAbility()
return checkIsAccessTokenAvailable() || checkIsRefreshTokenAvailable()

reissued && checkIsAccessTokenAvailable()
}
}

private fun refreshTokenAbility() {
_isCachedAccessTokenAvailable.value = checkIsAccessTokenAvailable()
_isCachedRefreshTokenAvailable.value = checkIsRefreshTokenAvailable()
_isCachedAccessTokenAvailable.value =
checkIsAccessTokenAvailable()
_isCachedRefreshTokenAvailable.value =
checkIsRefreshTokenAvailable()
}

private fun checkIsAccessTokenAvailable(): Boolean {
if (this._cachedAccessToken == null) {
return false
}
return !_cachedAccessToken!!.isExpired()
}
private fun checkIsAccessTokenAvailable(): Boolean =
_cachedAccessToken?.let { accessToken ->
!accessToken.isExpired()
} ?: false

private fun checkIsRefreshTokenAvailable(): Boolean {
if (this._cachedRefreshToken == null) {
return false
}
return !_cachedRefreshToken!!.isExpired()
}
private fun checkIsRefreshTokenAvailable(): Boolean =
_cachedRefreshToken?.let { refreshToken ->
!refreshToken.isExpired()
} ?: false

private suspend fun reissueTokensLocked(): Boolean {
val refreshToken = _cachedRefreshToken
Expand All @@ -138,13 +144,20 @@ internal class JwtProviderImpl @Inject constructor(
}

return try {
val tokens = jwtReissueManager(refreshToken = refreshToken.value)
val tokens = jwtReissueManager(
refreshToken = refreshToken.value,
)

updateTokensLocked(tokens = tokens)
true
} catch (exception: CannotReissueTokenException) {
if (exception.statusCode == 401 || exception.statusCode == 404) {
if (
exception.statusCode == 401 ||
exception.statusCode == 404
) {
clearCachesLocked()
}

false
}
}
Expand All @@ -153,39 +166,55 @@ internal class JwtProviderImpl @Inject constructor(
val previousAccessToken = _cachedAccessToken
val previousRefreshToken = _cachedRefreshToken

this._cachedAccessToken = tokens.accessToken
this._cachedRefreshToken = tokens.refreshToken
this.refreshTokenAbility()
_cachedAccessToken = tokens.accessToken
_cachedRefreshToken = tokens.refreshToken
refreshTokenAbility()

runCatchingCancellable {
jwtDataStoreDataSource.storeTokens(tokens = tokens)
}.onFailure { exception ->
this._cachedAccessToken = previousAccessToken
this._cachedRefreshToken = previousRefreshToken
this.refreshTokenAbility()
_cachedAccessToken = previousAccessToken
_cachedRefreshToken = previousRefreshToken
refreshTokenAbility()

Log.e("JwtProvider", "Failed to store tokens", exception)
throw IllegalStateException("Failed to persist tokens", exception)
Log.e(
"JwtProvider",
"Failed to store tokens",
exception,
)

throw IllegalStateException(
"Failed to persist tokens",
exception,
)
}
}

private suspend fun clearCachesLocked() {
val previousAccessToken = _cachedAccessToken
val previousRefreshToken = _cachedRefreshToken

this._cachedAccessToken = null
this._cachedRefreshToken = null
this.refreshTokenAbility()
_cachedAccessToken = null
_cachedRefreshToken = null
refreshTokenAbility()

runCatchingCancellable {
jwtDataStoreDataSource.clearTokens()
}.onFailure { exception ->
this._cachedAccessToken = previousAccessToken
this._cachedRefreshToken = previousRefreshToken
this.refreshTokenAbility()
_cachedAccessToken = previousAccessToken
_cachedRefreshToken = previousRefreshToken
refreshTokenAbility()

Log.e(
"JwtProvider",
"Failed to clear tokens",
exception,
)

Log.e("JwtProvider", "Failed to clear tokens", exception)
throw IllegalStateException("Failed to clear persisted tokens", exception)
throw IllegalStateException(
"Failed to clear persisted tokens",
exception,
)
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -10,14 +10,15 @@ import okhttp3.RequestBody.Companion.toRequestBody
import okhttp3.ResponseBody
import okhttp3.logging.HttpLoggingInterceptor
import team.aliens.dms.android.core.jwt.Tokens
import team.aliens.dms.android.core.jwt.di.TokenReissueUrl
import team.aliens.dms.android.core.jwt.network.exception.CannotReissueTokenException
import team.aliens.dms.android.core.jwt.network.model.TokensResponse
import team.aliens.dms.android.core.jwt.toModel
import team.aliens.dms.android.core.network.di.DefaultHttpLoggingInterceptor
import javax.inject.Inject

class JwtReissueManager @Inject constructor(
private val reissueUrl: String,
@TokenReissueUrl private val reissueUrl: String,
@DefaultHttpLoggingInterceptor private val httpLoggingInterceptor: HttpLoggingInterceptor,
baseHttpClient: OkHttpClient,
) {
Expand All @@ -27,27 +28,40 @@ class JwtReissueManager @Inject constructor(
}.build()
}

suspend operator fun invoke(refreshToken: String): Tokens = withContext(Dispatchers.IO) {
suspend operator fun invoke(
refreshToken: String,
): Tokens = withContext(Dispatchers.IO) {
val request = buildTokenReissueRequest(refreshToken)
val response = client.newCall(request).execute()

if (response.isSuccessful) {
response.body.toTokensResponse().toModel()
} else {
throw CannotReissueTokenException(statusCode = response.code)
client.newCall(request).execute().use { response ->
if (response.isSuccessful) {
response.body.toTokensResponse().toModel()
} else {
throw CannotReissueTokenException(
statusCode = response.code,
)
}
}
}

private fun ResponseBody?.toTokensResponse(): TokensResponse {
requireNotNull(this)
return Gson().fromJson(this.string(), TokensResponse::class.java)
return Gson().fromJson(string(), TokensResponse::class.java)
}

private fun buildTokenReissueRequest(refreshToken: String): Request =
Request.Builder().url(reissueUrl).put(
body = String().toRequestBody("application/json".toMediaType()),
).addHeader(
name = "refresh-token",
value = refreshToken,
).build()
private fun buildTokenReissueRequest(
refreshToken: String,
): Request =
Request.Builder()
.url(reissueUrl)
.put(
body = String().toRequestBody(
"application/json".toMediaType(),
),
)
.addHeader(
name = "refresh-token",
value = refreshToken,
)
.build()
}
Original file line number Diff line number Diff line change
Expand Up @@ -21,26 +21,48 @@ class JwtAuthenticator @Inject constructor(
): Request? {
val request = response.request

if (request.shouldBeIgnored() || response.retryCount >= MAX_AUTH_RETRY_COUNT) {
if (
request.shouldBeIgnored() ||
response.retryCount >= MAX_AUTH_RETRY_COUNT
) {
return null
}

val newAuthorization = refreshAuthorization()
val failedAuthorization = request.header(AUTHORIZATION_HEADER)
val newAuthorization = refreshAuthorization(
failedAuthorization = failedAuthorization,
)

return newAuthorization
?.takeUnless { refreshedAuthorization ->
refreshedAuthorization == request.header(AUTHORIZATION_HEADER)
?.takeUnless { authorization ->
authorization == failedAuthorization
}
?.let { refreshedAuthorization ->
?.let { authorization ->
request.newBuilder()
.header(AUTHORIZATION_HEADER, refreshedAuthorization)
.header(AUTHORIZATION_HEADER, authorization)
.build()
}
}

private fun refreshAuthorization(): String? {
@Synchronized
private fun refreshAuthorization(
failedAuthorization: String?,
): String? {
val currentAuthorization = runCatching {
"Bearer ${jwtProvider.cachedAccessToken.value}"
}.getOrNull()

if (
currentAuthorization != null &&
currentAuthorization != failedAuthorization
) {
return currentAuthorization
}

val refreshed = runCatching {
runBlocking { jwtProvider.refreshSession() }
runBlocking {
jwtProvider.refreshSession()
}
}.getOrDefault(false)

if (!refreshed) {
Expand All @@ -52,12 +74,15 @@ class JwtAuthenticator @Inject constructor(
}.getOrNull()
}

private fun Request.shouldBeIgnored(): Boolean = ignoreRequests.requests.any { ignoreRequest ->
val path = this@shouldBeIgnored.url.encodedPath
val method = this@shouldBeIgnored.method.toHttpMethod()
private fun Request.shouldBeIgnored(): Boolean =
checkS3Request(url.toString()) ||
ignoreRequests.requests.any { ignoreRequest ->
val path = url.encodedPath
val requestMethod = method.toHttpMethod()

path.contains(ignoreRequest.path) && method == ignoreRequest.method || checkS3Request(url = this@shouldBeIgnored.url.toString())
}
path.contains(ignoreRequest.path) &&
requestMethod == ignoreRequest.method
}

private val Response.retryCount: Int
get() {
Expand All @@ -72,7 +97,8 @@ class JwtAuthenticator @Inject constructor(
return count
}

private fun checkS3Request(url: String): Boolean = url.contains(ResourceKeys.IMAGE_URL)
private fun checkS3Request(url: String): Boolean =
url.contains(ResourceKeys.IMAGE_URL)

private companion object {
const val AUTHORIZATION_HEADER = "authorization"
Expand Down
Loading
Loading