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
@@ -0,0 +1,125 @@
package co.nilin.opex.api.app.interceptor

import co.nilin.opex.common.security.JwtUtils
import com.fasterxml.jackson.databind.ObjectMapper
import org.reactivestreams.Publisher
import org.slf4j.LoggerFactory
import org.springframework.core.Ordered
import org.springframework.core.annotation.Order
import org.springframework.core.io.buffer.DataBufferUtils
import org.springframework.http.HttpHeaders
import org.springframework.http.server.reactive.ServerHttpRequestDecorator
import org.springframework.http.server.reactive.ServerHttpResponseDecorator
import org.springframework.stereotype.Component
import org.springframework.web.server.ServerWebExchange
import org.springframework.web.server.WebFilter
import org.springframework.web.server.WebFilterChain
import reactor.core.publisher.Flux
import reactor.core.publisher.Mono
import java.nio.charset.StandardCharsets
import java.time.OffsetDateTime
import java.time.ZoneOffset

@Component
@Order(Ordered.LOWEST_PRECEDENCE)
class RequestAuditFilter(
private val objectMapper: ObjectMapper
) : WebFilter {

private val logger = LoggerFactory.getLogger(RequestAuditFilter::class.java)
private val maxPayloadSize = 10_000

override fun filter(exchange: ServerWebExchange, chain: WebFilterChain): Mono<Void> {
val request = exchange.request
val sourceIp = resolveClientIp(exchange)
val token = extractBearerToken(request.headers)
val mobile = extractClaim(token, "mobile", "phone_number")
val email = extractClaim(token, "email")
val deviceUuid = extractClaim(token, "deviceUuid", "device_uuid")

return DataBufferUtils.join(request.body)
.defaultIfEmpty(exchange.response.bufferFactory().wrap(ByteArray(0)))
.flatMap { requestBuffer ->
val requestBytes = ByteArray(requestBuffer.readableByteCount())
requestBuffer.read(requestBytes)
DataBufferUtils.release(requestBuffer)
val requestBody = truncateBody(String(requestBytes, StandardCharsets.UTF_8))
Comment thread
fatemeh-i marked this conversation as resolved.

val decoratedRequest = object : ServerHttpRequestDecorator(request) {
override fun getBody() = Flux.just(exchange.response.bufferFactory().wrap(requestBytes))
}

val responseBody = StringBuilder()
val decoratedResponse = object : ServerHttpResponseDecorator(exchange.response) {
override fun writeWith(body: Publisher<out org.springframework.core.io.buffer.DataBuffer>): Mono<Void> {
val wrapped = Flux.from(body).map { dataBuffer ->
val bytes = ByteArray(dataBuffer.readableByteCount())
dataBuffer.read(bytes)
DataBufferUtils.release(dataBuffer)
responseBody.append(String(bytes, StandardCharsets.UTF_8))
bufferFactory().wrap(bytes)
}
return super.writeWith(wrapped)
}

override fun writeAndFlushWith(body: Publisher<out Publisher<out org.springframework.core.io.buffer.DataBuffer>>): Mono<Void> {
return writeWith(Flux.from(body).flatMapSequential { it })
}
}

val updatedExchange = exchange.mutate()
.request(decoratedRequest)
.response(decoratedResponse)
.build()

return@flatMap chain.filter(updatedExchange)
.doFinally {
val payload = mapOf(
"date" to OffsetDateTime.now(ZoneOffset.UTC).toString(),
"ip" to sourceIp,
"mobile" to mobile,
"email" to email,
"deviceUuid" to deviceUuid,
"method" to request.method.name(),
"url" to request.uri.toString(),
"requestData" to requestBody,
"responseStatus" to (decoratedResponse.statusCode?.value() ?: updatedExchange.response.statusCode?.value()),
"responseBody" to truncateBody(responseBody.toString())
)
runCatching {
logger.info("API_REQUEST_AUDIT {}", objectMapper.writeValueAsString(payload))
}.onFailure {
logger.warn("Failed to write request audit log", it)
}
}
}
}

private fun extractBearerToken(headers: HttpHeaders): String? {
val header = headers.getFirst(HttpHeaders.AUTHORIZATION) ?: return null
if (!header.startsWith("Bearer ", true)) return null
return header.substringAfter("Bearer ").trim().takeIf { it.isNotBlank() }
}

private fun extractClaim(token: String?, vararg names: String): String? {
if (token.isNullOrBlank()) return null
val payload = runCatching { JwtUtils.decodePayload(token) }.getOrNull() ?: return null
return names.firstNotNullOfOrNull { name ->
payload[name]?.toString()?.takeIf { it.isNotBlank() }
}
}

private fun resolveClientIp(exchange: ServerWebExchange): String? {
val forwardedFor = exchange.request.headers.getFirst("X-Forwarded-For")
if (!forwardedFor.isNullOrBlank()) {
return forwardedFor.substringBefore(",").trim()
}
return exchange.request.headers.getFirst("X-Real-IP")?.takeIf { it.isNotBlank() }
?: exchange.request.remoteAddress?.address?.hostAddress
}

private fun truncateBody(body: String): String {
if (body.length <= maxPayloadSize) return body
return body.take(maxPayloadSize) + "...(truncated)"
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import io.swagger.v3.oas.annotations.responses.ApiResponse
import io.swagger.v3.oas.annotations.security.SecurityRequirement
import io.swagger.v3.oas.annotations.tags.Tag
import org.springframework.http.ResponseEntity
import org.springframework.http.server.reactive.ServerHttpRequest
import org.springframework.security.core.annotation.CurrentSecurityContext
import org.springframework.security.core.context.SecurityContext
import org.springframework.web.bind.annotation.PostMapping
Expand Down Expand Up @@ -48,7 +49,11 @@ Allowed values:
)
]
)
suspend fun requestGetToken(@RequestBody tokenRequest: PasswordFlowTokenRequest): ResponseEntity<TokenResponse> {
suspend fun requestGetToken(
@RequestBody tokenRequest: PasswordFlowTokenRequest,
request: ServerHttpRequest
Comment thread
fatemeh-i marked this conversation as resolved.
): ResponseEntity<TokenResponse> {
tokenRequest.ipAddress = resolveClientIp(request)
val tokenResponse = loginService.requestGetToken(tokenRequest)
return ResponseEntity.ok().body(tokenResponse)
}
Expand All @@ -69,7 +74,11 @@ Behavior: Completes password-flow login after OTP verification.""",
)
]
)
suspend fun confirmGetToken(@RequestBody tokenRequest: ConfirmPasswordFlowTokenRequest): ResponseEntity<TokenResponse> {
suspend fun confirmGetToken(
@RequestBody tokenRequest: ConfirmPasswordFlowTokenRequest,
request: ServerHttpRequest
): ResponseEntity<TokenResponse> {
tokenRequest.ipAddress = resolveClientIp(request)
val tokenResponse = loginService.confirmGetToken(tokenRequest)
return ResponseEntity.ok().body(tokenResponse)
}
Expand Down Expand Up @@ -137,8 +146,21 @@ Behavior: Issues a new access token from a valid refresh token.""",
)
]
)
suspend fun refreshToken(@RequestBody tokenRequest: RefreshTokenRequest): ResponseEntity<TokenResponse> {
suspend fun refreshToken(
@RequestBody tokenRequest: RefreshTokenRequest,
request: ServerHttpRequest
): ResponseEntity<TokenResponse> {
tokenRequest.ipAddress = resolveClientIp(request)
val tokenResponse = loginService.refreshToken(tokenRequest)
return ResponseEntity.ok().body(tokenResponse)
}

private fun resolveClientIp(request: ServerHttpRequest): String? {
val forwardedFor = request.headers.getFirst("X-Forwarded-For")
if (!forwardedFor.isNullOrBlank()) {
return forwardedFor.substringBefore(",").trim()
}
return request.headers.getFirst("X-Real-IP")?.takeIf { it.isNotBlank() }
?: request.remoteAddress?.address?.hostAddress
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import io.swagger.v3.oas.annotations.responses.ApiResponse
import io.swagger.v3.oas.annotations.tags.Tag
import jakarta.validation.Valid
import org.springframework.http.ResponseEntity
import org.springframework.http.server.reactive.ServerHttpRequest
import org.springframework.web.bind.annotation.PostMapping
import org.springframework.web.bind.annotation.RequestBody
import org.springframework.web.bind.annotation.RequestMapping
Expand Down Expand Up @@ -112,7 +113,11 @@ Behavior: Completes registration and returns login token data.""",
)
]
)
suspend fun confirmRegister(@RequestBody request: ConfirmRegisterRequest): ResponseEntity<Token> {
suspend fun confirmRegister(
@RequestBody request: ConfirmRegisterRequest,
@io.swagger.v3.oas.annotations.Parameter(hidden = true) serverRequest: ServerHttpRequest
): ResponseEntity<Token> {
request.ipAddress = resolveClientIp(serverRequest)
val loginToken = registerService.confirmRegister(request)
return ResponseEntity.ok(loginToken)
}
Expand Down Expand Up @@ -224,4 +229,13 @@ Response body: No response body.""",
forgetPasswordService.confirmForget(request)
return ResponseEntity.ok().build()
}

private fun resolveClientIp(request: ServerHttpRequest): String? {
val forwardedFor = request.headers.getFirst("X-Forwarded-For")
if (!forwardedFor.isNullOrBlank()) {
return forwardedFor.substringBefore(",").trim()
}
return request.headers.getFirst("X-Real-IP")?.takeIf { it.isNotBlank() }
?: request.remoteAddress?.address?.hostAddress
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -11,4 +11,5 @@ open class Device {
var pushToken: String? = null
var deviceUuid: String? = null
var buildNumber: Int? = null
var ipAddress: String? = null
}
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ data class LoginEvent(
val platform: Platform?,
val agent: String?,
val buildNumber: Int?,
val ipAddress: String?,
val sessionId: String,
val expireDate: LocalDateTime
) : AuthEvent()
Original file line number Diff line number Diff line change
Expand Up @@ -230,6 +230,7 @@ class LoginService(
platform = request.platform,
agent = request.agent,
buildNumber = request.buildNumber,
ipAddress = request.ipAddress,
sessionId = sessionState ?: "",
expireDate = LocalDateTime.now().plusSeconds(expiresIn.toLong())
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,7 @@ class RegisterService(
platform = request.platform,
agent = request.agent,
buildNumber = request.buildNumber,
ipAddress = request.ipAddress,
sessionId = sessionState ?: "",
expireDate = LocalDateTime.now().plusSeconds(expiresIn.toLong())
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ data class LoginEvent(
val platform: Platform?,
val agent: String?,
val buildNumber: Int?,
val ipAddress: String? = null,
val sessionId: String,
val expireDate: LocalDateTime
) : SessionEvent()

Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ data class Session(
val sessionState: String,
val userId: String,
val deviceId: Long,
val ipAddress: String? = null,
val status: SessionStatus = SessionStatus.ACTIVE,
val createDate: LocalDateTime? = LocalDateTime.now(),
val expireDate: LocalDateTime?
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ class UserSessionDeviceService(
sessionState = loginEvent.sessionId,
userId = loginEvent.uuid,
deviceId = device.id,
ipAddress = loginEvent.ipAddress,
expireDate = loginEvent.expireDate
)
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,11 +26,12 @@ class SessionManagerImpl(
private val logger = LoggerFactory.getLogger(SessionManagerImpl::class.java)

override suspend fun createOrUpdateSession(session: Session): Session? {
val newOrUpdatedSession = sessionRepository.findBySessionState(session.sessionState)
.awaitFirstOrNull()?.copy(
expireDate = session.expireDate,
status = SessionStatus.ACTIVE
) ?: session.toModel()
val existingSession = sessionRepository.findBySessionState(session.sessionState).awaitFirstOrNull()
val newOrUpdatedSession = existingSession?.copy(
expireDate = session.expireDate,
ipAddress = session.ipAddress ?: existingSession.ipAddress,
status = SessionStatus.ACTIVE
) ?: session.toModel()
sessionRepository.save(newOrUpdatedSession).awaitSingle()
return newOrUpdatedSession.toDto()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,9 @@ data class SessionModel(
@Column("session_state") val sessionState: String,
@Column("uuid") val userId: String,
@Column("device_id") val deviceId: Long,
@Column("ip_address") val ipAddress: String? = null,
@Column("status") val status: SessionStatus,
@Column("create_date") val createDate: LocalDateTime? = LocalDateTime.now(),
@Column("expire_date") val expireDate: LocalDateTime? = LocalDateTime.now(),
@Column @Version var version: Long? = null
)

Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ fun Session.toModel(): SessionModel {
sessionState = sessionState,
userId = userId,
deviceId = deviceId,
ipAddress = ipAddress,
status = status,
createDate = createDate ?: LocalDateTime.now(),
expireDate = expireDate
Expand All @@ -78,6 +79,7 @@ fun SessionModel.toDto(): Session {
sessionState = sessionState,
userId = userId,
deviceId = deviceId,
ipAddress = ipAddress,
status = status,
createDate = createDate,
expireDate = expireDate
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
ALTER TABLE sessions
ADD COLUMN IF NOT EXISTS ip_address VARCHAR(64);
Loading