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 80a0d79100e2354c0c1126d0dac1e0998de62d38
Author: YONGJAE LEE <[email protected]>
AuthorDate: Thu Oct 8 19:18:01 2026 +0900

    Close assistant resources during server shutdown
---
 .../org/apache/zeppelin/server/ZeppelinServer.java |  1 +
 .../service/assistant/AssistantService.java        | 43 +++++++++++---
 .../zeppelin/service/assistant/ChatModel.java      |  5 +-
 .../service/assistant/OpenAiChatModel.java         | 13 ++++
 .../org/apache/zeppelin/socket/NotebookServer.java | 17 ++++++
 .../service/assistant/AssistantServiceTest.java    | 50 ++++++++++++++++
 .../assistant/OpenAiChatModelLifecycleTest.java    | 53 +++++++++++++++++
 .../NotebookServerAssistantLifecycleTest.java      | 69 ++++++++++++++++++++++
 8 files changed, 243 insertions(+), 8 deletions(-)

diff --git 
a/zeppelin-server/src/main/java/org/apache/zeppelin/server/ZeppelinServer.java 
b/zeppelin-server/src/main/java/org/apache/zeppelin/server/ZeppelinServer.java
index a4b483e9d6..61088a7932 100644
--- 
a/zeppelin-server/src/main/java/org/apache/zeppelin/server/ZeppelinServer.java
+++ 
b/zeppelin-server/src/main/java/org/apache/zeppelin/server/ZeppelinServer.java
@@ -394,6 +394,7 @@ public class ZeppelinServer implements AutoCloseable {
         if (sharedServiceLocator != null) {
           // Stop after Jetty so no new connection can restart the heartbeat 
scheduler.
           
sharedServiceLocator.getService(NotebookServer.class).stopHeartbeatScheduler();
+          
sharedServiceLocator.getService(NotebookServer.class).stopAssistantRuns();
           if (!zConf.isRecoveryEnabled()) {
             
sharedServiceLocator.getService(InterpreterSettingManager.class).close();
           }
diff --git 
a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/AssistantService.java
 
b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/AssistantService.java
index eae06d9acd..505869d0dc 100644
--- 
a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/AssistantService.java
+++ 
b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/AssistantService.java
@@ -32,6 +32,8 @@ import java.util.Set;
 import java.util.UUID;
 import java.util.stream.IntStream;
 import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.CancellationException;
+import java.util.concurrent.atomic.AtomicBoolean;
 import jakarta.ws.rs.BadRequestException;
 import jakarta.ws.rs.ClientErrorException;
 import jakarta.ws.rs.ForbiddenException;
@@ -45,7 +47,7 @@ import 
org.apache.zeppelin.rest.exception.NoteNotFoundException;
 import org.apache.zeppelin.service.NotebookService;
 import org.apache.zeppelin.user.AuthenticationInfo;
 
-public class AssistantService {
+public class AssistantService implements AutoCloseable {
 
   private static final Logger LOGGER = 
LoggerFactory.getLogger(AssistantService.class);
   private static final Gson GSON = new Gson();
@@ -60,6 +62,7 @@ public class AssistantService {
 
   private final Set<String> busyConversations = ConcurrentHashMap.newKeySet();
 
+  private final AtomicBoolean closed = new AtomicBoolean();
   private final boolean available;
   private final Notebook notebook;
   private final ChatModel modelClient;
@@ -86,11 +89,24 @@ public class AssistantService {
     );
   }
 
+  @Override
+  public void close() {
+    if (closed.compareAndSet(false, true) && modelClient != null) {
+      modelClient.close();
+    }
+  }
+
+  private void checkRunActive() {
+    if (closed.get() || Thread.currentThread().isInterrupted()) {
+      throw new CancellationException("Assistant run stopped");
+    }
+  }
+
   public List<Conversation> listConversations(
       String noteId,
       Set<String> userAndRoles
   ) throws IOException {
-    if (!available) throw new ServiceUnavailableException();
+    if (!available || closed.get()) throw new ServiceUnavailableException();
     if (!authorizationService.isReader(noteId, userAndRoles)) {
       throw new ForbiddenException();
     }
@@ -107,7 +123,7 @@ public class AssistantService {
       AuthenticationInfo authInfo,
       Set<String> userAndRoles
   ) throws IOException {
-    if (!available) throw new ServiceUnavailableException();
+    if (!available || closed.get()) throw new ServiceUnavailableException();
     if (!authorizationService.isReader(noteId, userAndRoles)) throw new 
ForbiddenException();
     if (AuthenticationInfo.isAnonymous(authInfo)) throw new 
ForbiddenException();
 
@@ -124,7 +140,7 @@ public class AssistantService {
       String conversationId,
       Set<String> userAndRoles
   ) throws IOException {
-    if (!available) throw new ServiceUnavailableException();
+    if (!available || closed.get()) throw new ServiceUnavailableException();
     if (!authorizationService.isReader(noteId, userAndRoles)) throw new 
ForbiddenException();
 
     return notebook.processNote(noteId, note -> {
@@ -140,7 +156,7 @@ public class AssistantService {
       String userId,
       Set<String> userAndRoles
   ) throws IOException {
-    if (!available) throw new ServiceUnavailableException();
+    if (!available || closed.get()) throw new ServiceUnavailableException();
     if (!busyConversations.add(conversationId)) {
       throw new ClientErrorException("", Response.Status.CONFLICT);
     }
@@ -163,7 +179,7 @@ public class AssistantService {
       String userId,
       Set<String> userAndRoles
   ) throws IOException {
-    if (!available) throw new ServiceUnavailableException();
+    if (!available || closed.get()) throw new ServiceUnavailableException();
     if (!busyConversations.add(conversationId)) {
       throw new ClientErrorException("", Response.Status.CONFLICT);
     }
@@ -218,7 +234,8 @@ public class AssistantService {
     var runId = "run_" + UUID.randomUUID().toString().replace("-", 
"").substring(0, 16);
     boolean acquired = false;
     try {
-      if (!available) throw new ServiceUnavailableException();
+      if (!available || closed.get()) throw new ServiceUnavailableException();
+      checkRunActive();
       if (StringUtils.isBlank(userContent)) throw new BadRequestException();
       if (!busyConversations.add(conversationId)) {
         throw new ClientErrorException("", Response.Status.CONFLICT);
@@ -229,6 +246,7 @@ public class AssistantService {
       if (!conversation.isOwner(authInfo.getUser())) throw new 
ForbiddenException();
 
       notebook.processNote(noteId, note -> {
+        checkRunActive();
         if (note == null) throw new NoteNotFoundException(noteId);
         conversationRepository.update(
             conversation, c -> c.addMessage(Message.user(Message.id(), 
userContent))
@@ -246,6 +264,7 @@ public class AssistantService {
       tokens.put("output", 0);
 
       runLoop(noteId, conversation, authInfo, userAndRoles, sink, tokens);
+      checkRunActive();
 
       sink.onEvent(
           AssistantEventType.RUN_COMPLETED,
@@ -254,7 +273,10 @@ public class AssistantService {
               new AssistantEventPayload.Usage(tokens.get("input"), 
tokens.get("output"))
           )
       );
+    } catch (CancellationException e) {
+      // Shutdown has already closed the connection and must not publish a 
successful turn.
     } catch (Exception e) {
+      if (closed.get()) return;
       LOGGER.error("Error during Assistant run", e);
       sink.onEvent(
           AssistantEventType.RUN_FAILED,
@@ -274,6 +296,7 @@ public class AssistantService {
       Map<String, Integer> tokens
   ) throws IOException {
     for (int iteration = 0; iteration < 10; iteration++) {
+      checkRunActive();
       String assistantId = Message.id();
       var textBuffer = new StringBuilder();
       List<ToolCall> toolCalls = new ArrayList<>();
@@ -283,6 +306,7 @@ public class AssistantService {
           conversation.getRecentMessages(HISTORY_MESSAGE_LIMIT),
           toolExecutor.specs(),
           event -> {
+            checkRunActive();
             if (event instanceof AssistantEvent.TextDelta) {
               String delta = ((AssistantEvent.TextDelta) event).delta;
               textBuffer.append(delta);
@@ -301,6 +325,7 @@ public class AssistantService {
           }
       );
 
+      checkRunActive();
       if (!toolCalls.isEmpty()) {
         Message.Assistant assistantMsg = Message.assistant(assistantId, 
textBuffer.toString());
         toolCalls.forEach(assistantMsg::addToolCall);
@@ -316,6 +341,7 @@ public class AssistantService {
         }
 
         for (ToolCall tc : assistantMsg.getToolCalls()) {
+          checkRunActive();
           sink.onEvent(
               AssistantEventType.TOOL_CALL_STARTED,
               new AssistantEventPayload.ToolCallStarted(tc.getId(), 
tc.getName(), tc.getArguments())
@@ -323,6 +349,7 @@ public class AssistantService {
           ToolResult result = toolExecutor.callTool(
               noteId, tc.getName(), tc.getArguments(), authInfo, userAndRoles
           );
+          checkRunActive();
           tc.setResult(result);
           turn.add(Message.tool(Message.id(), tc.getId(), 
GSON.toJson(result)));
           sink.onEvent(
@@ -332,6 +359,7 @@ public class AssistantService {
         }
 
         notebook.processNote(noteId, note -> {
+          checkRunActive();
           if (note == null) throw new NoteNotFoundException(noteId);
           conversationRepository.update(conversation, c -> 
turn.forEach(c::addMessage));
           return null;
@@ -340,6 +368,7 @@ public class AssistantService {
         // No tool calls; finish the run.
         String assistantText = textBuffer.toString();
         notebook.processNote(noteId, note -> {
+          checkRunActive();
           if (note == null) throw new NoteNotFoundException(noteId);
           conversationRepository.update(
               conversation,
diff --git 
a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/ChatModel.java
 
b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/ChatModel.java
index 7c4ac2134a..f15dca1a41 100644
--- 
a/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/ChatModel.java
+++ 
b/zeppelin-server/src/main/java/org/apache/zeppelin/service/assistant/ChatModel.java
@@ -20,7 +20,10 @@ package org.apache.zeppelin.service.assistant;
 import java.util.List;
 import java.util.function.Consumer;
 
-public interface ChatModel {
+public interface ChatModel extends AutoCloseable {
+
+  @Override
+  default void close() { }
 
   void stream(
       String systemPrompt,
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 ad6b93e581..8808f81f69 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
@@ -45,6 +45,7 @@ public class OpenAiChatModel implements ChatModel {
   private final String apiKey;
   private final String model;
   private OpenAIClient cachedClient;
+  private boolean closed;
 
   public OpenAiChatModel(String baseUrl, String apiKey, String model) {
     this.baseUrl = baseUrl;
@@ -53,12 +54,24 @@ public class OpenAiChatModel implements ChatModel {
   }
 
   private synchronized OpenAIClient client() {
+    if (closed) {
+      throw new IllegalStateException("Assistant model is closed");
+    }
     if (cachedClient == null) {
       cachedClient = 
OpenAIOkHttpClient.builder().baseUrl(baseUrl).apiKey(apiKey).build();
     }
     return cachedClient;
   }
 
+  @Override
+  public synchronized void close() {
+    if (closed) return;
+    closed = true;
+    if (cachedClient != null) {
+      cachedClient.close();
+    }
+  }
+
   @Override
   public void stream(
       String instruction,
diff --git 
a/zeppelin-server/src/main/java/org/apache/zeppelin/socket/NotebookServer.java 
b/zeppelin-server/src/main/java/org/apache/zeppelin/socket/NotebookServer.java
index 0f57899631..f0fa255e37 100644
--- 
a/zeppelin-server/src/main/java/org/apache/zeppelin/socket/NotebookServer.java
+++ 
b/zeppelin-server/src/main/java/org/apache/zeppelin/socket/NotebookServer.java
@@ -40,6 +40,7 @@ import java.util.Set;
 import java.util.concurrent.ConcurrentHashMap;
 import java.util.concurrent.ExecutorService;
 import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
 import java.util.concurrent.ScheduledExecutorService;
 import java.util.concurrent.TimeUnit;
 import java.util.concurrent.atomic.AtomicReference;
@@ -328,6 +329,22 @@ public class NotebookServer implements 
AngularObjectRegistryListener,
     heartbeatInitialized = false;
   }
 
+  public void stopAssistantRuns() {
+    for (Runnable queued : assistantExecutor.shutdownNow()) {
+      if (queued instanceof Future<?>) {
+        ((Future<?>) queued).cancel(false);
+      }
+    }
+    getAssistantService().close();
+    try {
+      if (!assistantExecutor.awaitTermination(10, TimeUnit.SECONDS)) {
+        LOGGER.warn("Assistant runs did not stop within 10 seconds");
+      }
+    } catch (InterruptedException e) {
+      Thread.currentThread().interrupt();
+    }
+  }
+
   /**
    * Sends a WebSocket protocol ping frame to every connected session. Writing 
to a session
    * resets Jetty's idle timeout (and any intermediate proxy's idle timer), 
which is the whole
diff --git 
a/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/AssistantServiceTest.java
 
b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/AssistantServiceTest.java
index 6d8fb618ae..db205cc42e 100644
--- 
a/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/AssistantServiceTest.java
+++ 
b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/AssistantServiceTest.java
@@ -465,4 +465,54 @@ class AssistantServiceTest {
       executor.shutdownNow();
     }
   }
+  @Test
+  void closingBlockedRunPreventsPersistenceAndNewRequests() throws Exception {
+    var notebook = mock(Notebook.class);
+    when(notebook.processNote(eq("noteId"), any())).thenAnswer(invocation ->
+        ((NoteProcessor<?>) invocation.getArgument(1)).process(new Note()));
+    var authorization = mock(AuthorizationService.class);
+    when(authorization.isReader("noteId", userAndRoles)).thenReturn(true);
+    var repository = mock(ConversationRepository.class);
+    var conversation = Conversation.create("noteId", "test", 
authInfo.getUser());
+    when(repository.find("noteId", 
conversation.getId())).thenReturn(Optional.of(conversation));
+    var model = mock(ChatModel.class);
+    var notebookService = mock(NotebookService.class);
+    var sut = new AssistantService(
+        true, notebook, model, notebookService, authorization, repository);
+    var started = new CountDownLatch(1);
+    var release = new CountDownLatch(1);
+    var executor = Executors.newSingleThreadExecutor();
+    var events = new ArrayList<AssistantEventType>();
+    doAnswer(invocation -> {
+      invocation.<Consumer<AssistantEvent>>getArgument(3).accept(
+          new AssistantEvent.ToolCall("call_1", "list_paragraphs", "{}"));
+      started.countDown();
+      assertTrue(release.await(5, TimeUnit.SECONDS));
+      return null;
+    }).when(model).stream(any(), any(), any(), any());
+    doAnswer(invocation -> {
+      release.countDown();
+      return null;
+    }).when(model).close();
+    try {
+      var run = executor.submit(() -> sut.sendMessage("noteId", 
conversation.getId(), "hi",
+          authInfo, userAndRoles, (type, payload) -> events.add(type)));
+      assertTrue(started.await(5, TimeUnit.SECONDS));
+      sut.close();
+      run.get(5, TimeUnit.SECONDS);
+      sut.close();
+      verify(model, times(1)).close();
+      verify(repository, times(1)).update(any(), any());
+      verifyNoInteractions(notebookService);
+      assertEquals(List.of(AssistantEventType.RUN_STARTED), events);
+      assertThrows(ServiceUnavailableException.class,
+          () -> sut.listConversations("noteId", userAndRoles));
+      assertThrows(ServiceUnavailableException.class,
+          () -> sut.createConversation("noteId", "title", authInfo, 
userAndRoles));
+    } finally {
+      release.countDown();
+      executor.shutdownNow();
+    }
+  }
+
 }
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
new file mode 100644
index 0000000000..b33ca2aa91
--- /dev/null
+++ 
b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelLifecycleTest.java
@@ -0,0 +1,53 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.zeppelin.service.assistant;
+
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.times;
+import static org.mockito.Mockito.verify;
+
+import com.openai.client.OpenAIClient;
+import java.util.List;
+import org.junit.jupiter.api.Test;
+
+class OpenAiChatModelLifecycleTest {
+  @Test
+  void closesOwnedClientOnceAndRejectsReuse() throws Exception {
+    var model = new OpenAiChatModel("http://unused.invalid";, "test-key", 
"test-model");
+    var client = mock(OpenAIClient.class);
+    var field = OpenAiChatModel.class.getDeclaredField("cachedClient");
+    field.setAccessible(true);
+    field.set(model, client);
+
+    model.close();
+    model.close();
+
+    verify(client, times(1)).close();
+    assertThrows(IllegalStateException.class,
+        () -> model.stream("instruction", List.of(), List.of(), event -> { }));
+  }
+
+  @Test
+  void closingUnusedModelDoesNotInitializeClient() {
+    var model = new OpenAiChatModel("http://unused.invalid";, "test-key", 
"test-model");
+    model.close();
+    assertThrows(IllegalStateException.class,
+        () -> model.stream("instruction", List.of(), List.of(), event -> { }));
+  }
+}
diff --git 
a/zeppelin-server/src/test/java/org/apache/zeppelin/socket/NotebookServerAssistantLifecycleTest.java
 
b/zeppelin-server/src/test/java/org/apache/zeppelin/socket/NotebookServerAssistantLifecycleTest.java
new file mode 100644
index 0000000000..54c9700946
--- /dev/null
+++ 
b/zeppelin-server/src/test/java/org/apache/zeppelin/socket/NotebookServerAssistantLifecycleTest.java
@@ -0,0 +1,69 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.zeppelin.socket;
+
+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.verify;
+
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.RejectedExecutionException;
+import java.util.concurrent.TimeUnit;
+import org.apache.zeppelin.service.assistant.AssistantService;
+import org.junit.jupiter.api.Test;
+
+class NotebookServerAssistantLifecycleTest {
+  @Test
+  void shutdownInterruptsActiveRunsCancelsQueueAndRejectsSubmissions() throws 
Exception {
+    var server = new NotebookServer();
+    var service = mock(AssistantService.class);
+    server.setAssistantService(() -> service);
+    var field = NotebookServer.class.getDeclaredField("assistantExecutor");
+    field.setAccessible(true);
+    var executor = (ExecutorService) field.get(server);
+    var active = new CountDownLatch(10);
+    var interrupted = new CountDownLatch(10);
+    var release = new CountDownLatch(1);
+    try {
+      for (int i = 0; i < 10; i++) {
+        executor.submit(() -> {
+          active.countDown();
+          try {
+            release.await();
+          } catch (InterruptedException e) {
+            interrupted.countDown();
+            Thread.currentThread().interrupt();
+          }
+        });
+      }
+      assertTrue(active.await(5, TimeUnit.SECONDS));
+      var queued = executor.submit(() -> { throw new AssertionError("Queued 
run started"); });
+      server.stopAssistantRuns();
+      assertTrue(interrupted.await(5, TimeUnit.SECONDS));
+      assertTrue(queued.isCancelled());
+      assertTrue(executor.isTerminated());
+      verify(service).close();
+      assertThrows(RejectedExecutionException.class, () -> executor.submit(() 
-> { }));
+    } finally {
+      release.countDown();
+      executor.shutdownNow();
+    }
+  }
+}

Reply via email to