Rework token invalidation (#307)

This commit is contained in:
KirillPamPam
2023-09-25 16:02:36 +04:00
committed by GitHub
parent 109fb97e0a
commit 8e26d4a8cf
8 changed files with 51 additions and 104 deletions

View File

@@ -1,20 +1,22 @@
package io.emeraldpay.dshackle.auth package io.emeraldpay.dshackle.auth
import com.github.benmanes.caffeine.cache.Caffeine
import org.springframework.stereotype.Component
import java.time.Duration
import java.time.Instant import java.time.Instant
import java.util.concurrent.ConcurrentHashMap
@Component
class AuthContext { class AuthContext {
private val sessions = Caffeine.newBuilder()
.expireAfterAccess(Duration.ofDays(1))
.build<String, Boolean>()
companion object { fun putSessionInContext(tokenWrapper: TokenWrapper) {
val sessions = ConcurrentHashMap<String, TokenWrapper>() sessions.put(tokenWrapper.sessionId, true)
}
fun putTokenInContext(tokenWrapper: TokenWrapper) { fun containsSession(sessionId: String): Boolean {
sessions[tokenWrapper.sessionId] = tokenWrapper return sessions.asMap()[sessionId] != null
}
fun removeToken(sessionId: String) {
sessions.remove(sessionId)
}
} }
data class TokenWrapper( data class TokenWrapper(

View File

@@ -13,7 +13,9 @@ const val AUTH_METHOD_NAME = "emerald.Auth/Authenticate"
const val REFLECT_METHOD_NAME = "grpc.reflection.v1alpha.ServerReflection/ServerReflectionInfo" const val REFLECT_METHOD_NAME = "grpc.reflection.v1alpha.ServerReflection/ServerReflectionInfo"
@Component @Component
class AuthInterceptor : ServerInterceptor { class AuthInterceptor(
private val authContext: AuthContext
) : ServerInterceptor {
private val specialMethods = setOf(AUTH_METHOD_NAME, REFLECT_METHOD_NAME) private val specialMethods = setOf(AUTH_METHOD_NAME, REFLECT_METHOD_NAME)
override fun <ReqT : Any, RespT : Any> interceptCall( override fun <ReqT : Any, RespT : Any> interceptCall(
@@ -26,7 +28,7 @@ class AuthInterceptor : ServerInterceptor {
) )
val isOrdinaryMethod = !specialMethods.contains(call.methodDescriptor.fullMethodName) val isOrdinaryMethod = !specialMethods.contains(call.methodDescriptor.fullMethodName)
if (isOrdinaryMethod && (sessionId == null || !AuthContext.sessions.containsKey(sessionId))) { if (isOrdinaryMethod && (sessionId == null || !authContext.containsSession(sessionId))) {
val cause = if (sessionId == null) "sessionId is not passed" else "Session $sessionId does not exist" val cause = if (sessionId == null) "sessionId is not passed" else "Session $sessionId does not exist"
throw Status.UNAUTHENTICATED throw Status.UNAUTHENTICATED
.withDescription(cause) .withDescription(cause)

View File

@@ -2,6 +2,7 @@ package io.emeraldpay.dshackle.auth.processor
import com.auth0.jwt.JWT import com.auth0.jwt.JWT
import com.auth0.jwt.JWTVerifier import com.auth0.jwt.JWTVerifier
import com.auth0.jwt.RegisteredClaims
import com.auth0.jwt.algorithms.Algorithm import com.auth0.jwt.algorithms.Algorithm
import io.emeraldpay.dshackle.auth.AuthContext import io.emeraldpay.dshackle.auth.AuthContext
import io.emeraldpay.dshackle.auth.service.KeyReader import io.emeraldpay.dshackle.auth.service.KeyReader
@@ -13,6 +14,7 @@ import java.security.PublicKey
import java.security.interfaces.RSAPrivateKey import java.security.interfaces.RSAPrivateKey
import java.security.interfaces.RSAPublicKey import java.security.interfaces.RSAPublicKey
import java.time.Instant import java.time.Instant
import java.time.temporal.ChronoUnit
import java.util.UUID import java.util.UUID
const val SESSION_ID = "sessionId" const val SESSION_ID = "sessionId"
@@ -37,6 +39,9 @@ abstract class AuthProcessor(
try { try {
val verifier: JWTVerifier = JWT.require(verifyingAlgorithm(keys.externalPublicKey)) val verifier: JWTVerifier = JWT.require(verifyingAlgorithm(keys.externalPublicKey))
.withIssuer(authorizationConfig.publicKeyOwner) .withIssuer(authorizationConfig.publicKeyOwner)
.withClaim(RegisteredClaims.ISSUED_AT) { claim, _ ->
claim.asInstant().plus(1, ChronoUnit.MINUTES).isAfter(Instant.now())
}
.build() .build()
verifier.verify(token) verifier.verify(token)
} catch (e: Exception) { } catch (e: Exception) {

View File

@@ -1,26 +0,0 @@
package io.emeraldpay.dshackle.auth.processor
import io.emeraldpay.dshackle.auth.AuthContext
import org.slf4j.LoggerFactory
import org.springframework.scheduling.annotation.Scheduled
import org.springframework.stereotype.Component
import java.time.Instant
import java.time.temporal.ChronoUnit
@Component
open class TokenProcessor {
companion object {
private val log = LoggerFactory.getLogger(TokenProcessor::class.java)
}
@Scheduled(fixedRate = 30000)
fun invalidateTokens() {
AuthContext.sessions
.filter { Instant.now().isAfter(it.value.issuedAt.plus(1, ChronoUnit.HOURS)) }
.forEach {
log.info("Invalidate token with sessionId ${it.key}")
AuthContext.removeToken(it.key)
}
}
}

View File

@@ -11,7 +11,8 @@ import org.springframework.stereotype.Service
class AuthService( class AuthService(
private val authorizationConfig: AuthorizationConfig, private val authorizationConfig: AuthorizationConfig,
private val rsaKeyReader: KeyReader, private val rsaKeyReader: KeyReader,
private val authProcessorResolver: AuthProcessorResolver private val authProcessorResolver: AuthProcessorResolver,
private val authContext: AuthContext
) { ) {
fun authenticate(token: String): String { fun authenticate(token: String): String {
@@ -31,7 +32,7 @@ class AuthService(
.getAuthProcessor(decodedJwt) .getAuthProcessor(decodedJwt)
.process(keys, token) .process(keys, token)
.run { .run {
AuthContext.putTokenInContext(this) authContext.putSessionInContext(this)
this.token this.token
} }
} }

View File

@@ -17,9 +17,13 @@ import java.io.StringReader
import java.nio.file.Files import java.nio.file.Files
import java.nio.file.Paths import java.nio.file.Paths
import java.security.KeyFactory import java.security.KeyFactory
import java.security.PrivateKey
import java.security.PublicKey import java.security.PublicKey
import java.security.interfaces.RSAPrivateKey
import java.security.interfaces.RSAPublicKey import java.security.interfaces.RSAPublicKey
import java.security.spec.PKCS8EncodedKeySpec
import java.security.spec.X509EncodedKeySpec import java.security.spec.X509EncodedKeySpec
import java.time.Instant
class AuthProcessorV1Test { class AuthProcessorV1Test {
private val processor = AuthProcessorV1( private val processor = AuthProcessorV1(
@@ -31,12 +35,13 @@ class AuthProcessorV1Test {
) )
private val rsaKeyReader = RsaKeyReader() private val rsaKeyReader = RsaKeyReader()
private val privProviderPath = ResourceUtils.getFile("classpath:keys/priv.p8.key").path private val privProviderPath = ResourceUtils.getFile("classpath:keys/priv.p8.key").path
private val privDrpcPath = ResourceUtils.getFile("classpath:keys/priv-drpc.p8.key").path
private val publicDrpcPath = ResourceUtils.getFile("classpath:keys/public-drpc.pem").path private val publicDrpcPath = ResourceUtils.getFile("classpath:keys/public-drpc.pem").path
private val token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiJkcnBjIiwiaWF0IjoxNjkyMTg1OTMxLCJ2ZXJzaW9uI" + private val token = JWT.create()
"joiVjEifQ.BZILN0GQ7JzXGFz-GZIbFTT9E5L-miB4Nga0v4o_cQThk8gbDelBRzEfdsqxCq_ppPr3v_Own8M-vR9yQElx5nEdlI4xe5QAMdIvr3g" + .withIssuedAt(Instant.now())
"12fMckydX9IsW4sVQ1kJJY8RrHb-WL-uI0WSWqoMSwf-Psb-UyiEHAjc3oK7fA72lBaGT4waPHOxRBPvezwg7N934vCZvZMAftFfVgmeEtbCeD7bF" + .withIssuer("drpc")
"umEr0uEmkIKPTg4QwP-VMvqoLBYpMiJVzP_Ipg_wRHJ7fUN0BGEPjjMvhQ_6TWByiQUBz1kTMd0Ebf_kEuXFQeiwA-FXHJpWczzh66CbbmmWAWsi" + .withClaim(VERSION, AuthVersion.V1.toString())
"ehKw3KPZeBj0oQ" .sign(Algorithm.RSA256(generatePrivateKey(privDrpcPath) as RSAPrivateKey))
private val keyPair = rsaKeyReader.getKeyPair(privProviderPath, publicDrpcPath) private val keyPair = rsaKeyReader.getKeyPair(privProviderPath, publicDrpcPath)
@Test @Test
@@ -88,4 +93,14 @@ class AuthProcessorV1Test {
return KeyFactory.getInstance("RSA").generatePublic(publicKeySpec) return KeyFactory.getInstance("RSA").generatePublic(publicKeySpec)
} }
private fun generatePrivateKey(path: String): PrivateKey {
val privateKeyReader = StringReader(Files.readString(Paths.get(path)))
val privatePem = PEMParser(privateKeyReader).readPemObject()
val privateKeySpec = PKCS8EncodedKeySpec(privatePem.content)
return KeyFactory.getInstance("RSA").generatePrivate(privateKeySpec)
}
} }

View File

@@ -1,53 +0,0 @@
package io.emeraldpay.dshackle.auth.processor
import io.emeraldpay.dshackle.auth.AuthContext
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.BeforeEach
import org.junit.jupiter.api.Test
import java.time.Instant
import java.time.temporal.ChronoUnit
class TokenProcessorTest {
private val tokenProcessor = TokenProcessor()
@BeforeEach
fun removeSessions() {
AuthContext.sessions.clear()
}
@Test
fun `invalidate all tokens`() {
AuthContext.putTokenInContext(
AuthContext.TokenWrapper("token", Instant.now().minus(1, ChronoUnit.HOURS), "session1")
)
AuthContext.putTokenInContext(
AuthContext.TokenWrapper("token", Instant.now().minus(1, ChronoUnit.HOURS), "session2")
)
AuthContext.putTokenInContext(
AuthContext.TokenWrapper("token", Instant.now().minus(1, ChronoUnit.HOURS), "session3")
)
tokenProcessor.invalidateTokens()
assertTrue(AuthContext.sessions.isEmpty())
}
@Test
fun `tokens are still in the context after invalidation`() {
val token1 = AuthContext.TokenWrapper("token", Instant.now().minus(30, ChronoUnit.MINUTES), "session1")
val token2 = AuthContext.TokenWrapper("token", Instant.now().minus(30, ChronoUnit.MINUTES), "session2")
val token3 = AuthContext.TokenWrapper("token", Instant.now().minus(30, ChronoUnit.MINUTES), "session3")
AuthContext.putTokenInContext(token1)
AuthContext.putTokenInContext(token2)
AuthContext.putTokenInContext(token3)
tokenProcessor.invalidateTokens()
assertEquals(3, AuthContext.sessions.size)
assertEquals(
mapOf(token1.sessionId to token1, token2.sessionId to token2, token3.sessionId to token3),
AuthContext.sessions
)
}
}

View File

@@ -21,6 +21,7 @@ import java.util.concurrent.CompletableFuture
class AuthServiceTest { class AuthServiceTest {
private val rsaKeyReader = mock(KeyReader::class.java) private val rsaKeyReader = mock(KeyReader::class.java)
private val mockV1Processor = mock(AuthProcessor::class.java) private val mockV1Processor = mock(AuthProcessor::class.java)
private val authContext = AuthContext()
private val factory = AuthProcessorResolver(mockV1Processor) private val factory = AuthProcessorResolver(mockV1Processor)
private val token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiJkcnBjIiwiaWF0IjoxNjkyMTg1OTMxLCJ2ZXJzaW9uI" + private val token = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJpc3MiOiJkcnBjIiwiaWF0IjoxNjkyMTg1OTMxLCJ2ZXJzaW9uI" +
@@ -31,7 +32,7 @@ class AuthServiceTest {
@Test @Test
fun `unimplemented error if auth is disabled`() { fun `unimplemented error if auth is disabled`() {
val authService = AuthService(AuthorizationConfig.default(), rsaKeyReader, factory) val authService = AuthService(AuthorizationConfig.default(), rsaKeyReader, factory, authContext)
val e = assertThrows(StatusException::class.java) { authService.authenticate("") } val e = assertThrows(StatusException::class.java) { authService.authenticate("") }
assertEquals("UNIMPLEMENTED: Authentication process is not enabled", e.message) assertEquals("UNIMPLEMENTED: Authentication process is not enabled", e.message)
@@ -48,7 +49,7 @@ class AuthServiceTest {
AuthorizationConfig.ServerConfig("privPath", "pubPath"), AuthorizationConfig.ServerConfig("privPath", "pubPath"),
AuthorizationConfig.ClientConfig.default() AuthorizationConfig.ClientConfig.default()
), ),
rsaKeyReader, factory rsaKeyReader, factory, authContext
) )
val pair = KeyReader.Keys(mock(PrivateKey::class.java), mock(PublicKey::class.java)) val pair = KeyReader.Keys(mock(PrivateKey::class.java), mock(PublicKey::class.java))
@@ -59,7 +60,7 @@ class AuthServiceTest {
authService.authenticate(token) authService.authenticate(token)
verify(rsaKeyReader).getKeyPair("privPath", "pubPath") verify(rsaKeyReader).getKeyPair("privPath", "pubPath")
verify(mockV1Processor).process(pair, token) verify(mockV1Processor).process(pair, token)
assertTrue(AuthContext.sessions.containsKey(tokenWrapper.sessionId)) assertTrue(authContext.containsSession(tokenWrapper.sessionId))
} }
@Test @Test
@@ -77,7 +78,7 @@ class AuthServiceTest {
AuthorizationConfig.ServerConfig("privPath", "pubPath"), AuthorizationConfig.ServerConfig("privPath", "pubPath"),
AuthorizationConfig.ClientConfig.default() AuthorizationConfig.ClientConfig.default()
), ),
rsaKeyReader, factory rsaKeyReader, factory, authContext
) )
`when`(rsaKeyReader.getKeyPair("privPath", "pubPath")).thenReturn(pair) `when`(rsaKeyReader.getKeyPair("privPath", "pubPath")).thenReturn(pair)
@@ -93,7 +94,7 @@ class AuthServiceTest {
verify(rsaKeyReader, times(2)).getKeyPair("privPath", "pubPath") verify(rsaKeyReader, times(2)).getKeyPair("privPath", "pubPath")
verify(mockV1Processor, times(2)).process(pair, token) verify(mockV1Processor, times(2)).process(pair, token)
assertTrue(AuthContext.sessions.containsKey(tokenWrapper.sessionId)) assertTrue(authContext.containsSession(tokenWrapper.sessionId))
assertTrue(AuthContext.sessions.containsKey(tokenWrapper1.sessionId)) assertTrue(authContext.containsSession(tokenWrapper1.sessionId))
} }
} }