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
12 changes: 7 additions & 5 deletions build.gradle
Original file line number Diff line number Diff line change
Expand Up @@ -21,18 +21,20 @@ dependencies {
implementation 'org.springframework.boot:spring-boot-starter-web'
implementation 'org.springframework.boot:spring-boot-starter-data-jpa'
implementation 'org.springframework.boot:spring-boot-starter-validation'
implementation 'org.springframework.boot:spring-boot-starter-security'
implementation 'org.springframework.security:spring-security-crypto'
implementation 'io.jsonwebtoken:jjwt-api:0.13.0'
compileOnly 'org.projectlombok:lombok'
annotationProcessor 'org.projectlombok:lombok'
runtimeOnly 'com.h2database:h2'
runtimeOnly 'io.jsonwebtoken:jjwt-impl:0.13.0'
runtimeOnly 'io.jsonwebtoken:jjwt-jackson:0.13.0'
testCompileOnly 'org.projectlombok:lombok'
testAnnotationProcessor 'org.projectlombok:lombok'
testImplementation 'org.springframework.boot:spring-boot-starter-test'
testRuntimeOnly 'org.junit.platform:junit-platform-launcher'
}
testCompileOnly 'org.projectlombok:lombok'
testAnnotationProcessor 'org.projectlombok:lombok'
testImplementation 'org.springframework.boot:spring-boot-starter-test'
testImplementation 'org.springframework.security:spring-security-test'
testRuntimeOnly 'org.junit.platform:junit-platform-launcher'
}

