From 51063446f4f17f7452afdcf3277557cd5a5ea529 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:02:20 +0100 Subject: [PATCH 1/2] fix: clear async workload refresh state on sync failures --- .../com/openai/auth/WorkloadIdentityAuth.kt | 25 +++++++++++-------- 1 file changed, 15 insertions(+), 10 deletions(-) diff --git a/openai-java-core/src/main/kotlin/com/openai/auth/WorkloadIdentityAuth.kt b/openai-java-core/src/main/kotlin/com/openai/auth/WorkloadIdentityAuth.kt index 51d97e950..9b8ceffc1 100644 --- a/openai-java-core/src/main/kotlin/com/openai/auth/WorkloadIdentityAuth.kt +++ b/openai-java-core/src/main/kotlin/com/openai/auth/WorkloadIdentityAuth.kt @@ -155,7 +155,7 @@ internal class WorkloadIdentityAuth( return when (action) { is TokenAction.ReturnCached -> CompletableFuture.completedFuture(action.token) is TokenAction.BackgroundRefresh -> { - performRefreshAndComplete(action.future) + startAsyncRefresh(action.future) CompletableFuture.completedFuture(action.token) } is TokenAction.WaitForRefresh -> @@ -165,20 +165,25 @@ internal class WorkloadIdentityAuth( is TokenRefreshResult.Failure -> throw result.error } } - is TokenAction.ForegroundRefresh -> { - val refresh = refreshTokenAsync() - refresh.whenComplete { token, error -> - finishRefresh(action.future, token, unwrapCompletionException(error)) - } - refresh - } + is TokenAction.ForegroundRefresh -> startAsyncRefresh(action.future) } } - private fun performRefreshAndComplete(future: CompletableFuture) { - refreshTokenAsync().whenComplete { token, error -> + private fun startAsyncRefresh( + future: CompletableFuture + ): CompletableFuture { + val refresh = + try { + refreshTokenAsync() + } catch (error: Throwable) { + finishRefresh(future, null, error) + return CompletableFuture().also { it.completeExceptionally(error) } + } + + refresh.whenComplete { token, error -> finishRefresh(future, token, unwrapCompletionException(error)) } + return refresh } private fun finishRefresh( From 63098656b0f49541d6f7d287f91a0e016473c9a0 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Tue, 18 Aug 2026 14:03:06 +0100 Subject: [PATCH 2/2] test: cover synchronous async-provider failures --- ...ntityAuthSynchronousProviderFailureTest.kt | 121 ++++++++++++++++++ 1 file changed, 121 insertions(+) create mode 100644 openai-java-core/src/test/kotlin/com/openai/auth/WorkloadIdentityAuthSynchronousProviderFailureTest.kt diff --git a/openai-java-core/src/test/kotlin/com/openai/auth/WorkloadIdentityAuthSynchronousProviderFailureTest.kt b/openai-java-core/src/test/kotlin/com/openai/auth/WorkloadIdentityAuthSynchronousProviderFailureTest.kt new file mode 100644 index 000000000..ff2b59fca --- /dev/null +++ b/openai-java-core/src/test/kotlin/com/openai/auth/WorkloadIdentityAuthSynchronousProviderFailureTest.kt @@ -0,0 +1,121 @@ +package com.openai.auth + +import com.fasterxml.jackson.databind.json.JsonMapper +import com.openai.core.http.HttpClient +import com.openai.core.http.HttpRequest +import com.openai.core.http.HttpResponse +import java.io.ByteArrayInputStream +import java.util.concurrent.CompletableFuture +import java.util.concurrent.CompletionException +import org.assertj.core.api.Assertions.assertThat +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.assertThrows +import org.mockito.kotlin.any +import org.mockito.kotlin.mock +import org.mockito.kotlin.times +import org.mockito.kotlin.verify +import org.mockito.kotlin.verifyNoInteractions +import org.mockito.kotlin.whenever + +internal class WorkloadIdentityAuthSynchronousProviderFailureTest { + + @Test + fun getTokenAsync_clearsForegroundRefreshAfterSynchronousProviderFailure() { + val failure = IllegalStateException("provider failed") + var providerCalls = 0 + val provider = + object : SubjectTokenProvider { + override fun tokenType() = SubjectTokenType.JWT + + override fun getToken(httpClient: HttpClient, jsonMapper: JsonMapper): String = + error("not used") + + override fun getTokenAsync( + httpClient: HttpClient, + jsonMapper: JsonMapper, + ): CompletableFuture { + providerCalls++ + throw failure + } + } + val httpClient = mock() + val auth = createAuth(provider, httpClient) + + val firstFailure = assertThrows { auth.getTokenAsync().join() } + val secondFailure = assertThrows { auth.getTokenAsync().join() } + + assertThat(firstFailure.cause).isSameAs(failure) + assertThat(secondFailure.cause).isSameAs(failure) + assertThat(providerCalls).isEqualTo(2) + verifyNoInteractions(httpClient) + } + + @Test + fun getTokenAsync_keepsCachedTokenAndAllowsAnotherBackgroundRefreshAfterSynchronousFailure() { + val subjectToken = "subject-token" + val accessToken = "access-token" + val failure = IllegalStateException("provider failed") + var asyncProviderCalls = 0 + val provider = + object : SubjectTokenProvider { + override fun tokenType() = SubjectTokenType.JWT + + override fun getToken(httpClient: HttpClient, jsonMapper: JsonMapper): String = + subjectToken + + override fun getTokenAsync( + httpClient: HttpClient, + jsonMapper: JsonMapper, + ): CompletableFuture { + asyncProviderCalls++ + throw failure + } + } + val httpClient = mock() + val response = + mockResponse( + 200, + """ + { + "access_token": "$accessToken", + "issued_token_type": "urn:ietf:params:oauth:token-type:access_token", + "token_type": "Bearer", + "expires_in": 60 + } + """ + .trimIndent(), + ) + whenever(httpClient.execute(any())).thenReturn(response) + val auth = createAuth(provider, httpClient) + + assertThat(auth.getToken()).isEqualTo(accessToken) + assertThat(auth.getTokenAsync().join()).isEqualTo(accessToken) + assertThat(auth.getTokenAsync().join()).isEqualTo(accessToken) + + assertThat(asyncProviderCalls).isEqualTo(2) + verify(httpClient, times(1)).execute(any()) + } + + private fun createAuth( + provider: SubjectTokenProvider, + httpClient: HttpClient, + ): WorkloadIdentityAuth = + WorkloadIdentityAuth( + config = + WorkloadIdentity.builder() + .clientId("client-id") + .identityProviderId("provider-id") + .serviceAccountId("service-account-id") + .provider(provider) + .build(), + httpClient = httpClient, + jsonMapper = JsonMapper(), + ) + + private fun mockResponse(statusCode: Int, body: String): HttpResponse { + val response = mock() + whenever(response.statusCode()).thenReturn(statusCode) + whenever(response.body()).thenAnswer { ByteArrayInputStream(body.toByteArray()) } + return response + } +}