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 91414031ea7807c69a84371bd72dec60a2b48635 Author: YONGJAE LEE <[email protected]> AuthorDate: Thu Oct 8 20:24:39 2026 +0900 Bound assistant queue and verify stream events --- .../org/apache/zeppelin/socket/NotebookServer.java | 44 +++++-- .../assistant/OpenAiChatModelEventMappingTest.java | 146 +++++++++++++++++++++ .../NotebookServerAssistantLifecycleTest.java | 61 +++++++++ 3 files changed, 238 insertions(+), 13 deletions(-) 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 f0fa255e37..ddbacc8d9d 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 @@ -37,11 +37,15 @@ import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Set; +import java.util.UUID; +import java.util.concurrent.ArrayBlockingQueue; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; import java.util.concurrent.Future; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; import jakarta.inject.Inject; @@ -94,6 +98,7 @@ import org.apache.zeppelin.scheduler.Job.Status; import org.apache.zeppelin.service.JobManagerService; import org.apache.zeppelin.service.NotebookService; import org.apache.zeppelin.service.assistant.AssistantEventListener; +import org.apache.zeppelin.service.assistant.AssistantEventType; import org.apache.zeppelin.service.assistant.AssistantService; import org.apache.zeppelin.service.ServiceContext; import org.apache.zeppelin.service.SimpleServiceCallback; @@ -152,9 +157,10 @@ public class NotebookServer implements AngularObjectRegistryListener, private final ExecutorService executorService = Executors.newFixedThreadPool(10); - private final ExecutorService assistantExecutor = Executors.newFixedThreadPool( - 10, - new ThreadFactoryBuilder().setNameFormat("assistant-run-%d").setDaemon(true).build() + private final ExecutorService assistantExecutor = new ThreadPoolExecutor( + 10, 10, 0L, TimeUnit.MILLISECONDS, new ArrayBlockingQueue<>(64), + new ThreadFactoryBuilder().setNameFormat("assistant-run-%d").setDaemon(true).build(), + new ThreadPoolExecutor.AbortPolicy() ); // Package-private (not private) so NotebookServerHeartbeatTest can observe scheduler @@ -1333,16 +1339,28 @@ public class NotebookServer implements AngularObjectRegistryListener, LOGGER.warn("Failed to send assistant event to connection", e); } }; - assistantExecutor.submit( - () -> getAssistantService().sendMessage( - noteId, - conversationId, - content, - context.getAutheInfo(), - context.getUserAndRoles(), - sink - ) - ); + try { + assistantExecutor.submit( + () -> getAssistantService().sendMessage( + noteId, + conversationId, + content, + context.getAutheInfo(), + context.getUserAndRoles(), + sink + ) + ); + } catch (RejectedExecutionException e) { + try { + conn.send(serializeMessage(new Message(OP.ASSISTANT_EVENT) + .put("conversationId", conversationId) + .put("type", AssistantEventType.RUN_FAILED.wireName) + .put("payload", Map.of("runId", UUID.randomUUID().toString(), + "error", Map.of("status", 429))))); + } catch (IOException sendError) { + LOGGER.warn("Failed to send assistant rejection to connection", sendError); + } + } } private void cloneNote(NotebookSocket conn, diff --git a/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelEventMappingTest.java b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelEventMappingTest.java new file mode 100644 index 0000000000..e5c1e015c5 --- /dev/null +++ b/zeppelin-server/src/test/java/org/apache/zeppelin/service/assistant/OpenAiChatModelEventMappingTest.java @@ -0,0 +1,146 @@ +/* + * 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.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.google.gson.Gson; +import com.sun.net.httpserver.HttpServer; +import java.net.InetSocketAddress; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.function.Consumer; +import org.junit.jupiter.api.Test; + +class OpenAiChatModelEventMappingTest { + private static final Gson GSON = new Gson(); + + @Test + void mapsTextRefusalToolCallAndCompletionUsageInOrder() throws Exception { + var events = new ArrayList<AssistantEvent>(); + withStream(List.of( + delta("response.output_text.delta", "Hello"), + delta("response.refusal.delta", "I cannot help with that"), + Map.of("type", "response.output_item.done", "output_index", 0, + "sequence_number", 2, "item", Map.of("type", "function_call", "id", "item", + "call_id", "call", "name", "listNotes", "arguments", "{\"limit\":2}", + "status", "completed")), + completed(Map.of("input_tokens", 12, "output_tokens", 7, "total_tokens", 19, + "input_tokens_details", Map.of("cached_tokens", 0), + "output_tokens_details", Map.of("reasoning_tokens", 0)))), + model -> model.stream("instruction", List.of(), List.of(), events::add)); + + assertEquals(4, events.size()); + assertEquals("Hello", ((AssistantEvent.TextDelta) events.get(0)).delta); + assertEquals("I cannot help with that", ((AssistantEvent.TextDelta) events.get(1)).delta); + var tool = (AssistantEvent.ToolCall) events.get(2); + assertEquals("call", tool.id); + assertEquals("listNotes", tool.name); + assertEquals(Map.of("limit", 2.0), tool.arguments); + var usage = (AssistantEvent.Usage) events.get(3); + assertEquals(12, usage.inputTokens); + assertEquals(7, usage.outputTokens); + } + + @Test + void errorEventFailsWithoutPublishingAReply() throws Exception { + var events = new ArrayList<AssistantEvent>(); + withStream(List.of(Map.of("type", "error", "code", "server_error", + "message", "upstream unavailable", "sequence_number", 0)), model -> { + var error = assertThrows(IllegalStateException.class, + () -> model.stream("instruction", List.of(), List.of(), events::add)); + assertEquals("OpenAI stream error: upstream unavailable", error.getMessage()); + }); + assertTrue(events.isEmpty()); + } + + @Test + void failedAndIncompleteEventsFailInsteadOfCompleting() throws Exception { + for (String status : List.of("failed", "incomplete")) { + var events = new ArrayList<AssistantEvent>(); + withStream(List.of(Map.of("type", "response." + status, "sequence_number", 0, + "response", Map.of("id", "response", "status", status))), model -> { + var error = assertThrows(IllegalStateException.class, + () -> model.stream("instruction", List.of(), List.of(), events::add)); + assertEquals("OpenAI response " + status, error.getMessage()); + }); + assertTrue(events.isEmpty()); + } + } + + @Test + void streamWithoutCompletionFailsAfterPartialText() throws Exception { + var events = new ArrayList<AssistantEvent>(); + withStream(List.of(delta("response.output_text.delta", "partial")), model -> + assertThrows(IllegalStateException.class, + () -> model.stream("instruction", List.of(), List.of(), events::add))); + assertEquals(1, events.size()); + assertEquals("partial", ((AssistantEvent.TextDelta) events.get(0)).delta); + } + + @Test + void completionWithoutUsageDoesNotInventUsage() throws Exception { + var events = new ArrayList<AssistantEvent>(); + withStream(List.of(Map.of("type", "response.completed", "sequence_number", 0, + "response", Map.of("id", "response", "status", "completed"))), model -> + model.stream("instruction", List.of(), List.of(), events::add)); + assertTrue(events.isEmpty()); + } + + private static Map<String, Object> delta(String type, String text) { + return Map.of("type", type, "delta", text, "item_id", "item", "content_index", 0, + "output_index", 0, "sequence_number", 0, "logprobs", List.of()); + } + + private static Map<String, Object> completed(Map<String, Object> usage) { + return Map.of("type", "response.completed", "sequence_number", 3, + "response", Map.of("id", "response", "status", "completed", "usage", usage)); + } + + private static void withStream(List<Map<String, Object>> events, + Consumer<OpenAiChatModel> check) throws Exception { + var server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); + var stream = new StringBuilder(); + for (var event : events) { + stream.append("event: ").append(event.get("type")) + .append("\ndata: ").append(GSON.toJson(event)).append("\n\n"); + } + var bytes = stream.toString().getBytes(StandardCharsets.UTF_8); + server.createContext("/responses", exchange -> { + exchange.getRequestBody().readAllBytes(); + exchange.getResponseHeaders().set("Content-Type", "text/event-stream"); + exchange.sendResponseHeaders(200, bytes.length); + try (var body = exchange.getResponseBody()) { + body.write(bytes); + } + }); + server.start(); + var model = new OpenAiChatModel("http://127.0.0.1:" + server.getAddress().getPort(), + "test-key", "test-model"); + try { + check.accept(model); + } finally { + model.close(); + server.stop(0); + } + } +} 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 54c9700946..203d91a11f 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 @@ -17,15 +17,25 @@ package org.apache.zeppelin.socket; +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.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import com.google.gson.JsonParser; +import java.util.Set; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.TimeUnit; +import java.util.concurrent.ThreadPoolExecutor; +import org.apache.zeppelin.common.Message; +import org.apache.zeppelin.common.Message.OP; +import org.apache.zeppelin.service.ServiceContext; +import org.apache.zeppelin.user.AuthenticationInfo; +import org.mockito.ArgumentCaptor; import org.apache.zeppelin.service.assistant.AssistantService; import org.junit.jupiter.api.Test; @@ -66,4 +76,55 @@ class NotebookServerAssistantLifecycleTest { executor.shutdownNow(); } } + @Test + void saturatedAssistantExecutorReportsTooManyRequests() throws Exception { + var server = new NotebookServer(); + var service = mock(AssistantService.class); + server.setAssistantService(() -> service); + var socket = mock(NotebookSocket.class); + var field = NotebookServer.class.getDeclaredField("assistantExecutor"); + field.setAccessible(true); + var executor = (ThreadPoolExecutor) field.get(server); + var active = 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) { + Thread.currentThread().interrupt(); + } + }); + } + assertTrue(active.await(5, TimeUnit.SECONDS)); + assertEquals(64, executor.getQueue().remainingCapacity()); + for (int i = 0; i < 64; i++) { + executor.submit(() -> { throw new AssertionError("Queued run started"); }); + } + var send = NotebookServer.class.getDeclaredMethod("sendAssistantMessage", + NotebookSocket.class, ServiceContext.class, Message.class); + send.setAccessible(true); + send.invoke(server, socket, + new ServiceContext(AuthenticationInfo.ANONYMOUS, Set.of()), + new Message(OP.ASSISTANT_SEND_MESSAGE).put("noteId", "note") + .put("conversationId", "conversation").put("content", "hello")); + var response = ArgumentCaptor.forClass(String.class); + verify(socket).send(response.capture()); + var message = JsonParser.parseString(response.getValue()).getAsJsonObject(); + assertEquals("ASSISTANT_EVENT", message.get("op").getAsString()); + var data = message.getAsJsonObject("data"); + assertEquals("conversation", data.get("conversationId").getAsString()); + assertEquals("run.failed", data.get("type").getAsString()); + var payload = data.getAsJsonObject("payload"); + assertTrue(!payload.get("runId").getAsString().isEmpty()); + assertEquals(429, payload.getAsJsonObject("error").get("status").getAsInt()); + assertEquals(64, executor.getQueue().size()); + verifyNoInteractions(service); + } finally { + server.stopAssistantRuns(); + release.countDown(); + } + } }
