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 2d9eefa2a2bea848fb39127b407fbb3b3c2f933c Author: YONGJAE LEE <[email protected]> AuthorDate: Thu Oct 8 21:10:08 2026 +0900 Expose assistant run status for recovery --- .../rest/AssistantConversationRestApi.java | 12 +++-- .../rest/message/ConversationMetadata.java | 9 ++-- .../rest/message/ConversationResponse.java | 8 ++-- .../service/assistant/AssistantService.java | 22 +++++++++- .../org/apache/zeppelin/socket/NotebookServer.java | 21 ++++----- .../rest/message/ConversationRunStatusTest.java | 51 ++++++++++++++++++++++ .../service/assistant/AssistantServiceTest.java | 34 ++++++++++++++- .../NotebookServerAssistantLifecycleTest.java | 45 ++++++++++++++++++- 8 files changed, 176 insertions(+), 26 deletions(-) diff --git a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/AssistantConversationRestApi.java b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/AssistantConversationRestApi.java index 3463bdc12a..949c217b2e 100644 --- a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/AssistantConversationRestApi.java +++ b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/AssistantConversationRestApi.java @@ -75,7 +75,8 @@ public class AssistantConversationRestApi extends AbstractRestApi { Response.Status.OK, "", assistantService.listConversations(noteId, context.getUserAndRoles()).stream() - .map(c -> ConversationMetadata.of(c, context.getAutheInfo().getUser())) + .map(c -> ConversationMetadata.of( + c, context.getAutheInfo().getUser(), assistantService.isConversationRunning(c.getId()))) .collect(Collectors.toList()) ).build(); } @@ -96,7 +97,8 @@ public class AssistantConversationRestApi extends AbstractRestApi { return new JsonResponse<>( Response.Status.OK, "", - ConversationResponse.of(conversation, context.getAutheInfo().getUser()) + ConversationResponse.of(conversation, context.getAutheInfo().getUser(), + assistantService.isConversationRunning(conversation.getId())) ).build(); } @@ -116,7 +118,8 @@ public class AssistantConversationRestApi extends AbstractRestApi { return new JsonResponse<>( Response.Status.CREATED, "", - ConversationMetadata.of(conversation, context.getAutheInfo().getUser()) + ConversationMetadata.of(conversation, context.getAutheInfo().getUser(), + assistantService.isConversationRunning(conversation.getId())) ).build(); } @@ -138,7 +141,8 @@ public class AssistantConversationRestApi extends AbstractRestApi { return new JsonResponse<>( Response.Status.OK, "", - ConversationMetadata.of(conversation, context.getAutheInfo().getUser()) + ConversationMetadata.of(conversation, context.getAutheInfo().getUser(), + assistantService.isConversationRunning(conversation.getId())) ).build(); } diff --git a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationMetadata.java b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationMetadata.java index 408428c7d4..4244a2b3b3 100644 --- a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationMetadata.java +++ b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationMetadata.java @@ -27,6 +27,7 @@ public final class ConversationMetadata { private final String createdAt; private final String updatedAt; private final boolean canSendMessage; + private final boolean running; private ConversationMetadata( String id, @@ -35,7 +36,8 @@ public final class ConversationMetadata { String title, String createdAt, String updatedAt, - boolean canSendMessage + boolean canSendMessage, + boolean running ) { this.id = id; this.noteId = noteId; @@ -44,12 +46,13 @@ public final class ConversationMetadata { this.createdAt = createdAt; this.updatedAt = updatedAt; this.canSendMessage = canSendMessage; + this.running = running; } - public static ConversationMetadata of(Conversation c, String userId) { + public static ConversationMetadata of(Conversation c, String userId, boolean running) { return new ConversationMetadata( c.getId(), c.getNoteId(), c.getOwnerId(), c.getTitle(), c.getCreatedAt(), c.getUpdatedAt(), - c.isOwner(userId) + c.isOwner(userId), running ); } } diff --git a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationResponse.java b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationResponse.java index e747be73b7..f23e6c4fc6 100644 --- a/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationResponse.java +++ b/zeppelin-server/src/main/java/org/apache/zeppelin/rest/message/ConversationResponse.java @@ -29,9 +29,10 @@ public final class ConversationResponse { private final String createdAt; private final String updatedAt; private final boolean canSendMessage; + private final boolean running; private final List<Message> messages; - private ConversationResponse(Conversation conversation, String userId) { + private ConversationResponse(Conversation conversation, String userId, boolean running) { this.id = conversation.getId(); this.noteId = conversation.getNoteId(); this.ownerId = conversation.getOwnerId(); @@ -39,10 +40,11 @@ public final class ConversationResponse { this.createdAt = conversation.getCreatedAt(); this.updatedAt = conversation.getUpdatedAt(); this.canSendMessage = conversation.isOwner(userId); + this.running = running; this.messages = List.copyOf(conversation.getMessages()); } - public static ConversationResponse of(Conversation conversation, String userId) { - return new ConversationResponse(conversation, userId); + public static ConversationResponse of(Conversation conversation, String userId, boolean running) { + return new ConversationResponse(conversation, userId, running); } } 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 094648b7fb..f3982e7435 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 @@ -61,6 +61,7 @@ public class AssistantService implements AutoCloseable { private final AuthorizationService authorizationService; private final Set<String> busyConversations = ConcurrentHashMap.newKeySet(); + private final Map<String, Integer> pendingRequests = new ConcurrentHashMap<>(); private final AtomicBoolean closed = new AtomicBoolean(); private final boolean available; @@ -91,8 +92,9 @@ public class AssistantService implements AutoCloseable { @Override public void close() { - if (closed.compareAndSet(false, true) && modelClient != null) { - modelClient.close(); + if (closed.compareAndSet(false, true)) { + pendingRequests.clear(); + if (modelClient != null) modelClient.close(); } } @@ -149,6 +151,22 @@ public class AssistantService implements AutoCloseable { }); } + public boolean isConversationRunning(String conversationId) { + return !closed.get() && (busyConversations.contains(conversationId) + || pendingRequests.containsKey(conversationId)); + } + + public void registerPendingRequest(String conversationId) { + if (conversationId == null) return; + pendingRequests.compute(conversationId, + (id, count) -> closed.get() ? null : count == null ? 1 : count + 1); + } + + public void releasePendingRequest(String conversationId) { + if (conversationId == null) return; + pendingRequests.computeIfPresent(conversationId, (id, count) -> count == 1 ? null : count - 1); + } + public Conversation updateTitle( String noteId, String conversationId, 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 ddbacc8d9d..5ba5d309a4 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 @@ -1339,18 +1339,19 @@ public class NotebookServer implements AngularObjectRegistryListener, LOGGER.warn("Failed to send assistant event to connection", e); } }; + var service = getAssistantService(); + service.registerPendingRequest(conversationId); try { - assistantExecutor.submit( - () -> getAssistantService().sendMessage( - noteId, - conversationId, - content, - context.getAutheInfo(), - context.getUserAndRoles(), - sink - ) - ); + assistantExecutor.submit(() -> { + try { + service.sendMessage(noteId, conversationId, content, + context.getAutheInfo(), context.getUserAndRoles(), sink); + } finally { + service.releasePendingRequest(conversationId); + } + }); } catch (RejectedExecutionException e) { + service.releasePendingRequest(conversationId); try { conn.send(serializeMessage(new Message(OP.ASSISTANT_EVENT) .put("conversationId", conversationId) diff --git a/zeppelin-server/src/test/java/org/apache/zeppelin/rest/message/ConversationRunStatusTest.java b/zeppelin-server/src/test/java/org/apache/zeppelin/rest/message/ConversationRunStatusTest.java new file mode 100644 index 0000000000..354695ef12 --- /dev/null +++ b/zeppelin-server/src/test/java/org/apache/zeppelin/rest/message/ConversationRunStatusTest.java @@ -0,0 +1,51 @@ +/* + * 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.rest.message; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.google.gson.Gson; +import org.apache.zeppelin.service.assistant.Conversation; +import org.junit.jupiter.api.Test; + +class ConversationRunStatusTest { + private final Gson gson = new Gson(); + + @Test + void exposesTransientRunStatusWithoutChangingStoredConversation() { + var conversation = Conversation.create("note", "title", "owner"); + var runningMetadata = gson.toJsonTree(ConversationMetadata.of(conversation, "owner", true)) + .getAsJsonObject(); + var idleMetadata = gson.toJsonTree(ConversationMetadata.of(conversation, "owner", false)) + .getAsJsonObject(); + var runningResponse = gson.toJsonTree(ConversationResponse.of(conversation, "reader", true)) + .getAsJsonObject(); + var idleResponse = gson.toJsonTree(ConversationResponse.of(conversation, "owner", false)) + .getAsJsonObject(); + + assertTrue(runningMetadata.get("running").getAsBoolean()); + assertFalse(idleMetadata.get("running").getAsBoolean()); + assertTrue(runningResponse.get("running").getAsBoolean()); + assertFalse(idleResponse.get("running").getAsBoolean()); + assertFalse(runningResponse.get("canSendMessage").getAsBoolean()); + assertEquals(conversation.getId(), runningResponse.get("id").getAsString()); + assertFalse(gson.toJsonTree(conversation).getAsJsonObject().has("running")); + } +} 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 612bc8393b..5f1cb7ff02 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 @@ -49,6 +49,8 @@ import org.apache.zeppelin.service.NotebookService; import org.apache.zeppelin.user.AuthenticationInfo; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; class AssistantServiceTest { private final AuthenticationInfo authInfo = new AuthenticationInfo("holden"); @@ -441,7 +443,24 @@ class AssistantServiceTest { } @Test - void guardsRunningConversation() throws Exception { + void pendingRequestCountsRemainRunningUntilEveryRequestFinishes() { + var sut = new AssistantService(true, null, null, null, null, null); + sut.registerPendingRequest("conversation"); + sut.registerPendingRequest("conversation"); + assertTrue(sut.isConversationRunning("conversation")); + sut.releasePendingRequest("conversation"); + assertTrue(sut.isConversationRunning("conversation")); + sut.releasePendingRequest("conversation"); + assertFalse(sut.isConversationRunning("conversation")); + sut.registerPendingRequest("conversation"); + sut.close(); + assertFalse(sut.isConversationRunning("conversation")); + sut.releasePendingRequest("conversation"); + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void guardsRunningConversation(boolean modelFails) throws Exception { var notebook = mock(Notebook.class); when(notebook.processNote(eq("noteId"), any())).thenAnswer(invocation -> ((NoteProcessor<?>) invocation.getArgument(1)).process(new Note())); @@ -465,19 +484,26 @@ class AssistantServiceTest { doAnswer(invocation -> { streaming.countDown(); // signal streaming has started assertTrue(finish.await(5, TimeUnit.SECONDS)); // wait for finish + if (modelFails) throw new IOException("Model failed"); return null; }).when(chatModel).stream(any(), any(), any(), any()); try { + sut.registerPendingRequest(conversation.getId()); + assertTrue(sut.isConversationRunning(conversation.getId())); + var runEvents = new ArrayList<AssistantEventType>(); var running = executor.submit(() -> sut.sendMessage( "noteId", conversation.getId(), "hi", authInfo, userAndRoles, - (type, payload) -> { }) + (type, payload) -> runEvents.add(type)) ); assertTrue(streaming.await(5, TimeUnit.SECONDS)); + sut.releasePendingRequest(conversation.getId()); + assertTrue(sut.isConversationRunning(conversation.getId())); // other conversations can be mutated. var second = Conversation.create("noteId", "independent", authInfo.getUser()); when(repository.find("noteId", second.getId())).thenReturn(Optional.of(second)); + assertFalse(sut.isConversationRunning(second.getId())); assertEquals("updated", sut.updateTitle("noteId", second.getId(), "updated", authInfo.getUser(), userAndRoles) .getTitle() @@ -498,9 +524,13 @@ class AssistantServiceTest { (type, payload) -> events.add(payload) ); assertEquals(409, ((AssistantEventPayload.RunFailed) events.get(0)).error.status); + assertTrue(sut.isConversationRunning(conversation.getId())); finish.countDown(); running.get(5, TimeUnit.SECONDS); // wait for message processing to complete. + assertFalse(sut.isConversationRunning(conversation.getId())); + assertEquals(modelFails ? AssistantEventType.RUN_FAILED : AssistantEventType.RUN_COMPLETED, + runEvents.get(runEvents.size() - 1)); assertEquals("released", sut.updateTitle("noteId", conversation.getId(), "released", authInfo.getUser(), userAndRoles) 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 index 203d91a11f..43e80450f6 100644 --- a/zeppelin-server/src/test/java/org/apache/zeppelin/socket/NotebookServerAssistantLifecycleTest.java +++ b/zeppelin-server/src/test/java/org/apache/zeppelin/socket/NotebookServerAssistantLifecycleTest.java @@ -21,8 +21,11 @@ import static org.junit.jupiter.api.Assertions.assertEquals; 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.doAnswer; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.verify; -import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.verifyNoMoreInteractions; import com.google.gson.JsonParser; import java.util.Set; @@ -38,6 +41,8 @@ import org.apache.zeppelin.user.AuthenticationInfo; import org.mockito.ArgumentCaptor; import org.apache.zeppelin.service.assistant.AssistantService; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; class NotebookServerAssistantLifecycleTest { @Test @@ -76,6 +81,40 @@ class NotebookServerAssistantLifecycleTest { executor.shutdownNow(); } } + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void releasesPendingRequestAfterWorkerReturnsOrThrows(boolean fails) throws Exception { + var server = new NotebookServer(); + var service = mock(AssistantService.class); + server.setAssistantService(() -> service); + var released = new CountDownLatch(1); + doAnswer(invocation -> { + if (fails) throw new IllegalStateException("Run failed"); + return null; + }).when(service).sendMessage(eq("note"), eq("conversation"), eq("hello"), + any(), any(), any()); + doAnswer(invocation -> { + released.countDown(); + return null; + }).when(service).releasePendingRequest("conversation"); + var send = NotebookServer.class.getDeclaredMethod("sendAssistantMessage", + NotebookSocket.class, ServiceContext.class, Message.class); + send.setAccessible(true); + try { + send.invoke(server, mock(NotebookSocket.class), + new ServiceContext(AuthenticationInfo.ANONYMOUS, Set.of()), + new Message(OP.ASSISTANT_SEND_MESSAGE).put("noteId", "note") + .put("conversationId", "conversation").put("content", "hello")); + assertTrue(released.await(5, TimeUnit.SECONDS)); + verify(service).registerPendingRequest("conversation"); + verify(service).releasePendingRequest("conversation"); + verify(service).sendMessage(eq("note"), eq("conversation"), eq("hello"), + any(), any(), any()); + } finally { + server.stopAssistantRuns(); + } + } + @Test void saturatedAssistantExecutorReportsTooManyRequests() throws Exception { var server = new NotebookServer(); @@ -121,7 +160,9 @@ class NotebookServerAssistantLifecycleTest { assertTrue(!payload.get("runId").getAsString().isEmpty()); assertEquals(429, payload.getAsJsonObject("error").get("status").getAsInt()); assertEquals(64, executor.getQueue().size()); - verifyNoInteractions(service); + verify(service).registerPendingRequest("conversation"); + verify(service).releasePendingRequest("conversation"); + verifyNoMoreInteractions(service); } finally { server.stopAssistantRuns(); release.countDown();