tasks.named('test') {
useJUnitPlatform()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -86,13 +86,22 @@ private void saveRefreshTokenSession(Long userId, RefreshToken refreshToken, Str
}

private void validateRefreshTokenSession(RefreshTokenSession session, RefreshTokenClaims claims, Instant now) {
if (!session.getUserId().equals(claims.userId())
|| !session.getExpiresAt().equals(claims.expiresAt())
|| !session.isAvailableAt(now)) {
if (!session.getUserId().equals(claims.userId()) || !session.getExpiresAt().equals(claims.expiresAt())) {
revokeTokenFamilyAndReject(session, now);
}
if (session.isRevoked()) {
revokeTokenFamilyAndReject(session, now);
}
if (session.isExpiredAt(now)) {
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}
}

private void revokeTokenFamilyAndReject(RefreshTokenSession session, Instant now) {
refreshTokenSessionRepository.revokeFamily(session.getFamilyId(), now);
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}

private void validatePassword(String rawPassword, String encodedPassword) {
if (!passwordEncoder.matches(rawPassword, encodedPassword)) {
throw new AuthDomainException(AuthErrorCode.INVALID_CREDENTIALS);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,5 +6,10 @@ public interface AccessTokenProvider {

String createAccessToken(AuthenticatedUserInfo user);

AccessTokenClaims parseAccessToken(String accessToken);

Long getAccessTokenExpirationSeconds();

record AccessTokenClaims(Long userId, String role) {
}
}
Original file line number Diff line number Diff line change
@@ -1,11 +1,14 @@
package com.promsearch.auth.application.port.out;

import com.promsearch.auth.domain.RefreshTokenSession;
import java.time.Instant;
import java.util.Optional;

public interface RefreshTokenSessionRepository {

RefreshTokenSession save(RefreshTokenSession session);

Optional<RefreshTokenSession> findByTokenHashForUpdate(String tokenHash);

void revokeFamily(String familyId, Instant revokedAt);
}
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,14 @@ public boolean isAvailableAt(Instant now) {
return revokedAt == null && now.isBefore(expiresAt);
}

public boolean isRevoked() {
return revokedAt != null;
}

public boolean isExpiredAt(Instant now) {
return !now.isBefore(expiresAt);
}

public void revoke(Instant now) {
if (revokedAt == null) {
revokedAt = now;
Expand Down
Original file line number Diff line number Diff line change
@@ -1,11 +1,26 @@
package com.promsearch.auth.infrastructure.jwt;

import jakarta.validation.constraints.NotBlank;
import jakarta.validation.constraints.NotNull;
import jakarta.validation.constraints.Positive;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.validation.annotation.Validated;

@Validated
@ConfigurationProperties(prefix = "auth.jwt")
public record JwtProperties(
String secret,
@NotBlank(message = "auth.jwt.access-secret is required")
String accessSecret,

@NotBlank(message = "auth.jwt.refresh-secret is required")
String refreshSecret,

@NotNull(message = "auth.jwt.access-token-expiration-seconds is required")
@Positive(message = "auth.jwt.access-token-expiration-seconds must be positive")
Long accessTokenExpirationSeconds,

@NotNull(message = "auth.jwt.refresh-token-expiration-seconds is required")
@Positive(message = "auth.jwt.refresh-token-expiration-seconds must be positive")
Long refreshTokenExpirationSeconds
) {
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,12 @@

import com.promsearch.auth.application.AuthenticatedUserInfo;
import com.promsearch.auth.application.port.out.AccessTokenProvider;
import com.promsearch.auth.application.port.out.AccessTokenProvider.AccessTokenClaims;
import com.promsearch.auth.application.port.out.RefreshTokenProvider;
import com.promsearch.auth.domain.exception.AuthDomainException;
import com.promsearch.auth.domain.exception.AuthErrorCode;
import io.jsonwebtoken.Claims;
import io.jsonwebtoken.ExpiredJwtException;
import io.jsonwebtoken.JwtException;
import io.jsonwebtoken.JwtParser;
import io.jsonwebtoken.Jwts;
Expand All @@ -17,17 +19,23 @@
import java.util.UUID;
import javax.crypto.SecretKey;
import org.springframework.beans.factory.annotation.Autowired;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.stereotype.Component;

@Component
public class JwtTokenProvider implements AccessTokenProvider, RefreshTokenProvider {

private static final Logger log = LoggerFactory.getLogger(JwtTokenProvider.class);
private static final String REFRESH_TOKEN_TYPE = "refresh";
private static final int MIN_HMAC_KEY_BYTES = 32;

private final JwtProperties jwtProperties;
private final Clock clock;
private final SecretKey secretKey;
private final JwtParser jwtParser;
private final SecretKey accessSecretKey;
private final SecretKey refreshSecretKey;
private final JwtParser accessJwtParser;
private final JwtParser refreshJwtParser;

@Autowired
public JwtTokenProvider(JwtProperties jwtProperties) {
Expand All @@ -37,9 +45,14 @@ public JwtTokenProvider(JwtProperties jwtProperties) {
JwtTokenProvider(JwtProperties jwtProperties, Clock clock) {
this.jwtProperties = jwtProperties;
this.clock = clock;
this.secretKey = Keys.hmacShaKeyFor(jwtProperties.secret().getBytes(StandardCharsets.UTF_8));
this.jwtParser = Jwts.parser()
.verifyWith(secretKey)
this.accessSecretKey = createSecretKey("auth.jwt.access-secret", jwtProperties.accessSecret());
this.refreshSecretKey = createSecretKey("auth.jwt.refresh-secret", jwtProperties.refreshSecret());
this.accessJwtParser = Jwts.parser()
.verifyWith(accessSecretKey)
.clock(() -> Date.from(Instant.now(clock)))
.build();
this.refreshJwtParser = Jwts.parser()
.verifyWith(refreshSecretKey)
.clock(() -> Date.from(Instant.now(clock)))
.build();
}
Expand All @@ -55,10 +68,24 @@ public String createAccessToken(AuthenticatedUserInfo user) {
.claim("role", user.role())
.issuedAt(Date.from(now))
.expiration(Date.from(expiresAt))
.signWith(secretKey, Jwts.SIG.HS256)
.signWith(accessSecretKey, Jwts.SIG.HS256)
.compact();
}

@Override
public AccessTokenClaims parseAccessToken(String accessToken) {
try {
Claims claims = accessJwtParser.parseSignedClaims(accessToken).getPayload();
return new AccessTokenClaims(getUserId(claims), getRole(claims));
} catch (ExpiredJwtException e) {
log.warn("Expired access token rejected.");
throw new AuthDomainException(AuthErrorCode.ACCESS_TOKEN_EXPIRED);
} catch (JwtException | IllegalArgumentException e) {
log.warn("Invalid access token rejected. reason={}", e.getClass().getSimpleName());
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}
}

@Override
public Long getAccessTokenExpirationSeconds() {
return jwtProperties.accessTokenExpirationSeconds();
Expand All @@ -76,7 +103,7 @@ public RefreshToken createRefreshToken(AuthenticatedUserInfo user) {
.id(UUID.randomUUID().toString())
.issuedAt(Date.from(now))
.expiration(Date.from(expiresAt))
.signWith(secretKey, Jwts.SIG.HS256)
.signWith(refreshSecretKey, Jwts.SIG.HS256)
.compact();

return new RefreshToken(token, expiresAt);
Expand All @@ -85,14 +112,15 @@ public RefreshToken createRefreshToken(AuthenticatedUserInfo user) {
@Override
public RefreshTokenClaims parse(String refreshToken) {
try {
Claims claims = jwtParser.parseSignedClaims(refreshToken).getPayload();
Claims claims = refreshJwtParser.parseSignedClaims(refreshToken).getPayload();
validateRefreshTokenClaims(claims);
return new RefreshTokenClaims(
getUserId(claims),
claims.getId(),
claims.getExpiration().toInstant()
);
} catch (JwtException | IllegalArgumentException e) {
log.warn("Invalid refresh token rejected. reason={}", e.getClass().getSimpleName());
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}
}
Expand Down Expand Up @@ -120,4 +148,23 @@ private Long getUserId(Claims claims) {
}
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}

private String getRole(Claims claims) {
String role = claims.get("role", String.class);
if (role == null || role.isBlank()) {
throw new AuthDomainException(AuthErrorCode.INVALID_TOKEN);
}
return role;
}

private SecretKey createSecretKey(String propertyName, String secret) {
if (secret == null || secret.isBlank()) {
throw new IllegalStateException(propertyName + " is required.");
}
byte[] keyBytes = secret.getBytes(StandardCharsets.UTF_8);
if (keyBytes.length < MIN_HMAC_KEY_BYTES) {
throw new IllegalStateException(propertyName + " must be at least 32 bytes for HS256.");
}
return Keys.hmacShaKeyFor(keyBytes);
}
}
Original file line number Diff line number Diff line change
@@ -1,14 +1,25 @@
package com.promsearch.auth.infrastructure.persistence;

import jakarta.persistence.LockModeType;
import java.time.Instant;
import java.util.Optional;
import org.springframework.data.jpa.repository.JpaRepository;
import org.springframework.data.jpa.repository.Lock;
import org.springframework.data.jpa.repository.Modifying;
import org.springframework.data.jpa.repository.Query;
import org.springframework.data.repository.query.Param;

public interface RefreshTokenSessionJpaRepository extends JpaRepository<RefreshTokenSessionJpaEntity, Long> {
@Lock(LockModeType.PESSIMISTIC_WRITE)
@Query("select session from RefreshTokenSessionJpaEntity session where session.tokenHash = :tokenHash")
Optional<RefreshTokenSessionJpaEntity> findByTokenHashForUpdate(@Param("tokenHash") String tokenHash);

@Modifying(clearAutomatically = true, flushAutomatically = true)
@Query("""
update RefreshTokenSessionJpaEntity session
set session.revokedAt = :revokedAt
where session.familyId = :familyId
and session.revokedAt is null
""")
void revokeFamily(@Param("familyId") String familyId, @Param("revokedAt") Instant revokedAt);
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import com.promsearch.auth.application.port.out.RefreshTokenSessionRepository;
import com.promsearch.auth.domain.RefreshTokenSession;
import java.time.Instant;
import java.util.Optional;
import lombok.RequiredArgsConstructor;
import org.springframework.stereotype.Repository;
Expand All @@ -25,4 +26,9 @@ public RefreshTokenSession save(RefreshTokenSession session) {
public Optional<RefreshTokenSession> findByTokenHashForUpdate(String tokenHash) {
return repository.findByTokenHashForUpdate(tokenHash).map(RefreshTokenSessionJpaEntity::toDomain);
}

@Override
public void revokeFamily(String familyId, Instant revokedAt) {
repository.revokeFamily(familyId, revokedAt);
}
}
Loading
Loading