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 @@ -4,28 +4,39 @@ package com.openai.core

import com.openai.errors.OpenAIException
import java.lang.reflect.InvocationTargetException
import java.util.concurrent.atomic.AtomicBoolean

/**
* Closes [closeable] when [observed] becomes only phantom reachable.
*
* This is a wrapper around a Java 9+ [java.lang.ref.Cleaner], or a no-op in older Java versions.
* The returned handle performs the same cleanup explicitly, at most once across both paths.
*/
@JvmSynthetic
internal fun closeWhenPhantomReachable(observed: Any, closeable: AutoCloseable) {
internal fun closeWhenPhantomReachable(observed: Any, closeable: AutoCloseable): AutoCloseable {
check(observed !== closeable) {
"`observed` cannot be the same object as `closeable` because it would never become phantom reachable"
}
closeWhenPhantomReachable(observed, closeable::close)
return closeWhenPhantomReachable(observed, closeable::close)
}

/**
* Calls [close] when [observed] becomes only phantom reachable.
*
* This is a wrapper around a Java 9+ [java.lang.ref.Cleaner], or a no-op in older Java versions.
* Calling the returned handle performs the same cleanup explicitly, at most once across both paths.
*/
@JvmSynthetic
internal fun closeWhenPhantomReachable(observed: Any, close: () -> Unit) {
closeWhenPhantomReachable?.let { it(observed, close) }
internal fun closeWhenPhantomReachable(observed: Any, close: () -> Unit): AutoCloseable {
val closed = AtomicBoolean(false)
val closeOnce = {
if (closed.compareAndSet(false, true)) {
close()
}
}

closeWhenPhantomReachable?.let { it(observed, closeOnce) }
return AutoCloseable { closeOnce() }
}

private val closeWhenPhantomReachable: ((Any, () -> Unit) -> Unit)? by lazy {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,14 +10,12 @@ import java.util.concurrent.CompletableFuture
*/
internal class PhantomReachableSleeper(private val sleeper: Sleeper) : Sleeper {

init {
closeWhenPhantomReachable(this, sleeper)
}
private val closeHandle = closeWhenPhantomReachable(this, sleeper)

override fun sleep(duration: Duration) = sleeper.sleep(duration)

override fun sleepAsync(duration: Duration): CompletableFuture<Void> =
sleeper.sleepAsync(duration)

override fun close() = sleeper.close()
override fun close() = closeHandle.close()
}
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,8 @@ internal class PhantomReachableClosingAsyncStreamResponse<T>(
*/
private val reachabilityTracker = Object()

init {
private val closeHandle =
closeWhenPhantomReachable(reachabilityTracker, asyncStreamResponse::close)
}

override fun subscribe(handler: Handler<T>): AsyncStreamResponse<T> = apply {
asyncStreamResponse.subscribe(TrackedHandler(handler, reachabilityTracker))
Expand All @@ -37,7 +36,7 @@ internal class PhantomReachableClosingAsyncStreamResponse<T>(
override fun onCompleteFuture(): CompletableFuture<Void?> =
asyncStreamResponse.onCompleteFuture()

override fun close() = asyncStreamResponse.close()
override fun close() = closeHandle.close()
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,9 +10,7 @@ import java.util.concurrent.CompletableFuture
* This class ensures the `HttpClient` is closed even if the user forgets to close it.
*/
internal class PhantomReachableClosingHttpClient(private val httpClient: HttpClient) : HttpClient {
init {
closeWhenPhantomReachable(this, httpClient)
}
private val closeHandle = closeWhenPhantomReachable(this, httpClient)

override fun execute(request: HttpRequest, requestOptions: RequestOptions): HttpResponse =
httpClient.execute(request, requestOptions)
Expand All @@ -22,5 +20,5 @@ internal class PhantomReachableClosingHttpClient(private val httpClient: HttpCli
requestOptions: RequestOptions,
): CompletableFuture<HttpResponse> = httpClient.executeAsync(request, requestOptions)

override fun close() = httpClient.close()
override fun close() = closeHandle.close()
}
Original file line number Diff line number Diff line change
Expand Up @@ -12,15 +12,13 @@ import java.util.concurrent.CompletableFuture
internal class PhantomReachableClosingHttpRequestAuthenticator(
private val authenticator: HttpRequestAuthenticator
) : HttpRequestAuthenticator {
init {
closeWhenPhantomReachable(this, authenticator)
}
private val closeHandle = closeWhenPhantomReachable(this, authenticator)

override fun authenticate(request: HttpRequest): HttpRequest =
authenticator.authenticate(request)

override fun authenticateAsync(request: HttpRequest): CompletableFuture<HttpRequest> =
authenticator.authenticateAsync(request)

override fun close() = authenticator.close()
override fun close() = closeHandle.close()
}
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,9 @@ import java.util.stream.Stream
internal class PhantomReachableClosingStreamResponse<T>(
private val streamResponse: StreamResponse<T>
) : StreamResponse<T> {
init {
closeWhenPhantomReachable(this, streamResponse)
}
private val closeHandle = closeWhenPhantomReachable(this, streamResponse)

override fun stream(): Stream<T> = streamResponse.stream()

override fun close() = streamResponse.close()
override fun close() = closeHandle.close()
}
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
package com.openai.core

import java.util.concurrent.atomic.AtomicInteger
import org.assertj.core.api.Assertions.assertThat
import org.junit.jupiter.api.Test

Expand All @@ -24,4 +25,18 @@ internal class PhantomReachableTest {

assertThat(closed).isTrue()
}

@Test
fun closeWhenPhantomReachable_explicitHandleClosesAtMostOnce() {
val closeCount = AtomicInteger()
val observed = Any()
val handle = closeWhenPhantomReachable(observed) { closeCount.incrementAndGet() }

handle.close()
handle.close()

assertThat(closeCount.get()).isEqualTo(1)
// Keep the observed object strongly reachable until after both explicit closes.
assertThat(observed).isNotNull()
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
package com.openai.core.http

import org.junit.jupiter.api.Test
import org.mockito.kotlin.mock
import org.mockito.kotlin.times
import org.mockito.kotlin.verify

internal class PhantomReachableClosingHttpClientTest {

@Test
fun close_closesDelegateAtMostOnce() {
val delegate = mock<HttpClient>()
val client = PhantomReachableClosingHttpClient(delegate)

client.close()
client.close()

verify(delegate, times(1)).close()
}
}