Skip to content
Open
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
Expand Up @@ -42,50 +42,35 @@ internal fun errorHandler(
errorBodyHandler: Handler<JsonField<ErrorObject>>
): Handler<HttpResponse> =
object : Handler<HttpResponse> {
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()
}
}
}
}
Original file line number Diff line number Diff line change
@@ -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<JsonField<ErrorObject>> {
override fun handle(response: HttpResponse): JsonField<ErrorObject> =
JsonMissing.of()
}
)

assertThrows<OpenAIException> { 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<JsonField<ErrorObject>> {
override fun handle(response: HttpResponse): JsonField<ErrorObject> {
throw failure
}
}
)

val thrown = assertThrows<IllegalStateException> { handler.handle(response) }

assertThat(thrown).isSameAs(failure)
assertThat(response.closed).isTrue()
}

@Test
fun leavesSuccessfulResponseOpenForCaller() {
val response = RecordingResponse(200)
val handler =
errorHandler(
object : HttpResponse.Handler<JsonField<ErrorObject>> {
override fun handle(response: HttpResponse): JsonField<ErrorObject> =
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
}
}
}