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(); + } + } +}
