diff --git a/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBodies.kt b/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBodies.kt index 2c672bb8e..86a2c76b6 100644 --- a/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBodies.kt +++ b/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBodies.kt @@ -11,6 +11,8 @@ import com.openai.errors.OpenAIInvalidDataException import java.io.ByteArrayInputStream import java.io.InputStream import java.io.OutputStream +import java.io.SequenceInputStream +import java.util.Collections import java.util.UUID import kotlin.jvm.optionals.getOrNull @@ -21,6 +23,8 @@ internal inline fun json(jsonMapper: JsonMapper, value: T): HttpRequ override fun writeTo(outputStream: OutputStream) = outputStream.write(bytes) + override fun content(): InputStream = bytes.inputStream() + override fun contentType(): String = "application/json" override fun contentLength(): Long = bytes.size.toLong() @@ -60,6 +64,8 @@ internal fun multipartFormData( outputStream.write(byteArray) } + override fun content(): InputStream = byteArray.inputStream() + override fun contentType(): String = field.contentType override fun contentLength(): Long = byteArray.size.toLong() @@ -75,6 +81,8 @@ internal fun multipartFormData( bytes.copyTo(outputStream) } + override fun content(): InputStream = bytes + override fun contentType(): String = field.contentType override fun contentLength(): Long = -1L @@ -147,6 +155,36 @@ private constructor(private val boundary: String, private val parts: List) outputStream.write(CRLF) } + // This must remain in sync with `writeTo`. + override fun content(): InputStream { + val streams = mutableListOf() + + parts.forEach { part -> + streams.add(DASHDASH.inputStream()) + streams.add(boundaryBytes.inputStream()) + streams.add(CRLF.inputStream()) + + streams.add(CONTENT_DISPOSITION.inputStream()) + streams.add(part.contentDisposition.toByteArray().inputStream()) + streams.add(CRLF.inputStream()) + + streams.add(CONTENT_TYPE.inputStream()) + streams.add(part.contentType.toByteArray().inputStream()) + streams.add(CRLF.inputStream()) + + streams.add(CRLF.inputStream()) + streams.add(part.body.content()) + streams.add(CRLF.inputStream()) + } + + streams.add(DASHDASH.inputStream()) + streams.add(boundaryBytes.inputStream()) + streams.add(DASHDASH.inputStream()) + streams.add(CRLF.inputStream()) + + return SequenceInputStream(Collections.enumeration(streams)) + } + override fun contentType(): String = contentType // This must remain in sync with `writeTo`. diff --git a/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBody.kt b/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBody.kt index 55bc68a30..b466e67d9 100644 --- a/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBody.kt +++ b/openai-java-core/src/main/kotlin/com/openai/core/http/HttpRequestBody.kt @@ -1,5 +1,7 @@ package com.openai.core.http +import java.io.ByteArrayOutputStream +import java.io.InputStream import java.io.OutputStream import java.lang.AutoCloseable @@ -7,6 +9,19 @@ interface HttpRequestBody : AutoCloseable { fun writeTo(outputStream: OutputStream) + /** + * Returns the request body content as an input stream. + * + * The default implementation buffers the bytes produced by [writeTo] so existing third-party + * implementations remain compatible. Implementations backed by an existing byte array or stream + * should override this method to avoid buffering. + */ + fun content(): InputStream { + val outputStream = ByteArrayOutputStream() + writeTo(outputStream) + return outputStream.toByteArray().inputStream() + } + fun contentType(): String? fun contentLength(): Long diff --git a/openai-java-core/src/test/java/com/openai/core/http/HttpRequestBodyJavaTest.java b/openai-java-core/src/test/java/com/openai/core/http/HttpRequestBodyJavaTest.java new file mode 100644 index 000000000..5c35183b0 --- /dev/null +++ b/openai-java-core/src/test/java/com/openai/core/http/HttpRequestBodyJavaTest.java @@ -0,0 +1,54 @@ +package com.openai.core.http; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.nio.charset.StandardCharsets; +import org.junit.jupiter.api.Test; + +final class HttpRequestBodyJavaTest { + + @Test + void contentIsJavaDefaultMethod() throws Exception { + assertThat(HttpRequestBody.class.getMethod("content").isDefault()).isTrue(); + + HttpRequestBody body = + new HttpRequestBody() { + @Override + public void writeTo(OutputStream outputStream) { + try { + outputStream.write("body".getBytes(StandardCharsets.UTF_8)); + } catch (IOException e) { + throw new RuntimeException(e); + } + } + + @Override + public String contentType() { + return "text/plain"; + } + + @Override + public long contentLength() { + return 4L; + } + + @Override + public boolean repeatable() { + return true; + } + + @Override + public void close() {} + }; + + try (InputStream content = body.content()) { + byte[] bytes = new byte[4]; + assertThat(content.read(bytes)).isEqualTo(bytes.length); + assertThat(bytes).isEqualTo("body".getBytes(StandardCharsets.UTF_8)); + assertThat(content.read()).isEqualTo(-1); + } + } +} diff --git a/openai-java-core/src/test/kotlin/com/openai/core/http/HttpRequestBodyContentTest.kt b/openai-java-core/src/test/kotlin/com/openai/core/http/HttpRequestBodyContentTest.kt new file mode 100644 index 000000000..38ee82c68 --- /dev/null +++ b/openai-java-core/src/test/kotlin/com/openai/core/http/HttpRequestBodyContentTest.kt @@ -0,0 +1,78 @@ +package com.openai.core.http + +import com.openai.core.MultipartField +import com.openai.core.jsonMapper +import java.io.ByteArrayOutputStream +import java.io.OutputStream +import org.assertj.core.api.Assertions.assertThat +import org.junit.jupiter.api.Test + +internal class HttpRequestBodyContentTest { + + @Test + fun content_defaultsToWriteTo() { + val body = + object : HttpRequestBody { + override fun writeTo(outputStream: OutputStream) { + outputStream.write("body".toByteArray()) + } + + override fun contentType(): String = "text/plain" + + override fun contentLength(): Long = 4L + + override fun repeatable(): Boolean = true + + override fun close() {} + } + + body.content().use { content -> assertThat(content.readBytes()).isEqualTo("body".toByteArray()) } + } + + @Test + fun multipartContent_matchesWriteTo() { + val body = + multipartFormData( + jsonMapper(), + mapOf( + "field" to + MultipartField.builder() + .value("value") + .contentType("text/plain") + .build(), + "binary" to + MultipartField.builder() + .value("abc".toByteArray()) + .contentType("application/octet-stream") + .build(), + ), + ) + + val output = ByteArrayOutputStream() + body.writeTo(output) + + body.content().use { content -> assertThat(content.readBytes()).isEqualTo(output.toByteArray()) } + } + + @Test + fun multipartContent_streamsInputStreamParts() { + val body = + multipartFormData( + jsonMapper(), + mapOf( + "data" to + MultipartField.builder() + .value("stream content".byteInputStream().buffered()) + .contentType("application/octet-stream") + .build() + ), + ) + + val content = body.content().use { it.readBytes().toString(Charsets.UTF_8) } + + assertThat(body.repeatable()).isFalse() + assertThat(content).contains("Content-Disposition: form-data; name=\"data\"") + assertThat(content).contains("Content-Type: application/octet-stream") + assertThat(content).contains("stream content") + } +}