diff --git a/openai-java-core/src/main/kotlin/com/openai/core/handlers/ErrorHandler.kt b/openai-java-core/src/main/kotlin/com/openai/core/handlers/ErrorHandler.kt index a235b051a..83070517a 100644 --- a/openai-java-core/src/main/kotlin/com/openai/core/handlers/ErrorHandler.kt +++ b/openai-java-core/src/main/kotlin/com/openai/core/handlers/ErrorHandler.kt @@ -42,50 +42,35 @@ internal fun errorHandler( errorBodyHandler: Handler> ): Handler = object : Handler { - override fun handle(response: HttpResponse): HttpResponse = - when (val statusCode = response.statusCode()) { - in 200..299 -> response - 400 -> - throw BadRequestException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - 401 -> - throw UnauthorizedException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - 403 -> - throw PermissionDeniedException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - 404 -> - throw NotFoundException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - 422 -> - throw UnprocessableEntityException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - 429 -> - throw RateLimitException.builder() - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - in 500..599 -> - throw InternalServerException.builder() - .statusCode(statusCode) - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() - else -> - throw UnexpectedStatusCodeException.builder() - .statusCode(statusCode) - .headers(response.headers()) - .error(errorBodyHandler.handle(response)) - .build() + override fun handle(response: HttpResponse): HttpResponse { + val statusCode = response.statusCode() + if (statusCode in 200..299) { + return response } + + return response.use { + val headers = response.headers() + val error = errorBodyHandler.handle(response) + throw when (statusCode) { + 400 -> BadRequestException.builder().headers(headers).error(error).build() + 401 -> UnauthorizedException.builder().headers(headers).error(error).build() + 403 -> PermissionDeniedException.builder().headers(headers).error(error).build() + 404 -> NotFoundException.builder().headers(headers).error(error).build() + 422 -> UnprocessableEntityException.builder().headers(headers).error(error).build() + 429 -> RateLimitException.builder().headers(headers).error(error).build() + in 500..599 -> + InternalServerException.builder() + .statusCode(statusCode) + .headers(headers) + .error(error) + .build() + else -> + UnexpectedStatusCodeException.builder() + .statusCode(statusCode) + .headers(headers) + .error(error) + .build() + } + } + } } diff --git a/openai-java-core/src/test/kotlin/com/openai/core/handlers/ErrorHandlerResponseCloseTest.kt b/openai-java-core/src/test/kotlin/com/openai/core/handlers/ErrorHandlerResponseCloseTest.kt new file mode 100644 index 000000000..0ad07615f --- /dev/null +++ b/openai-java-core/src/test/kotlin/com/openai/core/handlers/ErrorHandlerResponseCloseTest.kt @@ -0,0 +1,84 @@ +package com.openai.core.handlers + +import com.openai.core.JsonField +import com.openai.core.JsonMissing +import com.openai.core.http.Headers +import com.openai.core.http.HttpResponse +import com.openai.errors.OpenAIException +import com.openai.models.ErrorObject +import java.io.ByteArrayInputStream +import java.io.InputStream +import org.assertj.core.api.Assertions.assertThat +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.assertThrows +import org.junit.jupiter.params.ParameterizedTest +import org.junit.jupiter.params.provider.ValueSource + +internal class ErrorHandlerResponseCloseTest { + + @ParameterizedTest + @ValueSource(ints = [400, 401, 403, 404, 418, 422, 429, 500]) + fun closesEveryNonSuccessResponse(statusCode: Int) { + val response = RecordingResponse(statusCode) + val handler = + errorHandler( + object : HttpResponse.Handler> { + override fun handle(response: HttpResponse): JsonField = + JsonMissing.of() + } + ) + + assertThrows { handler.handle(response) } + + assertThat(response.closed).isTrue() + } + + @Test + fun closesResponseWhenErrorBodyHandlerThrows() { + val failure = IllegalStateException("error body failed") + val response = RecordingResponse(400) + val handler = + errorHandler( + object : HttpResponse.Handler> { + override fun handle(response: HttpResponse): JsonField { + throw failure + } + } + ) + + val thrown = assertThrows { handler.handle(response) } + + assertThat(thrown).isSameAs(failure) + assertThat(response.closed).isTrue() + } + + @Test + fun leavesSuccessfulResponseOpenForCaller() { + val response = RecordingResponse(200) + val handler = + errorHandler( + object : HttpResponse.Handler> { + override fun handle(response: HttpResponse): JsonField = + error("error body handler must not run for successful responses") + } + ) + + assertThat(handler.handle(response)).isSameAs(response) + assertThat(response.closed).isFalse() + } + + private class RecordingResponse(private val statusCode: Int) : HttpResponse { + var closed = false + private set + + override fun statusCode(): Int = statusCode + + override fun headers(): Headers = Headers.builder().build() + + override fun body(): InputStream = ByteArrayInputStream(ByteArray(0)) + + override fun close() { + closed = true + } + } +}