diff --git a/AGENTS.md b/AGENTS.md index 9f936f64..6840f0ec 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -131,7 +131,7 @@ Incoming message → ChannelMessageReceivedEvent (channel name, message text) - **`ChannelRegistry`**: Registers channels, tracks last-active channel so background task replies are routed correctly. - **`DiscordChannel`**: JDA `ListenerAdapter`; accepts DMs from the configured user and guild messages only when the bot is mentioned. - **`TelegramChannel`**: `SpringLongPollingBot`; filters by `allowedUsername`; stores `chatId` for routing background replies. -- **`ChatChannel`**: WebSocket-first delivery (`setWsSession()`/`clearWsSession()`); falls back to buffering replies in `ConcurrentLinkedQueue` exposed via `drainPendingMessages()` REST endpoint when no WebSocket session is active. +- **`ChatChannel`**: WebSocket-first delivery (`setWsSession()`/`clearWsSession()`); falls back to buffering replies in `ConcurrentLinkedQueue` exposed via `drainPendingMessages()` REST endpoint when no WebSocket session is active. Web chat responses stream live: `Agent.respondTo(conversationId, question, ResponseListener)` uses `ChatClient.stream()` and reports each token via the listener callback (`onToken`/`onComplete`/`onError`); `ChatChannel` supplies a listener that pushes JSON frames (`chunk`, `done`, `error` — see `StreamFrameType`) to the browser over the WebSocket. `Channel` itself knows nothing about streaming; all other channels use the blocking `respondTo`/`sendMessage()` path. --- diff --git a/app/src/main/java/ai/javaclaw/chat/ChatChannel.java b/app/src/main/java/ai/javaclaw/chat/ChatChannel.java index 2050c312..c4578b53 100644 --- a/app/src/main/java/ai/javaclaw/chat/ChatChannel.java +++ b/app/src/main/java/ai/javaclaw/chat/ChatChannel.java @@ -1,6 +1,7 @@ package ai.javaclaw.chat; import ai.javaclaw.agent.Agent; +import ai.javaclaw.agent.ResponseListener; import ai.javaclaw.channels.Channel; import ai.javaclaw.channels.ChannelMessageReceivedEvent; import ai.javaclaw.channels.ChannelRegistry; @@ -16,10 +17,13 @@ import org.springframework.stereotype.Component; import org.springframework.web.socket.TextMessage; import org.springframework.web.socket.WebSocketSession; +import tools.jackson.databind.ObjectMapper; import java.io.IOException; import java.util.ArrayList; +import java.util.LinkedHashMap; import java.util.List; +import java.util.Map; import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.atomic.AtomicReference; @@ -38,13 +42,15 @@ public class ChatChannel implements Channel { private final Agent agent; private final ChannelRegistry channelRegistry; private final ChatMemoryRepository chatMemoryRepository; + private final ObjectMapper objectMapper; private final ConcurrentLinkedQueue pendingMessages = new ConcurrentLinkedQueue<>(); private final AtomicReference wsSession = new AtomicReference<>(); - public ChatChannel(Agent agent, ChannelRegistry channelRegistry, ChatMemoryRepository chatMemoryRepository) { + public ChatChannel(Agent agent, ChannelRegistry channelRegistry, ChatMemoryRepository chatMemoryRepository, ObjectMapper objectMapper) { this.agent = agent; this.channelRegistry = channelRegistry; this.chatMemoryRepository = chatMemoryRepository; + this.objectMapper = objectMapper; channelRegistry.registerChannel(this); log.info("Started Web Chat channel"); } @@ -93,6 +99,18 @@ public void sendMessage(String message) { } } + /** + * Delivers messages buffered while no WebSocket session was active. + * Each buffered message is attempted once; a failed push re-buffers it. + */ + public void flushPendingMessages() { + for (int i = pendingMessages.size(); i > 0; i--) { + String message = pendingMessages.poll(); + if (message == null) break; + sendMessage(message); + } + } + /** * Returns all known conversation IDs, always with "web" first. */ @@ -124,10 +142,46 @@ public List loadHistoryAsHtml(String conversationId) { /** * Handles a chat message from the web UI for the given conversationId. + * The response is streamed to the WebSocket session as JSON frames + * ({@code chunk}/{@code done}/{@code error}); the full response text is returned. */ public String chat(String conversationId, String message) { channelRegistry.publishMessageReceivedEvent(new ChannelMessageReceivedEvent(getName(), message)); - return agent.respondTo(conversationId, message); + + return agent.respondTo(conversationId, message, ResponseListener.of( + token -> sendChunkFrame(conversationId, token), + () -> sendDoneFrame(conversationId), + error -> sendErrorFrame(conversationId, error))); + } + + private void sendChunkFrame(String conversationId, String token) { + sendFrame(frame(StreamFrameType.CHUNK, conversationId, token)); + } + + private void sendDoneFrame(String conversationId) { + sendFrame(frame(StreamFrameType.DONE, conversationId, null)); + } + + private void sendErrorFrame(String conversationId, String error) { + sendFrame(frame(StreamFrameType.ERROR, conversationId, error == null ? "Unknown error" : error)); + } + + private static Map frame(StreamFrameType type, String conversationId, Object payload) { + Map frame = new LinkedHashMap<>(); + frame.put("type", type.type()); + if (payload != null) frame.put("data", payload); + frame.put("conversationId", conversationId); + return frame; + } + + private void sendFrame(Map frame) { + WebSocketSession session = wsSession.get(); + if (session == null || !session.isOpen()) return; + try { + session.sendMessage(new TextMessage(objectMapper.writeValueAsString(frame))); + } catch (IOException e) { + log.warn("WS push failed, dropping stream frame: {}", e.getMessage()); + } } private static String buildBackgroundMessageHtml(String text) { diff --git a/app/src/main/java/ai/javaclaw/chat/StreamFrameType.java b/app/src/main/java/ai/javaclaw/chat/StreamFrameType.java new file mode 100644 index 00000000..87a48023 --- /dev/null +++ b/app/src/main/java/ai/javaclaw/chat/StreamFrameType.java @@ -0,0 +1,21 @@ +package ai.javaclaw.chat; + +/** + * Type discriminator of the JSON frames streamed to the web chat WebSocket client. + * Every frame carries its payload (if any) in a uniform {@code data} field. + */ +public enum StreamFrameType { + CHUNK("chunk"), + DONE("done"), + ERROR("error"); + + private final String type; + + StreamFrameType(String type) { + this.type = type; + } + + public String type() { + return type; + } +} diff --git a/app/src/main/java/ai/javaclaw/chat/ws/ChatWebSocketHandler.java b/app/src/main/java/ai/javaclaw/chat/ws/ChatWebSocketHandler.java index 2132c7ac..d3947b50 100644 --- a/app/src/main/java/ai/javaclaw/chat/ws/ChatWebSocketHandler.java +++ b/app/src/main/java/ai/javaclaw/chat/ws/ChatWebSocketHandler.java @@ -45,6 +45,7 @@ public void afterConnectionEstablished(WebSocketSession session) throws Exceptio Htmx.oobInnerHtml("channel-selector", conversationSelector), Htmx.oobInnerHtml("chat-messages", bubbles), Htmx.oobInnerHtml("chat-input-area", inputArea)); + chatChannel.flushPendingMessages(); } @Override @@ -91,11 +92,10 @@ private void handleUserMessage(Map payload) throws Exception { Htmx.oobReplace("typing-indicator", ChatHtml.typingDots())); try { - // Call agent (blocking — background tasks may push messages via ChatChannel during this) - String response = chatChannel.chat(conversationId, userMessage); - chatChannel.sendHtml( - Htmx.oobAppend("chat-messages", ChatHtml.agentBubble(response)), - Htmx.oobReplace("typing-indicator", "")); + // Call agent (blocking — the response is streamed to the client as JSON frames + // by ChatChannel while this call runs; background tasks may push messages too) + chatChannel.chat(conversationId, userMessage); + chatChannel.sendHtml(Htmx.oobReplace("typing-indicator", "")); } catch (RuntimeException ex) { log.warn("Chat request failed for conversation {}", conversationId, ex); chatChannel.sendHtml( diff --git a/app/src/main/resources/templates/chat.html.peb b/app/src/main/resources/templates/chat.html.peb index 4728c84f..c28dc325 100644 --- a/app/src/main/resources/templates/chat.html.peb +++ b/app/src/main/resources/templates/chat.html.peb @@ -149,6 +149,11 @@ .chat-readonly-notice strong { color: rgba(180, 195, 235, .7); } + +.ar-msg--error .ar-msg__bubble { + border-color: color-mix(in srgb, var(--bulma-danger) 45%, transparent); + color: var(--bulma-danger); +} {% endblock %} @@ -259,9 +264,71 @@ if (ta) { ta.value = ''; ta.style.height = ''; } }); document.body.addEventListener('htmx:wsAfterMessage', function () { + scrollToBottom(); + }); + + // --- Streaming frames ------------------------------------------------ + // The server streams the agent response as JSON frames (type: chunk / + // done / error). These are intercepted here and never reach htmx; HTML + // frames (no JSON type) keep the htmx OOB-swap behavior unchanged. + var streamBubbles = {}; // conversationId -> in-progress bubble element + + document.body.addEventListener('htmx:wsBeforeMessage', function (e) { + var frame = parseStreamFrame(e.detail.message); + if (!frame) return; + e.preventDefault(); + handleStreamFrame(frame); + scrollToBottom(); + }); + + function parseStreamFrame(text) { + if (!text || text.charAt(0) !== '{') return null; + try { + var frame = JSON.parse(text); + return frame && typeof frame === 'object' && frame.type ? frame : null; + } catch (err) { + return null; + } + } + + function handleStreamFrame(frame) { + var id = frame.conversationId || 'web'; + if (frame.type === 'chunk') { + streamBubble(id).textContent += frame.data; + } else if (frame.type === 'done') { + clearTypingIndicator(); + delete streamBubbles[id]; + } else if (frame.type === 'error') { + var bubble = streamBubble(id); + bubble.textContent = frame.data; + bubble.closest('.ar-msg').classList.add('ar-msg--error'); + delete streamBubbles[id]; + } + } + + function streamBubble(conversationId) { + var bubble = streamBubbles[conversationId]; + if (!bubble || !bubble.isConnected) { + var article = document.createElement('article'); + article.className = 'ar-msg ar-msg--agent'; + article.innerHTML = '
JC
'; + document.getElementById('chat-messages').appendChild(article); + bubble = article.querySelector('.ar-msg__bubble'); + streamBubbles[conversationId] = bubble; + clearTypingIndicator(); + } + return bubble; + } + + function clearTypingIndicator() { + var typing = document.getElementById('typing-indicator'); + if (typing) typing.innerHTML = ''; + } + + function scrollToBottom() { var body = document.querySelector('.chat-body'); if (body) body.scrollTop = body.scrollHeight; - }); + } }()); {% endblock %} diff --git a/app/src/test/java/ai/javaclaw/chat/ChatChannelTest.java b/app/src/test/java/ai/javaclaw/chat/ChatChannelTest.java index 206bf424..7db87c16 100644 --- a/app/src/test/java/ai/javaclaw/chat/ChatChannelTest.java +++ b/app/src/test/java/ai/javaclaw/chat/ChatChannelTest.java @@ -1,6 +1,7 @@ package ai.javaclaw.chat; import ai.javaclaw.agent.Agent; +import ai.javaclaw.agent.ResponseListener; import ai.javaclaw.channels.ChannelRegistry; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -12,9 +13,11 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.web.socket.TextMessage; import org.springframework.web.socket.WebSocketSession; +import tools.jackson.databind.ObjectMapper; import java.io.IOException; import java.util.List; +import java.util.Map; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; @@ -34,7 +37,7 @@ class ChatChannelTest { @BeforeEach void setUp() { - chatChannel = new ChatChannel(agent, new ChannelRegistry(), chatMemoryRepository); + chatChannel = new ChatChannel(agent, new ChannelRegistry(), chatMemoryRepository, new ObjectMapper()); } // ----------------------------------------------------------------------- @@ -131,21 +134,21 @@ void loadHistoryUsesSuppliedConversationId() { @Test void chatDelegatesToAgentWithConversationId() { - when(agent.respondTo("web", "hello")).thenReturn("hi"); + when(agent.respondTo(eq("web"), eq("hello"), any(ResponseListener.class))).thenReturn("hi"); String response = chatChannel.chat("web", "hello"); assertThat(response).isEqualTo("hi"); - verify(agent).respondTo(eq("web"), eq("hello")); + verify(agent).respondTo(eq("web"), eq("hello"), any(ResponseListener.class)); } @Test void chatUsesSuppliedConversationId() { - when(agent.respondTo(eq("telegram-42"), any())).thenReturn("reply"); + when(agent.respondTo(eq("telegram-42"), any(), any(ResponseListener.class))).thenReturn("reply"); chatChannel.chat("telegram-42", "hello"); - verify(agent).respondTo(eq("telegram-42"), eq("hello")); + verify(agent).respondTo(eq("telegram-42"), eq("hello"), any(ResponseListener.class)); } // ----------------------------------------------------------------------- @@ -214,4 +217,106 @@ void sendMessagePushesOobHtmlToActiveSession() throws IOException { verify(session).sendMessage(any(TextMessage.class)); } + + @Test + void flushPendingMessagesDeliversMessagesBufferedWhileSendFailed() throws IOException { + WebSocketSession failingSession = mock(WebSocketSession.class); + when(failingSession.isOpen()).thenReturn(true); + org.mockito.Mockito.doThrow(new IOException("connection gone")).when(failingSession).sendMessage(any()); + chatChannel.setWsSession(failingSession); + chatChannel.sendMessage("Background result"); + + WebSocketSession session = openSession(); + chatChannel.flushPendingMessages(); + + verify(session).sendMessage(any(TextMessage.class)); + } + + @Test + void flushPendingMessagesDoesNothingWhenBufferIsEmpty() throws IOException { + WebSocketSession session = mock(WebSocketSession.class); + chatChannel.setWsSession(session); + + chatChannel.flushPendingMessages(); + + verify(session, never()).sendMessage(any()); + } + + // ----------------------------------------------------------------------- + // streaming frames + // ----------------------------------------------------------------------- + + @Test + void chatStreamsTokensAsChunkFramesFollowedByDoneFrame() throws IOException { + WebSocketSession session = openSession(); + agentStreams(listener -> { + listener.onToken("Hello "); + listener.onToken("world"); + listener.onComplete(); + }); + + chatChannel.chat("web", "hello"); + + List> frames = capturedFrames(session, 3); + assertThat(frames.get(0)) + .containsEntry("type", "chunk") + .containsEntry("data", "Hello ") + .containsEntry("conversationId", "web"); + assertThat(frames.get(1)) + .containsEntry("type", "chunk") + .containsEntry("data", "world"); + assertThat(frames.get(2)) + .containsEntry("type", "done") + .containsEntry("conversationId", "web") + .doesNotContainKey("data"); + } + + @Test + void chatStreamsErrorFrameWhenResponseFails() throws IOException { + WebSocketSession session = openSession(); + agentStreams(listener -> listener.onError("boom")); + + chatChannel.chat("web", "hello"); + + List> frames = capturedFrames(session, 1); + assertThat(frames.get(0)) + .containsEntry("type", "error") + .containsEntry("data", "boom") + .containsEntry("conversationId", "web"); + } + + @Test + void chatDropsStreamFramesWhenNoSessionIsActive() { + agentStreams(listener -> { + listener.onToken("Hello"); + listener.onComplete(); + }); + + // should not throw + chatChannel.chat("web", "hello"); + } + + private WebSocketSession openSession() { + WebSocketSession session = mock(WebSocketSession.class); + when(session.isOpen()).thenReturn(true); + chatChannel.setWsSession(session); + return session; + } + + private void agentStreams(java.util.function.Consumer progress) { + when(agent.respondTo(eq("web"), eq("hello"), any(ResponseListener.class))).thenAnswer(invocation -> { + progress.accept(invocation.getArgument(2)); + return ""; + }); + } + + @SuppressWarnings("unchecked") + private static List> capturedFrames(WebSocketSession session, int expectedCount) throws IOException { + var messageCaptor = org.mockito.ArgumentCaptor.forClass(TextMessage.class); + verify(session, org.mockito.Mockito.times(expectedCount)).sendMessage(messageCaptor.capture()); + ObjectMapper objectMapper = new ObjectMapper(); + return messageCaptor.getAllValues().stream() + .map(message -> (Map) objectMapper.readValue(message.getPayload(), Map.class)) + .toList(); + } } diff --git a/app/src/test/java/ai/javaclaw/chat/ws/ChatStreamingIntegrationTest.java b/app/src/test/java/ai/javaclaw/chat/ws/ChatStreamingIntegrationTest.java new file mode 100644 index 00000000..cf177b09 --- /dev/null +++ b/app/src/test/java/ai/javaclaw/chat/ws/ChatStreamingIntegrationTest.java @@ -0,0 +1,139 @@ +package ai.javaclaw.chat.ws; + +import ai.javaclaw.agent.Agent; +import ai.javaclaw.agent.ResponseListener; +import ai.javaclaw.channels.ChannelRegistry; +import ai.javaclaw.chat.ChatChannel; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.memory.ChatMemoryRepository; +import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository; +import org.springframework.boot.autoconfigure.ImportAutoConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.boot.tomcat.autoconfigure.servlet.TomcatServletWebServerAutoConfiguration; +import org.springframework.boot.webmvc.autoconfigure.DispatcherServletAutoConfiguration; +import org.springframework.boot.webmvc.autoconfigure.WebMvcAutoConfiguration; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.Import; +import org.springframework.web.socket.TextMessage; +import org.springframework.web.socket.WebSocketSession; +import org.springframework.web.socket.client.standard.StandardWebSocketClient; +import org.springframework.web.socket.handler.TextWebSocketHandler; +import tools.jackson.databind.ObjectMapper; + +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.concurrent.BlockingQueue; +import java.util.concurrent.LinkedBlockingQueue; +import java.util.concurrent.TimeUnit; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Full WebSocket round trip: a userMessage frame is sent to /ws/chat and the streamed + * response must arrive as multiple JSON chunk frames followed by a done frame. + */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT, + classes = ChatStreamingIntegrationTest.TestConfig.class) +class ChatStreamingIntegrationTest { + + @LocalServerPort + int port; + + final ObjectMapper objectMapper = new ObjectMapper(); + + @Test + void userMessageIsAnsweredWithMultipleChunkFramesBeforeDone() throws Exception { + BlockingQueue received = new LinkedBlockingQueue<>(); + WebSocketSession session = new StandardWebSocketClient() + .execute(new TextWebSocketHandler() { + @Override + protected void handleTextMessage(WebSocketSession session, TextMessage message) { + received.add(message.getPayload()); + } + }, "ws://localhost:" + port + "/ws/chat") + .get(5, TimeUnit.SECONDS); + + try { + session.sendMessage(new TextMessage(objectMapper.writeValueAsString(Map.of( + "type", "userMessage", + "conversationId", "web", + "message", "hi" + )))); + + List frameTypes = collectFrameTypesUntilDone(received); + + assertThat(frameTypes) + .containsSubsequence("chunk", "chunk", "done") + .last().isEqualTo("done"); + } finally { + session.close(); + } + } + + /** + * Collects the {@code type} of every JSON frame until the done frame arrives. + * Non-JSON payloads (htmx HTML fragments) are ignored. + */ + private List collectFrameTypesUntilDone(BlockingQueue received) throws Exception { + List frameTypes = new ArrayList<>(); + long deadline = System.currentTimeMillis() + 10_000; + while (System.currentTimeMillis() < deadline) { + String payload = received.poll(500, TimeUnit.MILLISECONDS); + if (payload == null || !payload.startsWith("{")) continue; + + @SuppressWarnings("unchecked") + Map frame = objectMapper.readValue(payload, Map.class); + frameTypes.add((String) frame.get("type")); + if ("done".equals(frame.get("type"))) break; + } + return frameTypes; + } + + @Configuration(proxyBeanMethods = false) + @ImportAutoConfiguration({TomcatServletWebServerAutoConfiguration.class, + DispatcherServletAutoConfiguration.class, WebMvcAutoConfiguration.class}) + @Import({WebSocketConfig.class, ChatWebSocketHandler.class, ChatChannel.class}) + static class TestConfig { + + @Bean + ObjectMapper objectMapper() { + return new ObjectMapper(); + } + + @Bean + ChannelRegistry channelRegistry() { + return new ChannelRegistry(); + } + + @Bean + ChatMemoryRepository chatMemoryRepository() { + return new InMemoryChatMemoryRepository(); + } + + @Bean + Agent agent() { + return new Agent() { + @Override + public String respondTo(String conversationId, String question) { + return "Hello world"; + } + + @Override + public String respondTo(String conversationId, String question, ResponseListener listener) { + listener.onToken("Hello "); + listener.onToken("world"); + listener.onComplete(); + return "Hello world"; + } + + @Override + public T prompt(String conversationId, String input, Class result) { + return null; + } + }; + } + } +} diff --git a/app/src/test/java/ai/javaclaw/chat/ws/ChatWebSocketHandlerTest.java b/app/src/test/java/ai/javaclaw/chat/ws/ChatWebSocketHandlerTest.java index c150440b..90902957 100644 --- a/app/src/test/java/ai/javaclaw/chat/ws/ChatWebSocketHandlerTest.java +++ b/app/src/test/java/ai/javaclaw/chat/ws/ChatWebSocketHandlerTest.java @@ -85,6 +85,34 @@ void handleUserMessageShowsGenericProviderErrorForUnexpectedFailures() throws Ex .contains("Details: boom"); } + @Test + void handleUserMessageClearsTypingIndicatorWithoutAppendingBubbleWhenResponseWasStreamed() throws Exception { + ChatChannel chatChannel = mock(ChatChannel.class); + WebSocketSession session = mock(WebSocketSession.class); + ChatWebSocketHandler handler = new ChatWebSocketHandler(chatChannel, new ObjectMapper()); + + // the response is streamed to the client by ChatChannel while chat() runs + when(chatChannel.chat("web", "hello")).thenReturn("streamed response"); + + handler.handleTextMessage(session, new TextMessage(new ObjectMapper().writeValueAsString(Map.of( + "type", "userMessage", + "conversationId", "web", + "message", "hello" + )))); + + ArgumentCaptor htmlCaptor = ArgumentCaptor.forClass(String[].class); + var inOrder = inOrder(chatChannel); + inOrder.verify(chatChannel).sendHtml(htmlCaptor.capture()); + inOrder.verify(chatChannel).chat("web", "hello"); + inOrder.verify(chatChannel).sendHtml(htmlCaptor.capture()); + verifyNoMoreInteractions(chatChannel); + + assertThat(String.join("", htmlCaptor.getAllValues().get(1))) + .contains("typing-indicator") + .doesNotContain("streamed response") + .doesNotContain("ar-msg--agent"); + } + @Test void handleChannelChangedSendsHistoryAndInputArea() throws Exception { ChatChannel chatChannel = mock(ChatChannel.class); diff --git a/base/src/main/java/ai/javaclaw/agent/Agent.java b/base/src/main/java/ai/javaclaw/agent/Agent.java index c01efa20..13836ce9 100644 --- a/base/src/main/java/ai/javaclaw/agent/Agent.java +++ b/base/src/main/java/ai/javaclaw/agent/Agent.java @@ -4,6 +4,13 @@ public interface Agent { String respondTo(String conversationId, String question); + default String respondTo(String conversationId, String question, ResponseListener listener) { + String response = respondTo(conversationId, question); + listener.onToken(response); + listener.onComplete(); + return response; + } + T prompt(String conversationId, String input, Class result); } diff --git a/base/src/main/java/ai/javaclaw/agent/DefaultAgent.java b/base/src/main/java/ai/javaclaw/agent/DefaultAgent.java index abb50103..e63b4761 100644 --- a/base/src/main/java/ai/javaclaw/agent/DefaultAgent.java +++ b/base/src/main/java/ai/javaclaw/agent/DefaultAgent.java @@ -1,5 +1,7 @@ package ai.javaclaw.agent; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.memory.ChatMemory; import org.springframework.stereotype.Component; @@ -7,6 +9,8 @@ @Component public class DefaultAgent implements Agent { + private static final Logger log = LoggerFactory.getLogger(DefaultAgent.class); + private final ChatClient chatClient; public DefaultAgent(ChatClient chatClient) { @@ -22,6 +26,35 @@ public String respondTo(String conversationId, String question) { .content(); } + @Override + public String respondTo(String conversationId, String question, ResponseListener listener) { + StringBuilder fullResponse = new StringBuilder(); + try { + chatClient + .prompt(question) + .advisors(a -> a.param(ChatMemory.CONVERSATION_ID, conversationId)) + .stream() + .content() + .doOnNext(token -> { + fullResponse.append(token); + listener.onToken(token); + }) + .blockLast(); + } catch (UnsupportedOperationException e) { + // The configured model cannot stream — fall back to the blocking call + String response = respondTo(conversationId, question); + listener.onToken(response); + listener.onComplete(); + return response; + } catch (RuntimeException e) { + log.warn("Streaming response failed for conversation {}", conversationId, e); + listener.onError(summarizeError(e)); + return fullResponse.toString(); + } + listener.onComplete(); + return fullResponse.toString(); + } + @Override public T prompt(String conversationId, String input, Class result) { return chatClient @@ -30,4 +63,9 @@ public T prompt(String conversationId, String input, Class result) { .call() .entity(result); } + + private static String summarizeError(Throwable ex) { + String message = ex.getMessage(); + return message == null || message.isBlank() ? ex.getClass().getSimpleName() : message; + } } diff --git a/base/src/main/java/ai/javaclaw/agent/ResponseListener.java b/base/src/main/java/ai/javaclaw/agent/ResponseListener.java new file mode 100644 index 00000000..b2d05659 --- /dev/null +++ b/base/src/main/java/ai/javaclaw/agent/ResponseListener.java @@ -0,0 +1,50 @@ +package ai.javaclaw.agent; + +import java.util.function.Consumer; + +/** + * Callback for observing an agent response as it is produced. Callers that can + * render progress incrementally (e.g. the web chat) supply an implementation; + * where the tokens go (WebSocket, SSE, console) is entirely the caller's concern. + */ +public interface ResponseListener { + + /** + * Called for every token as it arrives from the model. + */ + void onToken(String token); + + /** + * Called once after the last token; no further callbacks follow. + */ + void onComplete(); + + /** + * Called when the response fails; no further callbacks follow. Tokens already + * delivered via {@link #onToken(String)} may precede the failure. + */ + void onError(String message); + + /** + * Creates a listener from three lambdas, avoiding an anonymous class at the call site. + */ + static ResponseListener of(Consumer onToken, Runnable onComplete, Consumer onError) { + return new ResponseListener() { + + @Override + public void onToken(String token) { + onToken.accept(token); + } + + @Override + public void onComplete() { + onComplete.run(); + } + + @Override + public void onError(String message) { + onError.accept(message); + } + }; + } +} diff --git a/base/src/test/java/ai/javaclaw/agent/DefaultAgentTest.java b/base/src/test/java/ai/javaclaw/agent/DefaultAgentTest.java new file mode 100644 index 00000000..4652fe33 --- /dev/null +++ b/base/src/test/java/ai/javaclaw/agent/DefaultAgentTest.java @@ -0,0 +1,83 @@ +package ai.javaclaw.agent; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.ai.chat.client.ChatClient; +import reactor.core.publisher.Flux; + +import java.util.function.Consumer; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.inOrder; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +class DefaultAgentTest { + + @Mock ChatClient chatClient; + @Mock ChatClient.ChatClientRequestSpec requestSpec; + @Mock ChatClient.StreamResponseSpec streamSpec; + @Mock ChatClient.CallResponseSpec callSpec; + @Mock ResponseListener listener; + + DefaultAgent agent; + + @BeforeEach + @SuppressWarnings("unchecked") + void setUp() { + agent = new DefaultAgent(chatClient); + when(chatClient.prompt("hello")).thenReturn(requestSpec); + when(requestSpec.advisors(any(Consumer.class))).thenReturn(requestSpec); + when(requestSpec.stream()).thenReturn(streamSpec); + } + + @Test + void reportsEachTokenAndCompletionToListener() { + when(streamSpec.content()).thenReturn(Flux.just("Hello ", "world")); + + String response = agent.respondTo("web", "hello", listener); + + assertThat(response).isEqualTo("Hello world"); + var inOrder = inOrder(listener); + inOrder.verify(listener).onToken("Hello "); + inOrder.verify(listener).onToken("world"); + inOrder.verify(listener).onComplete(); + verify(listener, never()).onError(any()); + } + + @Test + void reportsErrorAndReturnsPartialResponseWhenStreamFails() { + when(streamSpec.content()).thenReturn(Flux.concat( + Flux.just("Hello "), + Flux.error(new RuntimeException("boom")))); + + String response = agent.respondTo("web", "hello", listener); + + assertThat(response).isEqualTo("Hello "); + var inOrder = inOrder(listener); + inOrder.verify(listener).onToken("Hello "); + inOrder.verify(listener).onError("boom"); + verify(listener, never()).onComplete(); + } + + @Test + void fallsBackToBlockingCallWhenModelDoesNotSupportStreaming() { + when(streamSpec.content()).thenReturn(Flux.error(new UnsupportedOperationException("streaming is not supported"))); + when(requestSpec.call()).thenReturn(callSpec); + when(callSpec.content()).thenReturn("full response"); + + String response = agent.respondTo("web", "hello", listener); + + assertThat(response).isEqualTo("full response"); + var inOrder = inOrder(listener); + inOrder.verify(listener).onToken("full response"); + inOrder.verify(listener).onComplete(); + verify(listener, never()).onError(any()); + } +}