|
15 | 15 | import java.util.Collections; |
16 | 16 | import java.util.concurrent.CountDownLatch; |
17 | 17 | import java.util.concurrent.TimeUnit; |
| 18 | +import java.util.concurrent.atomic.AtomicInteger; |
18 | 19 | import java.util.concurrent.atomic.AtomicReference; |
19 | 20 |
|
| 21 | +import okhttp3.extension.logging.HttpLogLevel; |
| 22 | + |
20 | 23 | import static org.junit.jupiter.api.Assertions.assertEquals; |
21 | 24 | import static org.junit.jupiter.api.Assertions.assertNotNull; |
22 | 25 | import static org.junit.jupiter.api.Assertions.assertTrue; |
@@ -54,6 +57,44 @@ void standardChatChunkDoesNotRequireNestedDataField() throws Exception { |
54 | 57 | } |
55 | 58 | } |
56 | 59 |
|
| 60 | + @Test |
| 61 | + void malformedJsonAndConsumerFailureRemainIsolated() throws Exception { |
| 62 | + OkHttpClient client = new OkHttpClient.Builder().addInterceptor(chain -> { |
| 63 | + String body = "data: invalid\\n\\n" |
| 64 | + + "data: {\\\"delta\\\":\\\"hello\\\"}\\n\\n" |
| 65 | + + "data: [DONE]\\n\\n"; |
| 66 | + return new Response.Builder() |
| 67 | + .request(chain.request()) |
| 68 | + .protocol(Protocol.HTTP_1_1) |
| 69 | + .code(200) |
| 70 | + .message("OK") |
| 71 | + .body(ResponseBody.create(body, MediaType.get("text/event-stream"))) |
| 72 | + .build(); |
| 73 | + }).build(); |
| 74 | + |
| 75 | + HermesHttpClientConfig config = new HermesHttpClientConfig(); |
| 76 | + config.markUnsafeBaseUrlOverriddenForTest(true); |
| 77 | + config.getDebug().setEnabled(true); |
| 78 | + config.getDebug().setLevel(HttpLogLevel.BODY); |
| 79 | + config.getDebug().setMaxContentLength(4); |
| 80 | + |
| 81 | + ChatRequest request = new ChatRequest(); |
| 82 | + request.setMessages(Collections.singletonList(new ChatRequest.Message("user", "hello"))); |
| 83 | + AtomicInteger consumerCalls = new AtomicInteger(); |
| 84 | + |
| 85 | + try (HermesSseClient sse = new HermesSseClient(config, null, client)) { |
| 86 | + CountDownLatch complete = new CountDownLatch(1); |
| 87 | + sse.subscribeChat(request, ignored -> { |
| 88 | + consumerCalls.incrementAndGet(); |
| 89 | + throw new IllegalStateException("consumer failure"); |
| 90 | + }, complete::countDown, ignored -> { }); |
| 91 | + |
| 92 | + assertTrue(complete.await(3, TimeUnit.SECONDS)); |
| 93 | + assertEquals(1, consumerCalls.get(), |
| 94 | + "malformed JSON must not be delivered to the business consumer"); |
| 95 | + } |
| 96 | + } |
| 97 | + |
57 | 98 | @Test |
58 | 99 | void sessionDisconnectDoesNotReplayOriginalPost() throws Exception { |
59 | 100 | try (MockWebServer server = new MockWebServer()) { |
|
0 commit comments