This is an automated email from the ASF dual-hosted git repository. voidmatcha pushed a commit to branch temp/seung-00/assistant-transport-base in repository https://gitbox.apache.org/repos/asf/zeppelin.git
commit 40ae2be9bf6284184064aac90ca130ae3441fc15 Author: YONGJAE LEE <[email protected]> AuthorDate: Thu Oct 8 19:28:40 2026 +0900 Cancel assistant HTTP requests during shutdown --- zeppelin-distribution/src/bin_license/LICENSE | 3 ++ zeppelin-server/pom.xml | 12 ++++++ .../service/assistant/OpenAiChatModel.java | 46 ++++++++++++++++++---- .../assistant/OpenAiChatModelLifecycleTest.java | 46 ++++++++++++++++++++++ 4 files changed, 100 insertions(+), 7 deletions(-) diff --git a/zeppelin-distribution/src/bin_license/LICENSE b/zeppelin-distribution/src/bin_license/LICENSE index 3c6939d403..4a0382fb00 100644 --- a/zeppelin-distribution/src/bin_license/LICENSE +++ b/zeppelin-distribution/src/bin_license/LICENSE @@ -131,6 +131,9 @@ The following components are provided under Apache License. (Apache 2.0) gcsio.jar (com.google.cloud.bigdataoss:gcsio:1.4.5 - https://github.com/GoogleCloudPlatform/BigData-interop/gcsio/) (Apache 2.0) util (com.google.cloud.bigdataoss:util:1.4.5 - https://github.com/GoogleCloudPlatform/BigData-interop/util/) (Apache 2.0) Google Guice - Core Library (com.google.inject:guice:3.0 - http://code.google.com/p/google-guice/guice/) + (Apache 2.0) OpenAI Java SDK (com.openai:openai-java:4.69.3) - https://github.com/openai/openai-java/blob/main/LICENSE + (Apache 2.0) OkHttp (com.squareup.okhttp3:okhttp:4.12.0) - https://github.com/square/okhttp/blob/parent-4.12.0/LICENSE.txt + (Apache 2.0) Okio (com.squareup.okio:okio-jvm:3.6.0) - https://github.com/square/okio/blob/3.6.0/LICENSE.txt (Apache 2.0) OkHttp (com.squareup.okhttp:okhttp:2.5.0 - https://github.com/square/okhttp/okhttp) (Apache 2.0) Okio (com.squareup.okio:okio:1.6.0 - https://github.com/square/okio/okio) (Apache 2.0) OkHttp mockwebserver (com.squareup.okhttp3:mockwebserver:3.13.1) - https://github.com/square/okhttp/blob/master/LICENSE.txt diff --git a/zeppelin-server/pom.xml b/zeppelin-server/pom.xml index 7191c0bb01..ffe7d38eec 100644 --- a/zeppelin-server/pom.xml +++ b/zeppelin-server/pom.xml @@ -432,6 +432,18 @@ <artifactId>gson</artifactId> </dependency> + <dependency> + <groupId>com.squareup.okhttp3</groupId> + <artifactId>okhttp</artifactId> + <version>4.12.0</version> + <exclusions> + <exclusion> + <groupId>org.jetbrains.kotlin</groupId> + <artifactId>kotlin-stdlib-jdk8</artifactId> + </exclusion> + </exclusions> + </dependency> + <dependency> <groupId>com.openai</groupId> <artifactId>openai-java</artifactId> diff --git a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/OpenAiChatModel.java b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/OpenAiChatModel.java index 8808f81f69..10b64f7527 100644 --- a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/OpenAiChatModel.java +++ b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/OpenAiChatModel.java @@ -19,7 +19,8 @@ package org.apache.zeppelin.service.assistant; import com.google.gson.Gson; import com.openai.client.OpenAIClient; -import com.openai.client.okhttp.OpenAIOkHttpClient; +import com.openai.client.OpenAIClientImpl; +import com.openai.core.ClientOptions; import com.openai.core.JsonValue; import com.openai.core.http.StreamResponse; import com.openai.models.responses.EasyInputMessage; @@ -29,9 +30,12 @@ import com.openai.models.responses.ResponseFunctionToolCall; import com.openai.models.responses.ResponseInputItem; import com.openai.models.responses.ResponseStreamEvent; +import java.io.IOException; import java.util.ArrayList; import java.util.List; import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; import java.util.function.Consumer; /** @@ -46,6 +50,7 @@ public class OpenAiChatModel implements ChatModel { private final String model; private OpenAIClient cachedClient; private boolean closed; + private final Set<okhttp3.Call> activeCalls = ConcurrentHashMap.newKeySet(); public OpenAiChatModel(String baseUrl, String apiKey, String model) { this.baseUrl = baseUrl; @@ -58,7 +63,37 @@ public class OpenAiChatModel implements ChatModel { throw new IllegalStateException("Assistant model is closed"); } if (cachedClient == null) { - cachedClient = OpenAIOkHttpClient.builder().baseUrl(baseUrl).apiKey(apiKey).build(); + var options = ClientOptions.builder().baseUrl(baseUrl).apiKey(apiKey); + var timeout = options.timeout(); + var httpClient = new okhttp3.OkHttpClient.Builder() + .connectTimeout(timeout.connect()) + .readTimeout(timeout.read()) + .writeTimeout(timeout.write()) + .callTimeout(timeout.request()) + .followRedirects(false) + .retryOnConnectionFailure(false) + .eventListener(new okhttp3.EventListener() { + @Override + public void callStart(okhttp3.Call call) { + synchronized (OpenAiChatModel.this) { + if (closed) call.cancel(); + else activeCalls.add(call); + } + } + + @Override + public void callEnd(okhttp3.Call call) { + activeCalls.remove(call); + } + + @Override + public void callFailed(okhttp3.Call call, IOException error) { + activeCalls.remove(call); + } + }) + .build(); + cachedClient = new OpenAIClientImpl(options + .httpClient(new com.openai.client.okhttp.OkHttpClient(httpClient)).build()); } return cachedClient; } @@ -67,6 +102,7 @@ public class OpenAiChatModel implements ChatModel { public synchronized void close() { if (closed) return; closed = true; + activeCalls.forEach(okhttp3.Call::cancel); if (cachedClient != null) { cachedClient.close(); } @@ -100,11 +136,7 @@ public class OpenAiChatModel implements ChatModel { } boolean[] completed = {false}; - try ( - StreamResponse<ResponseStreamEvent> stream = client() - .responses() - .createStreaming(params.build()) - ) { + try (StreamResponse<ResponseStreamEvent> stream = client().responses().createStreaming(params.build())) { stream.stream().forEach(event -> { if (event.completed().isPresent()) completed[0] = true; handleEvent(event, consumer); diff --git a/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelLifecycleTest.java b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelLifecycleTest.java index b33ca2aa91..7d1d69869e 100644 --- a/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelLifecycleTest.java +++ b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelLifecycleTest.java @@ -18,11 +18,19 @@ package org.apache.zeppelin.service.assistant; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; import com.openai.client.OpenAIClient; +import com.sun.net.httpserver.HttpServer; +import java.net.InetSocketAddress; +import java.nio.charset.StandardCharsets; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; import java.util.List; import org.junit.jupiter.api.Test; @@ -43,6 +51,44 @@ class OpenAiChatModelLifecycleTest { () -> model.stream("instruction", List.of(), List.of(), event -> { })); } + @Test + void closingModelCancelsAStalledHttpStream() throws Exception { + var server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + var streaming = new CountDownLatch(1); + var release = new CountDownLatch(1); + server.createContext("/responses", exchange -> { + exchange.getRequestBody().readAllBytes(); + exchange.getResponseHeaders().set("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, 0); + try (var body = exchange.getResponseBody()) { + body.write(": waiting\n\n".getBytes(StandardCharsets.UTF_8)); + body.flush(); + streaming.countDown(); + try { + release.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + } + }); + server.start(); + var worker = Executors.newFixedThreadPool(2); + var model = new OpenAiChatModel( + "http://127.0.0.1:" + server.getAddress().getPort(), "test-key", "test-model"); + try { + var run = worker.submit(() -> model.stream("instruction", List.of(), List.of(), event -> { })); + assertTrue(streaming.await(5, TimeUnit.SECONDS)); + worker.submit(model::close).get(3, TimeUnit.SECONDS); + assertThrows(ExecutionException.class, () -> run.get(3, TimeUnit.SECONDS)); + } finally { + release.countDown(); + server.stop(0); + model.close(); + worker.shutdownNow(); + assertTrue(worker.awaitTermination(5, TimeUnit.SECONDS)); + } + } + @Test void closingUnusedModelDoesNotInitializeClient() { var model = new OpenAiChatModel("http://unused.invalid", "test-key", "test-model");
