001/* 002 * Copyright 2022-2026 Revetware LLC. 003 * 004 * Licensed under the Apache License, Version 2.0 (the "License"); 005 * you may not use this file except in compliance with the License. 006 * You may obtain a copy of the License at 007 * 008 * http://www.apache.org/licenses/LICENSE-2.0 009 * 010 * Unless required by applicable law or agreed to in writing, software 011 * distributed under the License is distributed on an "AS IS" BASIS, 012 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. 013 * See the License for the specific language governing permissions and 014 * limitations under the License. 015 */ 016 017package com.soklet; 018 019import com.soklet.SseRequestResult.HandshakeAccepted; 020import com.soklet.SseRequestResult.HandshakeRejected; 021import com.soklet.annotation.SseEventSource; 022import com.soklet.internal.spring.LinkedCaseInsensitiveMap; 023import org.jspecify.annotations.NonNull; 024import org.jspecify.annotations.Nullable; 025 026import javax.annotation.concurrent.NotThreadSafe; 027import javax.annotation.concurrent.ThreadSafe; 028import java.io.BufferedReader; 029import java.io.ByteArrayOutputStream; 030import java.io.IOException; 031import java.io.InputStreamReader; 032import java.io.Reader; 033import java.lang.reflect.InvocationTargetException; 034import java.nio.ByteBuffer; 035import java.nio.CharBuffer; 036import java.nio.charset.Charset; 037import java.nio.charset.CharsetEncoder; 038import java.nio.charset.CharacterCodingException; 039import java.nio.charset.CoderResult; 040import java.nio.charset.StandardCharsets; 041import java.time.Duration; 042import java.time.Instant; 043import java.util.ArrayList; 044import java.util.Collections; 045import java.util.EnumSet; 046import java.util.HashMap; 047import java.util.LinkedHashMap; 048import java.util.List; 049import java.util.Locale; 050import java.util.Map; 051import java.util.Map.Entry; 052import java.util.Objects; 053import java.util.Optional; 054import java.util.Set; 055import java.util.concurrent.ConcurrentHashMap; 056import java.util.concurrent.CopyOnWriteArrayList; 057import java.util.concurrent.CopyOnWriteArraySet; 058import java.util.concurrent.CountDownLatch; 059import java.util.concurrent.Flow; 060import java.util.concurrent.TimeUnit; 061import java.util.concurrent.atomic.AtomicBoolean; 062import java.util.concurrent.atomic.AtomicReference; 063import java.util.concurrent.locks.ReentrantLock; 064import java.util.function.BiConsumer; 065import java.util.function.Consumer; 066import java.util.function.Function; 067import java.util.stream.Collectors; 068 069import static com.soklet.Utilities.emptyByteArray; 070import static com.soklet.Utilities.extractContentTypeFromHeaders; 071import static java.lang.String.format; 072import static java.util.Objects.requireNonNull; 073 074/** 075 * Soklet's main class - manages one or more configured transport servers ({@link HttpServer}, {@link SseServer}, and/or {@link McpServer}) 076 * using the provided system configuration. 077 * <p> 078 * <pre>{@code // Use out-of-the-box defaults 079 * SokletConfig config = SokletConfig.withHttpServer( 080 * HttpServer.fromPort(8080) 081 * ).build(); 082 * 083 * try (Soklet soklet = Soklet.fromConfig(config)) { 084 * soklet.start(); 085 * System.out.println("Soklet started, press [enter] to exit"); 086 * soklet.awaitShutdown(ShutdownTrigger.ENTER_KEY); 087 * }}</pre> 088 * <p> 089 * Soklet also offers an off-network {@link Simulator} concept via {@link #runSimulator(SokletConfig, Consumer)}, useful for integration testing. 090 * <p> 091 * Given a <em>Resource Method</em>... 092 * <pre>{@code public class HelloResource { 093 * @GET("/hello") 094 * public String hello(@QueryParameter String name) { 095 * return String.format("Hello, %s", name); 096 * } 097 * }}</pre> 098 * ...we might test it like this: 099 * <pre>{@code @Test 100 * public void integrationTest() { 101 * // Just use your app's existing configuration 102 * SokletConfig config = obtainMySokletConfig(); 103 * 104 * // Instead of running on a real HTTP server that listens on a port, 105 * // a non-network Simulator is provided against which you can 106 * // issue requests and receive responses. 107 * Soklet.runSimulator(config, (simulator -> { 108 * // Construct a request 109 * Request request = Request.withPath(HttpMethod.GET, "/hello") 110 * .queryParameters(Map.of("name", Set.of("Mark"))) 111 * .build(); 112 * 113 * // Perform the request and get a handle to the response 114 * HttpRequestResult result = simulator.performHttpRequest(request); 115 * MarshaledResponse marshaledResponse = result.getMarshaledResponse(); 116 * 117 * // Verify status code 118 * Integer expectedCode = 200; 119 * Integer actualCode = marshaledResponse.getStatusCode(); 120 * assertEquals(expectedCode, actualCode, "Bad status code"); 121 * 122 * // Verify response body 123 * marshaledResponse.getBody().ifPresentOrElse(body -> { 124 * String expectedBody = "Hello, Mark"; 125 * byte[] bytes = ((MarshaledResponseBody.Bytes) body).getBytes(); 126 * String actualBody = new String(bytes, StandardCharsets.UTF_8); 127 * assertEquals(expectedBody, actualBody, "Bad response body"); 128 * }, () -> { 129 * Assertions.fail("No response body"); 130 * }); 131 * })); 132 * }}</pre> 133 * <p> 134 * The {@link Simulator} also supports Server-Sent Events. 135 * <p> 136 * Integration testing documentation is available at <a href="https://www.soklet.com/docs/testing">https://www.soklet.com/docs/testing</a>. 137 * 138 * @author <a href="https://www.revetkn.com">Mark Allen</a> 139 */ 140@ThreadSafe 141public final class Soklet implements AutoCloseable { 142 @NonNull 143 private static final Map<@NonNull String, @NonNull Set<@NonNull String>> DEFAULT_ACCEPTED_HANDSHAKE_HEADERS; 144 145 static { 146 // Generally speaking, we always want these headers for SSE streaming responses. 147 // Users can override if they think necessary 148 LinkedCaseInsensitiveMap<Set<String>> defaultAcceptedHandshakeHeaders = new LinkedCaseInsensitiveMap<>(4); 149 defaultAcceptedHandshakeHeaders.put("Content-Type", Set.of("text/event-stream; charset=UTF-8")); 150 defaultAcceptedHandshakeHeaders.put("Cache-Control", Set.of("no-cache", "no-transform")); 151 defaultAcceptedHandshakeHeaders.put("Connection", Set.of("keep-alive")); 152 defaultAcceptedHandshakeHeaders.put("X-Accel-Buffering", Set.of("no")); 153 154 DEFAULT_ACCEPTED_HANDSHAKE_HEADERS = Collections.unmodifiableMap(defaultAcceptedHandshakeHeaders); 155 } 156 157 /** 158 * Acquires a Soklet instance with the given configuration. 159 * 160 * @param sokletConfig configuration that drives the Soklet system 161 * @return a Soklet instance 162 */ 163 @NonNull 164 public static Soklet fromConfig(@NonNull SokletConfig sokletConfig) { 165 requireNonNull(sokletConfig); 166 return new Soklet(sokletConfig); 167 } 168 169 @NonNull 170 private final SokletConfig sokletConfig; 171 @NonNull 172 private final ReentrantLock lock; 173 @NonNull 174 private final AtomicReference<@NonNull CountDownLatch> awaitShutdownLatchReference; 175 @NonNull 176 private final DefaultMcpRuntime defaultMcpRuntime; 177 178 /** 179 * Creates a Soklet instance with the given configuration. 180 * 181 * @param sokletConfig configuration that drives the Soklet system 182 */ 183 private Soklet(@NonNull SokletConfig sokletConfig) { 184 requireNonNull(sokletConfig); 185 186 this.sokletConfig = sokletConfig; 187 this.lock = new ReentrantLock(); 188 this.awaitShutdownLatchReference = new AtomicReference<>(new CountDownLatch(1)); 189 this.defaultMcpRuntime = new DefaultMcpRuntime(this); 190 191 sokletConfig.getMcpServer() 192 .map(McpServer::getSessionStore) 193 .filter(DefaultMcpSessionStore.class::isInstance) 194 .map(DefaultMcpSessionStore.class::cast) 195 .ifPresent(sessionStore -> sessionStore.pinnedSessionPredicate(this.defaultMcpRuntime::hasActiveStream)); 196 197 sokletConfig.getMcpServer() 198 .map(mcpServer -> mcpServer instanceof McpServerProxy mcpServerProxy ? mcpServerProxy.getRealImplementation() : mcpServer) 199 .filter(DefaultMcpServer.class::isInstance) 200 .map(DefaultMcpServer.class::cast) 201 .ifPresent(defaultMcpServer -> defaultMcpServer.mcpRuntime(this.defaultMcpRuntime)); 202 203 Set<ResourceMethod> resourceMethods = sokletConfig.getResourceMethodResolver().getResourceMethods(); 204 205 // Fail fast in the event that Soklet appears misconfigured 206 if (resourceMethods.size() == 0 207 && sokletConfig.getMcpServer().isEmpty()) 208 throw new IllegalStateException(format("No Soklet Resource Methods were found. First, try to rebuild and see if that solves the problem. If not, please ensure your %s is configured correctly. " 209 + "See https://www.soklet.com/docs/request-handling#resource-method-resolution for details.", ResourceMethodResolver.class.getSimpleName())); 210 211 boolean hasStandardHttpResourceMethods = resourceMethods.stream() 212 .anyMatch(resourceMethod -> !resourceMethod.isSseEventSource()); 213 214 if (hasStandardHttpResourceMethods && sokletConfig.getHttpServer().isEmpty()) 215 throw new IllegalStateException(format("Resource Methods were found, but no %s is configured. See https://www.soklet.com/docs/server-configuration for details.", 216 HttpServer.class.getSimpleName())); 217 218 // SSE misconfiguration check: @SseEventSource resource methods are declared, but not SseServer exists 219 boolean hasSseResourceMethods = resourceMethods.stream() 220 .anyMatch(resourceMethod -> resourceMethod.isSseEventSource()); 221 222 if (hasSseResourceMethods && sokletConfig.getSseServer().isEmpty()) 223 throw new IllegalStateException(format("Resource Methods annotated with @%s were found, but no %s is configured. See https://www.soklet.com/docs/server-sent-events for details.", 224 SseEventSource.class.getSimpleName(), SseServer.class.getSimpleName())); 225 226 MetricsCollector metricsCollector = sokletConfig.getMetricsCollector(); 227 228 if (metricsCollector instanceof DefaultMetricsCollector defaultMetricsCollector) { 229 try { 230 defaultMetricsCollector.initialize(sokletConfig); 231 } catch (Throwable t) { 232 sokletConfig.getAggregateLifecycleObserver().didReceiveLogEvent( 233 LogEvent.with(LogEventType.METRICS_COLLECTOR_FAILED, 234 format("An exception occurred while initializing %s", metricsCollector.getClass().getSimpleName())) 235 .throwable(t) 236 .build()); 237 } 238 } 239 240 // Use a layer of indirection here so the Soklet type does not need to directly implement the `RequestHandler` interface. 241 // Reasoning: the `handleRequest` method for Soklet should not be public, which might lead to accidental invocation by users. 242 // That method should only be called by the managed `HttpServer` instance. 243 Soklet soklet = this; 244 245 sokletConfig.getHttpServer().ifPresent(server -> server.initialize(getSokletConfig(), (request, marshaledResponseConsumer) -> { 246 // Delegate to Soklet's internal request handling method 247 soklet.handleRequest(request, ServerType.STANDARD_HTTP, marshaledResponseConsumer); 248 })); 249 250 SseServer sseServer = sokletConfig.getSseServer().orElse(null); 251 252 if (sseServer != null) 253 sseServer.initialize(sokletConfig, (request, marshaledResponseConsumer) -> { 254 // Delegate to Soklet's internal request handling method 255 soklet.handleRequest(request, ServerType.SSE, marshaledResponseConsumer); 256 }); 257 258 McpServer mcpServer = sokletConfig.getMcpServer().orElse(null); 259 260 if (mcpServer != null) 261 mcpServer.initialize(sokletConfig, soklet::handleMcpRequest); 262 } 263 264 /** 265 * Starts the managed server instance[s]. 266 * <p> 267 * If the managed server[s] are already started, this is a no-op. 268 */ 269 public void start() { 270 getLock().lock(); 271 272 try { 273 if (isStarted()) 274 return; 275 276 getAwaitShutdownLatchReference().set(new CountDownLatch(1)); 277 278 SokletConfig sokletConfig = getSokletConfig(); 279 LifecycleObserver lifecycleObserver = sokletConfig.getAggregateLifecycleObserver(); 280 281 // 1. Notify global intent to start 282 lifecycleObserver.willStartSoklet(this); 283 284 HttpServer httpServer = sokletConfig.getHttpServer().orElse(null); 285 SseServer sseServer = sokletConfig.getSseServer().orElse(null); 286 McpServer mcpServer = sokletConfig.getMcpServer().orElse(null); 287 boolean httpServerStarted = false; 288 boolean sseServerStarted = false; 289 boolean mcpServerStarted = false; 290 291 try { 292 // 2. Attempt to start Main HttpServer 293 if (httpServer != null) { 294 lifecycleObserver.willStartHttpServer(httpServer); 295 try { 296 httpServer.start(); 297 httpServerStarted = true; 298 lifecycleObserver.didStartHttpServer(httpServer); 299 } catch (Throwable t) { 300 lifecycleObserver.didFailToStartHttpServer(httpServer, t); 301 throw t; // Rethrow to trigger outer catch block 302 } 303 } 304 305 // 3. Attempt to start SSE HttpServer (if present) 306 if (sseServer != null) { 307 lifecycleObserver.willStartSseServer(sseServer); 308 try { 309 sseServer.start(); 310 sseServerStarted = true; 311 lifecycleObserver.didStartSseServer(sseServer); 312 } catch (Throwable t) { 313 lifecycleObserver.didFailToStartSseServer(sseServer, t); 314 throw t; // Rethrow to trigger outer catch block 315 } 316 } 317 318 if (mcpServer != null) { 319 lifecycleObserver.willStartMcpServer(mcpServer); 320 try { 321 mcpServer.start(); 322 mcpServerStarted = true; 323 lifecycleObserver.didStartMcpServer(mcpServer); 324 } catch (Throwable t) { 325 lifecycleObserver.didFailToStartMcpServer(mcpServer, t); 326 throw t; 327 } 328 } 329 330 // 4. Global success 331 lifecycleObserver.didStartSoklet(this); 332 } catch (Throwable t) { 333 rollbackStartedServersAfterFailedStart(lifecycleObserver, 334 httpServer, httpServerStarted, 335 sseServer, sseServerStarted, 336 mcpServer, mcpServerStarted, 337 t); 338 339 // 5. Global failure 340 lifecycleObserver.didFailToStartSoklet(this, t); 341 342 // Ensure the exception bubbles up so the application knows startup failed 343 if (t instanceof RuntimeException) 344 throw (RuntimeException) t; 345 346 throw new RuntimeException(t); 347 } 348 } finally { 349 getLock().unlock(); 350 } 351 } 352 353 private void rollbackStartedServersAfterFailedStart(@NonNull LifecycleObserver lifecycleObserver, 354 @Nullable HttpServer httpServer, 355 boolean httpServerStarted, 356 @Nullable SseServer sseServer, 357 boolean sseServerStarted, 358 @Nullable McpServer mcpServer, 359 boolean mcpServerStarted, 360 @NonNull Throwable startupFailure) { 361 requireNonNull(lifecycleObserver); 362 requireNonNull(startupFailure); 363 364 if (mcpServerStarted && mcpServer != null) 365 stopStartedMcpServerForRollback(lifecycleObserver, mcpServer, startupFailure); 366 367 if (sseServerStarted && sseServer != null) 368 stopStartedSseServerForRollback(lifecycleObserver, sseServer, startupFailure); 369 370 if (httpServerStarted && httpServer != null) 371 stopStartedHttpServerForRollback(lifecycleObserver, httpServer, startupFailure); 372 373 CountDownLatch awaitShutdownLatch = getAwaitShutdownLatchReference().get(); 374 375 if (awaitShutdownLatch != null) 376 awaitShutdownLatch.countDown(); 377 } 378 379 private void stopStartedHttpServerForRollback(@NonNull LifecycleObserver lifecycleObserver, 380 @NonNull HttpServer httpServer, 381 @NonNull Throwable startupFailure) { 382 requireNonNull(lifecycleObserver); 383 requireNonNull(httpServer); 384 requireNonNull(startupFailure); 385 386 try { 387 lifecycleObserver.willStopHttpServer(httpServer); 388 } catch (Throwable t) { 389 startupFailure.addSuppressed(t); 390 } 391 392 try { 393 httpServer.stop(); 394 try { 395 lifecycleObserver.didStopHttpServer(httpServer); 396 } catch (Throwable t) { 397 startupFailure.addSuppressed(t); 398 } 399 } catch (Throwable t) { 400 startupFailure.addSuppressed(t); 401 402 try { 403 lifecycleObserver.didFailToStopHttpServer(httpServer, t); 404 } catch (Throwable t2) { 405 startupFailure.addSuppressed(t2); 406 } 407 } 408 } 409 410 private void stopStartedSseServerForRollback(@NonNull LifecycleObserver lifecycleObserver, 411 @NonNull SseServer sseServer, 412 @NonNull Throwable startupFailure) { 413 requireNonNull(lifecycleObserver); 414 requireNonNull(sseServer); 415 requireNonNull(startupFailure); 416 417 try { 418 lifecycleObserver.willStopSseServer(sseServer); 419 } catch (Throwable t) { 420 startupFailure.addSuppressed(t); 421 } 422 423 try { 424 sseServer.stop(); 425 try { 426 lifecycleObserver.didStopSseServer(sseServer); 427 } catch (Throwable t) { 428 startupFailure.addSuppressed(t); 429 } 430 } catch (Throwable t) { 431 startupFailure.addSuppressed(t); 432 433 try { 434 lifecycleObserver.didFailToStopSseServer(sseServer, t); 435 } catch (Throwable t2) { 436 startupFailure.addSuppressed(t2); 437 } 438 } 439 } 440 441 private void stopStartedMcpServerForRollback(@NonNull LifecycleObserver lifecycleObserver, 442 @NonNull McpServer mcpServer, 443 @NonNull Throwable startupFailure) { 444 requireNonNull(lifecycleObserver); 445 requireNonNull(mcpServer); 446 requireNonNull(startupFailure); 447 448 try { 449 lifecycleObserver.willStopMcpServer(mcpServer); 450 } catch (Throwable t) { 451 startupFailure.addSuppressed(t); 452 } 453 454 try { 455 mcpServer.stop(); 456 try { 457 lifecycleObserver.didStopMcpServer(mcpServer); 458 } catch (Throwable t) { 459 startupFailure.addSuppressed(t); 460 } 461 } catch (Throwable t) { 462 startupFailure.addSuppressed(t); 463 464 try { 465 lifecycleObserver.didFailToStopMcpServer(mcpServer, t); 466 } catch (Throwable t2) { 467 startupFailure.addSuppressed(t2); 468 } 469 } 470 } 471 472 /** 473 * Stops the managed server instance[s]. 474 * <p> 475 * If the managed server[s] are already stopped, this is a no-op. 476 */ 477 public void stop() { 478 getLock().lock(); 479 480 try { 481 if (isStarted()) { 482 SokletConfig sokletConfig = getSokletConfig(); 483 LifecycleObserver lifecycleObserver = sokletConfig.getAggregateLifecycleObserver(); 484 485 // 1. Notify global intent to stop 486 lifecycleObserver.willStopSoklet(this); 487 488 Throwable firstEncounteredException = null; 489 HttpServer httpServer = sokletConfig.getHttpServer().orElse(null); 490 491 // 2. Attempt to stop Main HttpServer 492 if (httpServer != null && httpServer.isStarted()) { 493 lifecycleObserver.willStopHttpServer(httpServer); 494 try { 495 httpServer.stop(); 496 lifecycleObserver.didStopHttpServer(httpServer); 497 } catch (Throwable t) { 498 firstEncounteredException = t; 499 lifecycleObserver.didFailToStopHttpServer(httpServer, t); 500 } 501 } 502 503 // 3. Attempt to stop SSE HttpServer 504 SseServer sseServer = sokletConfig.getSseServer().orElse(null); 505 506 if (sseServer != null && sseServer.isStarted()) { 507 lifecycleObserver.willStopSseServer(sseServer); 508 try { 509 sseServer.stop(); 510 lifecycleObserver.didStopSseServer(sseServer); 511 } catch (Throwable t) { 512 if (firstEncounteredException == null) 513 firstEncounteredException = t; 514 515 lifecycleObserver.didFailToStopSseServer(sseServer, t); 516 } 517 } 518 519 McpServer mcpServer = sokletConfig.getMcpServer().orElse(null); 520 521 if (mcpServer != null && mcpServer.isStarted()) { 522 lifecycleObserver.willStopMcpServer(mcpServer); 523 try { 524 mcpServer.stop(); 525 lifecycleObserver.didStopMcpServer(mcpServer); 526 } catch (Throwable t) { 527 if (firstEncounteredException == null) 528 firstEncounteredException = t; 529 530 lifecycleObserver.didFailToStopMcpServer(mcpServer, t); 531 } 532 } 533 534 // 4. Global completion (Success or Failure) 535 if (firstEncounteredException == null) 536 lifecycleObserver.didStopSoklet(this); 537 else 538 lifecycleObserver.didFailToStopSoklet(this, firstEncounteredException); 539 } 540 } finally { 541 try { 542 requireNonNull(getAwaitShutdownLatchReference().get()).countDown(); 543 } finally { 544 getLock().unlock(); 545 } 546 } 547 } 548 549 /** 550 * Blocks the current thread until JVM shutdown ({@code SIGTERM/SIGINT/System.exit(...)} and so forth), <strong>or</strong> if one of the provided {@code shutdownTriggers} occurs. 551 * <p> 552 * This method will automatically invoke this instance's {@link #stop()} method once it becomes unblocked. 553 * <p> 554 * <strong>Notes regarding {@link ShutdownTrigger#ENTER_KEY}:</strong> 555 * <ul> 556 * <li>It will invoke {@link #stop()} on <i>all</i> Soklet instances, as stdin is process-wide</li> 557 * <li>It requires usable standard input. If stdin is unavailable or reaches EOF before a keypress (e.g. running in a Docker container), Soklet will fire {@link LifecycleObserver#didReceiveLogEvent(LogEvent)} with an event of type {@link LogEventType#CONFIGURATION_UNSUPPORTED} and keep running</li> 558 * </ul> 559 * 560 * @param shutdownTriggers additional trigger[s] which signal that shutdown should occur, e.g. {@link ShutdownTrigger#ENTER_KEY} for "enter key pressed" 561 * @throws InterruptedException if the current thread has its interrupted status set on entry to this method, or is interrupted while waiting 562 */ 563 public void awaitShutdown(@Nullable ShutdownTrigger... shutdownTriggers) throws InterruptedException { 564 Thread shutdownHook = null; 565 boolean registeredEnterKeyShutdownTrigger = false; 566 Set<ShutdownTrigger> shutdownTriggersAsSet = shutdownTriggers == null || shutdownTriggers.length == 0 ? Set.of() : EnumSet.copyOf(Set.of(shutdownTriggers)); 567 568 try { 569 // Optionally listen for enter key 570 if (shutdownTriggersAsSet.contains(ShutdownTrigger.ENTER_KEY)) { 571 registeredEnterKeyShutdownTrigger = KeypressManager.register(this); // returns false if stdin unusable/disabled 572 573 if (!registeredEnterKeyShutdownTrigger) { 574 LogEvent logEvent = LogEvent.with( 575 LogEventType.CONFIGURATION_UNSUPPORTED, 576 format("Ignoring request for %s.%s - it is unsupported in this environment (stdin is unavailable)", ShutdownTrigger.class.getSimpleName(), ShutdownTrigger.ENTER_KEY.name()) 577 ).build(); 578 579 getSokletConfig().getAggregateLifecycleObserver().didReceiveLogEvent(logEvent); 580 } 581 } 582 583 // Always register a shutdown hook 584 shutdownHook = new Thread(() -> { 585 try { 586 stop(); 587 } catch (Throwable ignored) { 588 // Nothing to do 589 } 590 }, "soklet-shutdown-hook"); 591 592 Runtime.getRuntime().addShutdownHook(shutdownHook); 593 594 // Wait until "close" finishes 595 requireNonNull(getAwaitShutdownLatchReference().get()).await(); 596 } finally { 597 if (registeredEnterKeyShutdownTrigger) 598 KeypressManager.unregister(this); 599 600 try { 601 Runtime.getRuntime().removeShutdownHook(shutdownHook); 602 } catch (IllegalStateException ignored) { 603 // JVM shutting down 604 } 605 } 606 } 607 608 /** 609 * Handles "awaitShutdown" for {@link ShutdownTrigger#ENTER_KEY} by listening to stdin - all Soklet instances are terminated on keypress. 610 */ 611 @ThreadSafe 612 static final class KeypressManager { 613 @NonNull 614 private static final Set<@NonNull Soklet> SOKLET_REGISTRY; 615 @NonNull 616 private static final AtomicBoolean LISTENER_STARTED; 617 @NonNull 618 private static final AtomicReference<@Nullable Boolean> INTERACTIVE_CONSOLE_AVAILABLE_OVERRIDE; 619 620 static { 621 SOKLET_REGISTRY = new CopyOnWriteArraySet<>(); 622 LISTENER_STARTED = new AtomicBoolean(false); 623 INTERACTIVE_CONSOLE_AVAILABLE_OVERRIDE = new AtomicReference<>(); 624 } 625 626 /** 627 * Register a Soklet for Enter-to-stop support. Returns true iff a listener is (or was already) active. 628 * If System.in is not usable (or disabled), returns false and does nothing. 629 */ 630 @NonNull 631 synchronized static Boolean register(@NonNull Soklet soklet) { 632 requireNonNull(soklet); 633 634 // If stdin is unavailable, don't start a listener. 635 if (!canReadFromStdin()) 636 return false; 637 638 SOKLET_REGISTRY.add(soklet); 639 640 // Start a single process-wide listener once. 641 if (LISTENER_STARTED.compareAndSet(false, true)) { 642 Thread thread = new Thread(KeypressManager::runLoop, "soklet-keypress-shutdown-listener"); 643 thread.setDaemon(true); // never block JVM exit 644 thread.start(); 645 } 646 647 return true; 648 } 649 650 static void interactiveConsoleAvailableOverride(@Nullable Boolean interactiveConsoleAvailableOverride) { 651 INTERACTIVE_CONSOLE_AVAILABLE_OVERRIDE.set(interactiveConsoleAvailableOverride); 652 } 653 654 static boolean isListenerStarted() { 655 return LISTENER_STARTED.get(); 656 } 657 658 synchronized static void reset() { 659 SOKLET_REGISTRY.clear(); 660 LISTENER_STARTED.set(false); 661 INTERACTIVE_CONSOLE_AVAILABLE_OVERRIDE.set(null); 662 } 663 664 synchronized static void unregister(@NonNull Soklet soklet) { 665 SOKLET_REGISTRY.remove(soklet); 666 // We intentionally keep the listener alive; it's daemon and cheap. 667 // If stdin hits EOF, the listener exits on its own. 668 } 669 670 /** 671 * ENTER_KEY shutdown only requires readable stdin. Environments such as IntelliJ often provide stdin without 672 * a {@link System#console()}; EOF is handled by the listener without stopping the server. 673 */ 674 @NonNull 675 private static Boolean canReadFromStdin() { 676 if (System.in == null) 677 return false; 678 679 Boolean interactiveConsoleAvailableOverride = INTERACTIVE_CONSOLE_AVAILABLE_OVERRIDE.get(); 680 681 if (interactiveConsoleAvailableOverride != null) 682 return interactiveConsoleAvailableOverride; 683 684 return true; 685 } 686 687 /** 688 * Single blocking read on stdin. On any line, stop all registered servers. EOF means stdin is 689 * unusable for interactive shutdown and must not stop the process. 690 */ 691 private static void runLoop() { 692 try { 693 BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(System.in, StandardCharsets.UTF_8)); 694 String line = bufferedReader.readLine(); 695 696 if (line == null) { 697 logEnterKeyUnsupported("Ignoring request for ENTER_KEY shutdown - stdin reached EOF before an enter keypress"); 698 return; 699 } 700 701 stopAllSoklets(); 702 } catch (Throwable ignored) { 703 logEnterKeyUnsupported("Ignoring request for ENTER_KEY shutdown - stdin became unusable before an enter keypress"); 704 } finally { 705 LISTENER_STARTED.set(false); 706 } 707 } 708 709 synchronized private static void stopAllSoklets() { 710 for (Soklet soklet : SOKLET_REGISTRY) { 711 try { 712 soklet.stop(); 713 } catch (Throwable ignored) { 714 // Nothing to do 715 } 716 } 717 } 718 719 private static void logEnterKeyUnsupported(@NonNull String message) { 720 requireNonNull(message); 721 722 LogEvent logEvent = LogEvent.with(LogEventType.CONFIGURATION_UNSUPPORTED, message).build(); 723 724 for (Soklet soklet : SOKLET_REGISTRY) { 725 try { 726 soklet.getSokletConfig().getAggregateLifecycleObserver().didReceiveLogEvent(logEvent); 727 } catch (Throwable ignored) { 728 // Nothing to do 729 } 730 } 731 } 732 733 private KeypressManager() {} 734 } 735 736 /** 737 * Nonpublic "informal" implementation of {@link com.soklet.HttpServer.RequestHandler} so Soklet does not need to expose {@code handleRequest} publicly. 738 * Reasoning: users of this library should never call {@code handleRequest} directly - it should only be invoked in response to events 739 * provided by a {@link HttpServer} or {@link SseServer} implementation. 740 */ 741 protected void handleRequest(@NonNull Request request, 742 @NonNull ServerType serverType, 743 @NonNull Consumer<HttpRequestResult> requestResultConsumer) { 744 requireNonNull(request); 745 requireNonNull(serverType); 746 requireNonNull(requestResultConsumer); 747 748 Instant processingStarted = Instant.now(); 749 750 SokletConfig sokletConfig = getSokletConfig(); 751 ResourceMethodResolver resourceMethodResolver = sokletConfig.getResourceMethodResolver(); 752 ResponseMarshaler responseMarshaler = sokletConfig.getResponseMarshaler(); 753 LifecycleObserver lifecycleObserver = sokletConfig.getAggregateLifecycleObserver(); 754 RequestInterceptor requestInterceptor = sokletConfig.getRequestInterceptor(); 755 MetricsCollector metricsCollector = sokletConfig.getMetricsCollector(); 756 757 // Holders to permit mutable effectively-final variables 758 AtomicReference<MarshaledResponse> marshaledResponseHolder = new AtomicReference<>(); 759 AtomicReference<Throwable> resourceMethodResolutionExceptionHolder = new AtomicReference<>(); 760 AtomicReference<Request> requestHolder = new AtomicReference<>(request); 761 AtomicReference<ResourceMethod> resourceMethodHolder = new AtomicReference<>(); 762 AtomicReference<HttpRequestResult> requestResultHolder = new AtomicReference<>(); 763 764 // Holders to permit mutable effectively-final state tracking 765 AtomicBoolean willStartResponseWritingCompleted = new AtomicBoolean(false); 766 AtomicBoolean didFinishResponseWritingCompleted = new AtomicBoolean(false); 767 AtomicBoolean didFinishRequestHandlingCompleted = new AtomicBoolean(false); 768 AtomicBoolean didInvokeWrapRequestConsumer = new AtomicBoolean(false); 769 770 List<Throwable> throwables = new ArrayList<>(10); 771 772 Consumer<LogEvent> safelyLog = (logEvent -> { 773 try { 774 lifecycleObserver.didReceiveLogEvent(logEvent); 775 } catch (Throwable throwable) { 776 // The LifecycleObserver implementation errored out, but we can't let that affect us. 777 throwables.add(throwable); 778 } 779 }); 780 781 BiConsumer<String, Consumer<MetricsCollector>> safelyCollectMetrics = (message, metricsInvocation) -> { 782 if (metricsCollector == null) 783 return; 784 785 try { 786 metricsInvocation.accept(metricsCollector); 787 } catch (Throwable throwable) { 788 safelyLog.accept(LogEvent.with(LogEventType.METRICS_COLLECTOR_FAILED, message) 789 .throwable(throwable) 790 .request(requestHolder.get()) 791 .resourceMethod(resourceMethodHolder.get()) 792 .marshaledResponse(marshaledResponseHolder.get()) 793 .build()); 794 } 795 }; 796 797 requestHolder.set(request); 798 799 try { 800 requestInterceptor.wrapRequest(serverType, request, (wrappedRequest) -> { 801 didInvokeWrapRequestConsumer.set(true); 802 requestHolder.set(wrappedRequest); 803 804 try { 805 // Resolve after wrapping so path/method rewrites affect routing. 806 resourceMethodHolder.set(resourceMethodResolver.resourceMethodForRequest(requestHolder.get(), serverType).orElse(null)); 807 resourceMethodResolutionExceptionHolder.set(null); 808 } catch (Throwable t) { 809 safelyLog.accept(LogEvent.with(LogEventType.RESOURCE_METHOD_RESOLUTION_FAILED, "Unable to resolve Resource Method") 810 .throwable(t) 811 .request(requestHolder.get()) 812 .build()); 813 814 // If an exception occurs here, keep track of it - we will surface them after letting LifecycleObserver 815 // see that a request has come in. 816 throwables.add(t); 817 resourceMethodResolutionExceptionHolder.set(t); 818 resourceMethodHolder.set(null); 819 } 820 821 try { 822 lifecycleObserver.didStartRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get()); 823 } catch (Throwable t) { 824 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_START_REQUEST_HANDLING_FAILED, 825 format("An exception occurred while invoking %s::didStartRequestHandling", 826 LifecycleObserver.class.getSimpleName())) 827 .throwable(t) 828 .request(requestHolder.get()) 829 .resourceMethod(resourceMethodHolder.get()) 830 .build()); 831 832 throwables.add(t); 833 } 834 835 safelyCollectMetrics.accept( 836 format("An exception occurred while invoking %s::didStartRequestHandling", MetricsCollector.class.getSimpleName()), 837 (metricsInvocation) -> metricsInvocation.didStartRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get())); 838 839 try { 840 AtomicBoolean didInvokeMarshaledResponseConsumer = new AtomicBoolean(false); 841 842 requestInterceptor.interceptRequest(serverType, requestHolder.get(), resourceMethodHolder.get(), (interceptorRequest) -> { 843 requestHolder.set(interceptorRequest); 844 845 try { 846 if (resourceMethodResolutionExceptionHolder.get() != null) 847 throw resourceMethodResolutionExceptionHolder.get(); 848 849 HttpRequestResult requestResult = toHttpRequestResult(requestHolder.get(), resourceMethodHolder.get(), serverType); 850 requestResultHolder.set(requestResult); 851 852 MarshaledResponse originalMarshaledResponse = requestResult.getMarshaledResponse(); 853 MarshaledResponse updatedMarshaledResponse = requestResult.getMarshaledResponse(); 854 855 // A few special cases that are "global" in that they can affect all requests and 856 // need to happen after marshaling the response... 857 858 // 1. Customize response for HEAD (e.g. remove body, set Content-Length header) 859 updatedMarshaledResponse = applyHeadResponseIfApplicable(requestHolder.get(), updatedMarshaledResponse); 860 861 // 2. Apply other standard response customizations (CORS, Content-Length) 862 // Note that we don't want to write Content-Length for SSE "accepted" handshakes 863 SseHandshakeResult sseHandshakeResult = requestResult.getSseHandshakeResult().orElse(null); 864 boolean suppressContentLength = sseHandshakeResult != null && sseHandshakeResult instanceof SseHandshakeResult.Accepted; 865 866 updatedMarshaledResponse = applyCommonPropertiesToMarshaledResponse(requestHolder.get(), updatedMarshaledResponse, suppressContentLength); 867 868 // Update our result holder with the modified response if necessary 869 if (originalMarshaledResponse != updatedMarshaledResponse) { 870 marshaledResponseHolder.set(updatedMarshaledResponse); 871 requestResultHolder.set(requestResult.copy() 872 .marshaledResponse(updatedMarshaledResponse) 873 .finish()); 874 } 875 876 return updatedMarshaledResponse; 877 } catch (Throwable t) { 878 if (t != resourceMethodResolutionExceptionHolder.get()) { 879 throwables.add(t); 880 881 safelyLog.accept(LogEvent.with(LogEventType.REQUEST_PROCESSING_FAILED, 882 "An exception occurred while processing request") 883 .throwable(t) 884 .request(requestHolder.get()) 885 .resourceMethod(resourceMethodHolder.get()) 886 .build()); 887 } 888 889 // Unhappy path. Try to use configuration's exception response marshaler... 890 try { 891 MarshaledResponse marshaledResponse = responseMarshaler.forThrowable(requestHolder.get(), t, resourceMethodHolder.get()); 892 marshaledResponse = applyCommonPropertiesToMarshaledResponse(requestHolder.get(), marshaledResponse); 893 marshaledResponseHolder.set(marshaledResponse); 894 895 return marshaledResponse; 896 } catch (Throwable t2) { 897 throwables.add(t2); 898 899 safelyLog.accept(LogEvent.with(LogEventType.RESPONSE_MARSHALER_FOR_THROWABLE_FAILED, 900 format("An exception occurred while trying to write an exception response for %s", t)) 901 .throwable(t2) 902 .request(requestHolder.get()) 903 .resourceMethod(resourceMethodHolder.get()) 904 .build()); 905 906 // The configuration's exception response marshaler failed - provide a failsafe response to recover 907 return provideFailsafeMarshaledResponse(requestHolder.get(), t2); 908 } 909 } 910 }, (interceptorMarshaledResponse) -> { 911 requireNonNull(interceptorMarshaledResponse); 912 didInvokeMarshaledResponseConsumer.set(true); 913 marshaledResponseHolder.set(interceptorMarshaledResponse); 914 }); 915 916 if (!didInvokeMarshaledResponseConsumer.get()) { 917 requestResultHolder.set(null); 918 throw new IllegalStateException(format("%s::interceptRequest must call responseWriter", RequestInterceptor.class.getSimpleName())); 919 } 920 } catch (Throwable t) { 921 throwables.add(t); 922 923 try { 924 // In the event that an error occurs during processing of a RequestInterceptor method, for example 925 safelyLog.accept(LogEvent.with(LogEventType.REQUEST_INTERCEPTOR_INTERCEPT_REQUEST_FAILED, 926 format("An exception occurred while invoking %s::interceptRequest", RequestInterceptor.class.getSimpleName())) 927 .throwable(t) 928 .request(requestHolder.get()) 929 .resourceMethod(resourceMethodHolder.get()) 930 .build()); 931 932 MarshaledResponse marshaledResponse = responseMarshaler.forThrowable(requestHolder.get(), t, resourceMethodHolder.get()); 933 marshaledResponse = applyCommonPropertiesToMarshaledResponse(requestHolder.get(), marshaledResponse); 934 marshaledResponseHolder.set(marshaledResponse); 935 } catch (Throwable t2) { 936 throwables.add(t2); 937 938 safelyLog.accept(LogEvent.with(LogEventType.RESPONSE_MARSHALER_FOR_THROWABLE_FAILED, 939 format("An exception occurred while invoking %s::forThrowable when trying to write an exception response for %s", ResponseMarshaler.class.getSimpleName(), t)) 940 .throwable(t2) 941 .request(requestHolder.get()) 942 .resourceMethod(resourceMethodHolder.get()) 943 .build()); 944 945 marshaledResponseHolder.set(provideFailsafeMarshaledResponse(requestHolder.get(), t2)); 946 } 947 } finally { 948 try { 949 try { 950 lifecycleObserver.willWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get()); 951 } finally { 952 willStartResponseWritingCompleted.set(true); 953 } 954 955 safelyCollectMetrics.accept( 956 format("An exception occurred while invoking %s::willWriteResponse", MetricsCollector.class.getSimpleName()), 957 (metricsInvocation) -> metricsInvocation.willWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get())); 958 959 Instant responseWriteStarted = Instant.now(); 960 961 try { 962 HttpRequestResult requestResult = requestResultHolder.get(); 963 964 if (requestResult != null) 965 requestResultConsumer.accept(requestResult); 966 else 967 requestResultConsumer.accept(HttpRequestResult.withMarshaledResponse(marshaledResponseHolder.get()) 968 .resourceMethod(resourceMethodHolder.get()) 969 .build()); 970 971 Instant responseWriteFinished = Instant.now(); 972 Duration responseWriteDuration = Duration.between(responseWriteStarted, responseWriteFinished); 973 974 try { 975 lifecycleObserver.didWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), responseWriteDuration); 976 } catch (Throwable t) { 977 throwables.add(t); 978 979 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_WRITE_RESPONSE_FAILED, 980 format("An exception occurred while invoking %s::didWriteResponse", 981 LifecycleObserver.class.getSimpleName())) 982 .throwable(t) 983 .request(requestHolder.get()) 984 .resourceMethod(resourceMethodHolder.get()) 985 .marshaledResponse(marshaledResponseHolder.get()) 986 .build()); 987 } finally { 988 didFinishResponseWritingCompleted.set(true); 989 } 990 991 safelyCollectMetrics.accept( 992 format("An exception occurred while invoking %s::didWriteResponse", MetricsCollector.class.getSimpleName()), 993 (metricsInvocation) -> metricsInvocation.didWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), 994 marshaledResponseHolder.get(), responseWriteDuration)); 995 } catch (Throwable t) { 996 throwables.add(t); 997 998 Instant responseWriteFinished = Instant.now(); 999 Duration responseWriteDuration = Duration.between(responseWriteStarted, responseWriteFinished); 1000 1001 try { 1002 lifecycleObserver.didFailToWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), responseWriteDuration, t); 1003 } catch (Throwable t2) { 1004 throwables.add(t2); 1005 1006 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_WRITE_RESPONSE_FAILED, 1007 format("An exception occurred while invoking %s::didFailToWriteResponse", 1008 LifecycleObserver.class.getSimpleName())) 1009 .throwable(t2) 1010 .request(requestHolder.get()) 1011 .resourceMethod(resourceMethodHolder.get()) 1012 .marshaledResponse(marshaledResponseHolder.get()) 1013 .build()); 1014 } 1015 1016 safelyCollectMetrics.accept( 1017 format("An exception occurred while invoking %s::didFailToWriteResponse", MetricsCollector.class.getSimpleName()), 1018 (metricsInvocation) -> metricsInvocation.didFailToWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), 1019 marshaledResponseHolder.get(), responseWriteDuration, t)); 1020 } 1021 } finally { 1022 Duration processingDuration = Duration.between(processingStarted, Instant.now()); 1023 1024 safelyCollectMetrics.accept( 1025 format("An exception occurred while invoking %s::didFinishRequestHandling", MetricsCollector.class.getSimpleName()), 1026 (metricsInvocation) -> metricsInvocation.didFinishRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), processingDuration, Collections.unmodifiableList(throwables))); 1027 1028 try { 1029 lifecycleObserver.didFinishRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), processingDuration, Collections.unmodifiableList(throwables)); 1030 } catch (Throwable t) { 1031 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_FINISH_REQUEST_HANDLING_FAILED, 1032 format("An exception occurred while invoking %s::didFinishRequestHandling", 1033 LifecycleObserver.class.getSimpleName())) 1034 .throwable(t) 1035 .request(requestHolder.get()) 1036 .resourceMethod(resourceMethodHolder.get()) 1037 .marshaledResponse(marshaledResponseHolder.get()) 1038 .build()); 1039 } finally { 1040 didFinishRequestHandlingCompleted.set(true); 1041 } 1042 } 1043 } 1044 }); 1045 1046 if (!didInvokeWrapRequestConsumer.get()) 1047 throw new IllegalStateException(format("%s::wrapRequest must call requestProcessor", RequestInterceptor.class.getSimpleName())); 1048 } catch (Throwable t) { 1049 // If an error occurred during request wrapping, it's possible a response was never written/communicated back to LifecycleObserver. 1050 // Detect that here and inform LifecycleObserver accordingly. 1051 safelyLog.accept(LogEvent.with(LogEventType.REQUEST_INTERCEPTOR_WRAP_REQUEST_FAILED, 1052 format("An exception occurred while invoking %s::wrapRequest", 1053 RequestInterceptor.class.getSimpleName())) 1054 .throwable(t) 1055 .request(requestHolder.get()) 1056 .resourceMethod(resourceMethodHolder.get()) 1057 .marshaledResponse(marshaledResponseHolder.get()) 1058 .build()); 1059 1060 // If we don't have a response, let the marshaler try to make one for the exception. 1061 // If that fails, use the failsafe. 1062 if (marshaledResponseHolder.get() == null) { 1063 try { 1064 MarshaledResponse marshaledResponse = responseMarshaler.forThrowable(requestHolder.get(), t, resourceMethodHolder.get()); 1065 marshaledResponse = applyCommonPropertiesToMarshaledResponse(requestHolder.get(), marshaledResponse); 1066 marshaledResponseHolder.set(marshaledResponse); 1067 } catch (Throwable t2) { 1068 throwables.add(t2); 1069 1070 safelyLog.accept(LogEvent.with(LogEventType.RESPONSE_MARSHALER_FOR_THROWABLE_FAILED, 1071 format("An exception occurred during request wrapping while invoking %s::forThrowable", 1072 ResponseMarshaler.class.getSimpleName())) 1073 .throwable(t2) 1074 .request(requestHolder.get()) 1075 .resourceMethod(resourceMethodHolder.get()) 1076 .marshaledResponse(marshaledResponseHolder.get()) 1077 .build()); 1078 1079 marshaledResponseHolder.set(provideFailsafeMarshaledResponse(requestHolder.get(), t)); 1080 } 1081 } 1082 1083 if (!willStartResponseWritingCompleted.get()) { 1084 try { 1085 lifecycleObserver.willWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get()); 1086 } catch (Throwable t2) { 1087 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_WILL_WRITE_RESPONSE_FAILED, 1088 format("An exception occurred while invoking %s::willWriteResponse", 1089 LifecycleObserver.class.getSimpleName())) 1090 .throwable(t2) 1091 .request(requestHolder.get()) 1092 .resourceMethod(resourceMethodHolder.get()) 1093 .marshaledResponse(marshaledResponseHolder.get()) 1094 .build()); 1095 } 1096 1097 safelyCollectMetrics.accept( 1098 format("An exception occurred while invoking %s::willWriteResponse", MetricsCollector.class.getSimpleName()), 1099 (metricsInvocation) -> metricsInvocation.willWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get())); 1100 } 1101 1102 try { 1103 Instant responseWriteStarted = Instant.now(); 1104 1105 if (!didFinishResponseWritingCompleted.get()) { 1106 try { 1107 HttpRequestResult requestResult = requestResultHolder.get(); 1108 1109 if (requestResult != null) 1110 requestResultConsumer.accept(requestResult); 1111 else 1112 requestResultConsumer.accept(HttpRequestResult.withMarshaledResponse(marshaledResponseHolder.get()) 1113 .resourceMethod(resourceMethodHolder.get()) 1114 .build()); 1115 1116 Instant responseWriteFinished = Instant.now(); 1117 Duration responseWriteDuration = Duration.between(responseWriteStarted, responseWriteFinished); 1118 1119 try { 1120 lifecycleObserver.didWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), responseWriteDuration); 1121 } catch (Throwable t2) { 1122 throwables.add(t2); 1123 1124 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_WRITE_RESPONSE_FAILED, 1125 format("An exception occurred while invoking %s::didWriteResponse", 1126 LifecycleObserver.class.getSimpleName())) 1127 .throwable(t2) 1128 .request(requestHolder.get()) 1129 .resourceMethod(resourceMethodHolder.get()) 1130 .marshaledResponse(marshaledResponseHolder.get()) 1131 .build()); 1132 } 1133 1134 safelyCollectMetrics.accept( 1135 format("An exception occurred while invoking %s::didWriteResponse", MetricsCollector.class.getSimpleName()), 1136 (metricsInvocation) -> metricsInvocation.didWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), 1137 marshaledResponseHolder.get(), responseWriteDuration)); 1138 } catch (Throwable t2) { 1139 throwables.add(t2); 1140 1141 Instant responseWriteFinished = Instant.now(); 1142 Duration responseWriteDuration = Duration.between(responseWriteStarted, responseWriteFinished); 1143 1144 try { 1145 lifecycleObserver.didFailToWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), responseWriteDuration, t); 1146 } catch (Throwable t3) { 1147 throwables.add(t3); 1148 1149 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_WRITE_RESPONSE_FAILED, 1150 format("An exception occurred while invoking %s::didFailToWriteResponse", 1151 LifecycleObserver.class.getSimpleName())) 1152 .throwable(t3) 1153 .request(requestHolder.get()) 1154 .resourceMethod(resourceMethodHolder.get()) 1155 .marshaledResponse(marshaledResponseHolder.get()) 1156 .build()); 1157 } 1158 1159 safelyCollectMetrics.accept( 1160 format("An exception occurred while invoking %s::didFailToWriteResponse", MetricsCollector.class.getSimpleName()), 1161 (metricsInvocation) -> metricsInvocation.didFailToWriteResponse(serverType, requestHolder.get(), resourceMethodHolder.get(), 1162 marshaledResponseHolder.get(), responseWriteDuration, t)); 1163 } 1164 } 1165 } finally { 1166 if (!didFinishRequestHandlingCompleted.get()) { 1167 Duration processingDuration = Duration.between(processingStarted, Instant.now()); 1168 1169 safelyCollectMetrics.accept( 1170 format("An exception occurred while invoking %s::didFinishRequestHandling", MetricsCollector.class.getSimpleName()), 1171 (metricsInvocation) -> metricsInvocation.didFinishRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), processingDuration, Collections.unmodifiableList(throwables))); 1172 1173 try { 1174 lifecycleObserver.didFinishRequestHandling(serverType, requestHolder.get(), resourceMethodHolder.get(), marshaledResponseHolder.get(), processingDuration, Collections.unmodifiableList(throwables)); 1175 } catch (Throwable t2) { 1176 safelyLog.accept(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_FINISH_REQUEST_HANDLING_FAILED, 1177 format("An exception occurred while invoking %s::didFinishRequestHandling", 1178 LifecycleObserver.class.getSimpleName())) 1179 .throwable(t2) 1180 .request(requestHolder.get()) 1181 .resourceMethod(resourceMethodHolder.get()) 1182 .marshaledResponse(marshaledResponseHolder.get()) 1183 .build()); 1184 } 1185 } 1186 } 1187 } 1188 } 1189 1190 @NonNull 1191 protected HttpRequestResult toHttpRequestResult(@NonNull Request request, 1192 @Nullable ResourceMethod resourceMethod, 1193 @NonNull ServerType serverType) throws Throwable { 1194 requireNonNull(request); 1195 requireNonNull(serverType); 1196 1197 ResourceMethodParameterProvider resourceMethodParameterProvider = getSokletConfig().getResourceMethodParameterProvider(); 1198 InstanceProvider instanceProvider = getSokletConfig().getInstanceProvider(); 1199 CorsAuthorizer corsAuthorizer = getSokletConfig().getCorsAuthorizer(); 1200 ResourceMethodResolver resourceMethodResolver = getSokletConfig().getResourceMethodResolver(); 1201 ResponseMarshaler responseMarshaler = getSokletConfig().getResponseMarshaler(); 1202 CorsPreflight corsPreflight = request.getCorsPreflight().orElse(null); 1203 1204 // Special short-circuit for big requests 1205 if (request.isContentTooLarge()) 1206 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forContentTooLarge(request, resourceMethodResolver.resourceMethodForRequest(request, serverType).orElse(null))) 1207 .resourceMethod(resourceMethod) 1208 .build(); 1209 1210 // Special short-circuit for OPTIONS * 1211 if (isOptionsSplat(request.getResourcePath())) 1212 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forOptionsSplat(request)).build(); 1213 1214 // No resource method was found for this HTTP method and path. 1215 if (resourceMethod == null) { 1216 // If this was an OPTIONS request, do special processing. 1217 // If not, figure out if we should return a 404 or 405. 1218 if (request.getHttpMethod() == HttpMethod.OPTIONS) { 1219 // See what methods are available to us for this request's path 1220 Map<HttpMethod, ResourceMethod> matchingResourceMethodsByHttpMethod = resolveMatchingResourceMethodsByHttpMethod(request, resourceMethodResolver, serverType); 1221 1222 // Special handling for CORS preflight requests, if needed 1223 if (corsPreflight != null) { 1224 // Let configuration function determine if we should authorize this request. 1225 // Discard any OPTIONS references - see https://stackoverflow.com/a/68529748 1226 Map<HttpMethod, ResourceMethod> nonOptionsMatchingResourceMethodsByHttpMethod = matchingResourceMethodsByHttpMethod.entrySet().stream() 1227 .filter(entry -> entry.getKey() != HttpMethod.OPTIONS) 1228 .collect(Collectors.toMap(Entry::getKey, Entry::getValue)); 1229 1230 CorsPreflightResponse corsPreflightResponse = corsAuthorizer.authorizePreflight(request, corsPreflight, nonOptionsMatchingResourceMethodsByHttpMethod).orElse(null); 1231 1232 // Allow or reject CORS depending on what the function said to do 1233 if (corsPreflightResponse != null) { 1234 // Allow 1235 MarshaledResponse marshaledResponse = responseMarshaler.forCorsPreflightAllowed(request, corsPreflight, corsPreflightResponse); 1236 1237 return HttpRequestResult.withMarshaledResponse(marshaledResponse) 1238 .corsPreflightResponse(corsPreflightResponse) 1239 .build(); 1240 } 1241 1242 // Reject 1243 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forCorsPreflightRejected(request, corsPreflight)) 1244 .build(); 1245 } else { 1246 // Just a normal OPTIONS response (non-CORS-preflight). 1247 // If there's a matching OPTIONS resource method for this OPTIONS request, then invoke it. 1248 ResourceMethod optionsResourceMethod = matchingResourceMethodsByHttpMethod.get(HttpMethod.OPTIONS); 1249 1250 if (optionsResourceMethod != null) { 1251 resourceMethod = optionsResourceMethod; 1252 } else { 1253 Set<HttpMethod> allowedHttpMethods = allowedHttpMethodsForResponse(matchingResourceMethodsByHttpMethod, true); 1254 1255 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forOptions(request, allowedHttpMethods)) 1256 .build(); 1257 } 1258 } 1259 } else if (request.getHttpMethod() == HttpMethod.HEAD) { 1260 // If there's a matching GET resource method for this HEAD request, then invoke it 1261 Request headGetRequest = request.copy().httpMethod(HttpMethod.GET).finish(); 1262 ResourceMethod headGetResourceMethod = resourceMethodResolver.resourceMethodForRequest(headGetRequest, serverType).orElse(null); 1263 1264 if (headGetResourceMethod != null) 1265 resourceMethod = headGetResourceMethod; 1266 else 1267 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forNotFound(request)) 1268 .build(); 1269 } else { 1270 // Not an OPTIONS request, so it's possible we have a 405. See if other HTTP methods match... 1271 Map<HttpMethod, ResourceMethod> otherMatchingResourceMethodsByHttpMethod = resolveMatchingResourceMethodsByHttpMethod(request, resourceMethodResolver, serverType); 1272 1273 Set<HttpMethod> matchingNonOptionsHttpMethods = otherMatchingResourceMethodsByHttpMethod.keySet().stream() 1274 .filter(httpMethod -> httpMethod != HttpMethod.OPTIONS) 1275 .collect(Collectors.toSet()); 1276 1277 if (matchingNonOptionsHttpMethods.size() > 0) { 1278 // ...if some do, it's a 405 1279 Set<HttpMethod> allowedHttpMethods = allowedHttpMethodsForResponse(otherMatchingResourceMethodsByHttpMethod, true); 1280 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forMethodNotAllowed(request, allowedHttpMethods)) 1281 .build(); 1282 } else { 1283 // no matching resource method found, it's a 404 1284 return HttpRequestResult.withMarshaledResponse(responseMarshaler.forNotFound(request)) 1285 .build(); 1286 } 1287 } 1288 } 1289 1290 // Found a resource method - happy path. 1291 // 1. Get an instance of the resource class 1292 // 2. Get values to pass to the resource method on the resource class 1293 // 3. Invoke the resource method and use its return value to drive a response 1294 Class<?> resourceClass = resourceMethod.getMethod().getDeclaringClass(); 1295 Object resourceClassInstance; 1296 1297 try { 1298 resourceClassInstance = instanceProvider.provide(resourceClass); 1299 } catch (Exception e) { 1300 throw new IllegalArgumentException(format("Unable to acquire an instance of %s", resourceClass.getName()), e); 1301 } 1302 1303 List<Object> parameterValues = resourceMethodParameterProvider.parameterValuesForResourceMethod(request, resourceMethod); 1304 1305 Object responseObject; 1306 1307 try { 1308 responseObject = resourceMethod.getMethod().invoke(resourceClassInstance, parameterValues.toArray()); 1309 } catch (InvocationTargetException e) { 1310 if (e.getTargetException() != null) 1311 throw e.getTargetException(); 1312 1313 throw e; 1314 } 1315 1316 // Unwrap the Optional<T>, if one exists. We do not recurse deeper than one level 1317 if (responseObject instanceof Optional<?>) 1318 responseObject = ((Optional<?>) responseObject).orElse(null); 1319 1320 Response response; 1321 SseHandshakeResult sseHandshakeResult = null; 1322 1323 // If null/void return, it's a 204 1324 // If it's a MarshaledResponse object, no marshaling + return it immediately - caller knows exactly what it wants to write. 1325 // If it's a Response object, use as is. 1326 // If it's a non-Response type of object, assume it's the response body and wrap in a Response. 1327 if (responseObject == null) { 1328 response = Response.withStatusCode(204).build(); 1329 } else if (responseObject instanceof MarshaledResponse) { 1330 MarshaledResponse marshaledResponse = (MarshaledResponse) responseObject; 1331 enforceBodylessStatusCode(marshaledResponse.getStatusCode(), marshaledResponse.getBody().isPresent() || marshaledResponse.getStream().isPresent()); 1332 1333 return HttpRequestResult.withMarshaledResponse(marshaledResponse) 1334 .resourceMethod(resourceMethod) 1335 .build(); 1336 } else if (responseObject instanceof Response) { 1337 response = (Response) responseObject; 1338 } else if (responseObject instanceof SseHandshakeResult.Accepted accepted) { // SSE "accepted" handshake 1339 return HttpRequestResult.withMarshaledResponse(toMarshaledResponse(accepted)) 1340 .resourceMethod(resourceMethod) 1341 .sseHandshakeResult(accepted) 1342 .build(); 1343 } else if (responseObject instanceof SseHandshakeResult.Rejected rejected) { // SSE "rejected" handshake 1344 response = rejected.getResponse(); 1345 sseHandshakeResult = rejected; 1346 } else { 1347 response = Response.withStatusCode(200).body(responseObject).build(); 1348 } 1349 1350 enforceBodylessStatusCode(response.getStatusCode(), response.getBody().isPresent()); 1351 1352 MarshaledResponse marshaledResponse = responseMarshaler.forResourceMethod(request, response, resourceMethod); 1353 1354 enforceBodylessStatusCode(marshaledResponse.getStatusCode(), marshaledResponse.getBody().isPresent() || marshaledResponse.getStream().isPresent()); 1355 1356 return HttpRequestResult.withMarshaledResponse(marshaledResponse) 1357 .response(response) 1358 .resourceMethod(resourceMethod) 1359 .sseHandshakeResult(sseHandshakeResult) 1360 .build(); 1361 } 1362 1363 @NonNull 1364 private MarshaledResponse toMarshaledResponse(SseHandshakeResult.@NonNull Accepted accepted) { 1365 requireNonNull(accepted); 1366 1367 Map<String, Set<String>> acceptedHeaders = accepted.getHeaders(); 1368 Map<String, Set<String>> headers = acceptedHeaders == null ? Map.of() : acceptedHeaders; 1369 LinkedCaseInsensitiveMap<Set<String>> finalHeaders = new LinkedCaseInsensitiveMap<>(DEFAULT_ACCEPTED_HANDSHAKE_HEADERS.size() + headers.size()); 1370 1371 // Start with defaults 1372 for (Map.Entry<String, Set<String>> e : DEFAULT_ACCEPTED_HANDSHAKE_HEADERS.entrySet()) 1373 finalHeaders.put(e.getKey(), e.getValue()); // values already unmodifiable 1374 1375 // Overlay user-supplied headers (prefer user values on key collision) 1376 for (Map.Entry<String, Set<String>> e : headers.entrySet()) { 1377 // Defensively copy so callers can't mutate after construction 1378 Set<String> values = e.getValue() == null ? Set.of() : Set.copyOf(e.getValue()); 1379 finalHeaders.put(e.getKey(), values); 1380 } 1381 1382 return MarshaledResponse.withStatusCode(200) 1383 .headers(finalHeaders) 1384 .cookies(accepted.getCookies()) 1385 .build(); 1386 } 1387 1388 private static void enforceBodylessStatusCode(@NonNull Integer statusCode, 1389 @NonNull Boolean hasBody) { 1390 requireNonNull(statusCode); 1391 requireNonNull(hasBody); 1392 1393 if (hasBody && isBodylessStatusCode(statusCode)) 1394 throw new IllegalStateException(format("HTTP status code %d must not include a response body", statusCode)); 1395 } 1396 1397 private static boolean isBodylessStatusCode(@NonNull Integer statusCode) { 1398 requireNonNull(statusCode); 1399 return (statusCode >= 100 && statusCode < 200) || statusCode == 204 || statusCode == 304; 1400 } 1401 1402 @NonNull 1403 protected MarshaledResponse applyHeadResponseIfApplicable(@NonNull Request request, 1404 @NonNull MarshaledResponse marshaledResponse) { 1405 if (request.getHttpMethod() != HttpMethod.HEAD) 1406 return marshaledResponse; 1407 1408 return getSokletConfig().getResponseMarshaler().forHead(request, marshaledResponse); 1409 } 1410 1411 // Hat tip to Aslan Parçası and GrayStar 1412 @NonNull 1413 protected MarshaledResponse applyCommonPropertiesToMarshaledResponse(@NonNull Request request, 1414 @NonNull MarshaledResponse marshaledResponse) { 1415 requireNonNull(request); 1416 requireNonNull(marshaledResponse); 1417 1418 return applyCommonPropertiesToMarshaledResponse(request, marshaledResponse, false); 1419 } 1420 1421 protected void handleMcpRequest(@NonNull Request request, 1422 @NonNull Consumer<HttpRequestResult> requestResultConsumer) { 1423 requireNonNull(request); 1424 requireNonNull(requestResultConsumer); 1425 1426 requestResultConsumer.accept(this.defaultMcpRuntime.handleRequest(request)); 1427 } 1428 1429 protected void handleSimulatedMcpStreamDisconnect(@NonNull Request request, 1430 @NonNull String sessionId) { 1431 requireNonNull(request); 1432 requireNonNull(sessionId); 1433 1434 this.defaultMcpRuntime.handleClientDisconnectedStream(request, sessionId); 1435 } 1436 1437 @NonNull 1438 protected MarshaledResponse applyCommonPropertiesToMarshaledResponse(@NonNull Request request, 1439 @NonNull MarshaledResponse marshaledResponse, 1440 @NonNull Boolean suppressContentLength) { 1441 requireNonNull(request); 1442 requireNonNull(marshaledResponse); 1443 requireNonNull(suppressContentLength); 1444 1445 // Don't write Content-Length for an accepted SSE Handshake, for example 1446 if (!suppressContentLength) 1447 marshaledResponse = applyContentLengthIfApplicable(request, marshaledResponse); 1448 1449 // If the Date header is missing, add it using our cached provider 1450 if (!marshaledResponse.getHeaders().containsKey("Date")) 1451 marshaledResponse = marshaledResponse.copy() 1452 .headers(headers -> headers.put("Date", Set.of(HttpDate.currentSecondHeaderValue()))) 1453 .finish(); 1454 1455 marshaledResponse = applyCorsResponseIfApplicable(request, marshaledResponse); 1456 1457 return marshaledResponse; 1458 } 1459 1460 @NonNull 1461 protected MarshaledResponse applyContentLengthIfApplicable(@NonNull Request request, 1462 @NonNull MarshaledResponse marshaledResponse) { 1463 requireNonNull(request); 1464 requireNonNull(marshaledResponse); 1465 1466 if (marshaledResponse.isStreaming()) 1467 return marshaledResponse; 1468 1469 Set<String> normalizedHeaderNames = marshaledResponse.getHeaders().keySet().stream() 1470 .map(headerName -> headerName.toLowerCase(Locale.US)) 1471 .collect(Collectors.toSet()); 1472 1473 // If Content-Length is already specified, don't do anything 1474 if (normalizedHeaderNames.contains("content-length") || normalizedHeaderNames.contains("transfer-encoding")) 1475 return marshaledResponse; 1476 1477 if (shouldOmitAutomaticContentLength(request, marshaledResponse)) 1478 return marshaledResponse; 1479 1480 // If Content-Length is not specified, specify as the number of bytes in the body 1481 return marshaledResponse.copy() 1482 .headers((mutableHeaders) -> { 1483 String contentLengthHeaderValue = String.valueOf(marshaledResponse.getBodyLength()); 1484 mutableHeaders.put("Content-Length", Set.of(contentLengthHeaderValue)); 1485 }).finish(); 1486 } 1487 1488 private boolean shouldOmitAutomaticContentLength(@NonNull Request request, 1489 @NonNull MarshaledResponse marshaledResponse) { 1490 requireNonNull(request); 1491 requireNonNull(marshaledResponse); 1492 1493 int statusCode = marshaledResponse.getStatusCode(); 1494 1495 if ((statusCode >= 100 && statusCode < 200) || statusCode == 204 || statusCode == 304) 1496 return true; 1497 1498 return request.getHttpMethod() == HttpMethod.HEAD && marshaledResponse.getBodyLength() == 0L; 1499 } 1500 1501 @NonNull 1502 protected MarshaledResponse applyCorsResponseIfApplicable(@NonNull Request request, 1503 @NonNull MarshaledResponse marshaledResponse) { 1504 requireNonNull(request); 1505 requireNonNull(marshaledResponse); 1506 1507 Cors cors = request.getCors().orElse(null); 1508 1509 // If non-CORS request, nothing further to do (note that CORS preflight was handled earlier) 1510 if (cors == null) 1511 return marshaledResponse; 1512 1513 CorsAuthorizer corsAuthorizer = getSokletConfig().getCorsAuthorizer(); 1514 1515 // Does the authorizer say we are authorized? 1516 CorsResponse corsResponse = corsAuthorizer.authorize(request, cors).orElse(null); 1517 1518 // Not authorized - don't apply CORS headers to the response 1519 if (corsResponse == null) 1520 return marshaledResponse; 1521 1522 // Authorized - OK, let's apply the headers to the response 1523 return getSokletConfig().getResponseMarshaler().forCorsAllowed(request, cors, corsResponse, marshaledResponse); 1524 } 1525 1526 @NonNull 1527 protected Map<@NonNull HttpMethod, @NonNull ResourceMethod> resolveMatchingResourceMethodsByHttpMethod(@NonNull Request request, 1528 @NonNull ResourceMethodResolver resourceMethodResolver, 1529 @NonNull ServerType serverType) { 1530 requireNonNull(request); 1531 requireNonNull(resourceMethodResolver); 1532 requireNonNull(serverType); 1533 1534 // Special handling for OPTIONS * 1535 if (isOptionsSplat(request.getResourcePath())) 1536 return new LinkedHashMap<>(); 1537 1538 Map<HttpMethod, ResourceMethod> matchingResourceMethodsByHttpMethod = new LinkedHashMap<>(HttpMethod.values().length); 1539 1540 for (HttpMethod httpMethod : HttpMethod.values()) { 1541 // Make a quick copy of the request to see if other paths match 1542 Request otherRequest = Request.withPath(httpMethod, request.getPath()).build(); 1543 ResourceMethod resourceMethod = resourceMethodResolver.resourceMethodForRequest(otherRequest, serverType).orElse(null); 1544 1545 if (resourceMethod != null) 1546 matchingResourceMethodsByHttpMethod.put(httpMethod, resourceMethod); 1547 } 1548 1549 return matchingResourceMethodsByHttpMethod; 1550 } 1551 1552 @SuppressWarnings("ReferenceEquality") 1553 private static Boolean isOptionsSplat(@NonNull ResourcePath resourcePath) { 1554 requireNonNull(resourcePath); 1555 return resourcePath == ResourcePath.OPTIONS_SPLAT_RESOURCE_PATH; 1556 } 1557 1558 @NonNull 1559 private static Set<@NonNull HttpMethod> allowedHttpMethodsForResponse(@NonNull Map<@NonNull HttpMethod, @NonNull ResourceMethod> matchingResourceMethodsByHttpMethod, 1560 @NonNull Boolean includeOptions) { 1561 requireNonNull(matchingResourceMethodsByHttpMethod); 1562 requireNonNull(includeOptions); 1563 1564 Set<HttpMethod> allowedHttpMethods = EnumSet.noneOf(HttpMethod.class); 1565 allowedHttpMethods.addAll(matchingResourceMethodsByHttpMethod.keySet()); 1566 1567 if (includeOptions) 1568 allowedHttpMethods.add(HttpMethod.OPTIONS); 1569 1570 if (matchingResourceMethodsByHttpMethod.containsKey(HttpMethod.GET) || matchingResourceMethodsByHttpMethod.containsKey(HttpMethod.HEAD)) 1571 allowedHttpMethods.add(HttpMethod.HEAD); 1572 1573 return allowedHttpMethods; 1574 } 1575 1576 @NonNull 1577 protected MarshaledResponse provideFailsafeMarshaledResponse(@NonNull Request request, 1578 @NonNull Throwable throwable) { 1579 requireNonNull(request); 1580 requireNonNull(throwable); 1581 1582 Integer statusCode = 500; 1583 Charset charset = StandardCharsets.UTF_8; 1584 1585 return MarshaledResponse.withStatusCode(statusCode) 1586 .headers(Map.of("Content-Type", Set.of(format("text/plain; charset=%s", charset.name())))) 1587 .body(format("HTTP %d: %s", statusCode, StatusCode.fromStatusCode(statusCode).get().getReasonPhrase()).getBytes(charset)) 1588 .build(); 1589 } 1590 1591 /** 1592 * Synonym for {@link #stop()}. 1593 */ 1594 @Override 1595 public void close() { 1596 stop(); 1597 } 1598 1599 /** 1600 * Is any managed transport server started? 1601 * 1602 * @return {@code true} if at least one configured transport server is started, {@code false} otherwise 1603 */ 1604 @NonNull 1605 public Boolean isStarted() { 1606 getLock().lock(); 1607 1608 try { 1609 HttpServer httpServer = getSokletConfig().getHttpServer().orElse(null); 1610 1611 if (httpServer != null && httpServer.isStarted()) 1612 return true; 1613 1614 SseServer sseServer = getSokletConfig().getSseServer().orElse(null); 1615 if (sseServer != null && sseServer.isStarted()) 1616 return true; 1617 1618 McpServer mcpServer = getSokletConfig().getMcpServer().orElse(null); 1619 return mcpServer != null && mcpServer.isStarted(); 1620 } finally { 1621 getLock().unlock(); 1622 } 1623 } 1624 1625 /** 1626 * Runs Soklet with special non-network "simulator" implementations of the configured transport servers - useful for integration testing. 1627 * <p> 1628 * See <a href="https://www.soklet.com/docs/testing">https://www.soklet.com/docs/testing</a> for how to write these tests. 1629 * 1630 * @param sokletConfig configuration that drives the Soklet system 1631 * @param simulatorConsumer code to execute within the context of the simulator 1632 */ 1633 public static void runSimulator(@NonNull SokletConfig sokletConfig, 1634 @NonNull Consumer<Simulator> simulatorConsumer) { 1635 runSimulator(sokletConfig, SimulatorOptions.defaultInstance(), simulatorConsumer); 1636 } 1637 1638 /** 1639 * Runs Soklet with special non-network "simulator" implementations of the configured transport servers - useful for integration testing. 1640 * <p> 1641 * See <a href="https://www.soklet.com/docs/testing">https://www.soklet.com/docs/testing</a> for how to write these tests. 1642 * 1643 * @param sokletConfig configuration that drives the Soklet system 1644 * @param simulatorOptions simulator behavior options 1645 * @param simulatorConsumer code to execute within the context of the simulator 1646 */ 1647 public static void runSimulator(@NonNull SokletConfig sokletConfig, 1648 @NonNull SimulatorOptions simulatorOptions, 1649 @NonNull Consumer<Simulator> simulatorConsumer) { 1650 requireNonNull(sokletConfig); 1651 requireNonNull(simulatorOptions); 1652 requireNonNull(simulatorConsumer); 1653 1654 // Create Soklet instance - this initializes the REAL implementations through proxies 1655 Soklet soklet = Soklet.fromConfig(sokletConfig); 1656 1657 // Extract proxies (they're guaranteed to be proxies now) 1658 HttpServerProxy serverProxy = sokletConfig.getHttpServer() 1659 .map(server -> (HttpServerProxy) server) 1660 .orElse(null); 1661 SseServerProxy sseServerProxy = sokletConfig.getSseServer() 1662 .map(s -> (SseServerProxy) s) 1663 .orElse(null); 1664 McpServerProxy mcpServerProxy = sokletConfig.getMcpServer() 1665 .map(mcpServer -> (McpServerProxy) mcpServer) 1666 .orElse(null); 1667 1668 // Create mock implementations 1669 MockHttpServer mockServer = serverProxy == null ? null : new MockHttpServer(); 1670 MockSseServer mockSseServer = new MockSseServer(); 1671 MockMcpServer mockMcpServer = mcpServerProxy == null ? null : new MockMcpServer(mcpServerProxy.getRealImplementation()); 1672 1673 // Switch proxies to simulator mode 1674 if (serverProxy != null) 1675 serverProxy.enableSimulatorMode(mockServer); 1676 1677 if (sseServerProxy != null) 1678 sseServerProxy.enableSimulatorMode(mockSseServer); 1679 1680 if (mcpServerProxy != null) 1681 mcpServerProxy.enableSimulatorMode(mockMcpServer); 1682 1683 try { 1684 // Initialize mocks with request handlers that delegate to Soklet's processing 1685 if (mockServer != null) 1686 mockServer.initialize(sokletConfig, (request, marshaledResponseConsumer) -> { 1687 // Delegate to Soklet's internal request handling 1688 soklet.handleRequest(request, ServerType.STANDARD_HTTP, marshaledResponseConsumer); 1689 }); 1690 1691 if (mockSseServer != null) 1692 mockSseServer.initialize(sokletConfig, (request, marshaledResponseConsumer) -> { 1693 // Delegate to Soklet's internal request handling for SSE 1694 soklet.handleRequest(request, ServerType.SSE, marshaledResponseConsumer); 1695 }); 1696 1697 if (mockMcpServer != null) 1698 mockMcpServer.initialize(sokletConfig, soklet::handleMcpRequest); 1699 1700 if (mockMcpServer != null) 1701 mockMcpServer.onClientDisconnectedMcpStream(soklet::handleSimulatedMcpStreamDisconnect); 1702 1703 // Create and provide simulator 1704 Simulator simulator = new DefaultSimulator(mockServer, mockSseServer, mockMcpServer, simulatorOptions); 1705 simulatorConsumer.accept(simulator); 1706 } finally { 1707 // Always restore to real implementations 1708 if (serverProxy != null) 1709 serverProxy.disableSimulatorMode(); 1710 1711 if (sseServerProxy != null) 1712 sseServerProxy.disableSimulatorMode(); 1713 1714 if (mcpServerProxy != null) 1715 mcpServerProxy.disableSimulatorMode(); 1716 } 1717 } 1718 1719 @NonNull 1720 protected SokletConfig getSokletConfig() { 1721 return this.sokletConfig; 1722 } 1723 1724 @NonNull 1725 protected ReentrantLock getLock() { 1726 return this.lock; 1727 } 1728 1729 @NonNull 1730 protected AtomicReference<@NonNull CountDownLatch> getAwaitShutdownLatchReference() { 1731 return this.awaitShutdownLatchReference; 1732 } 1733 1734 @ThreadSafe 1735 static class DefaultSimulator implements Simulator { 1736 @Nullable 1737 private MockHttpServer server; 1738 @Nullable 1739 private MockSseServer sseServer; 1740 @Nullable 1741 private MockMcpServer mcpServer; 1742 @NonNull 1743 private final SimulatorOptions simulatorOptions; 1744 1745 public DefaultSimulator(@Nullable MockHttpServer server, 1746 @Nullable MockSseServer sseServer, 1747 @Nullable MockMcpServer mcpServer) { 1748 this(server, sseServer, mcpServer, SimulatorOptions.defaultInstance()); 1749 } 1750 1751 public DefaultSimulator(@Nullable MockHttpServer server, 1752 @Nullable MockSseServer sseServer, 1753 @Nullable MockMcpServer mcpServer, 1754 @NonNull SimulatorOptions simulatorOptions) { 1755 this.server = server; 1756 this.sseServer = sseServer; 1757 this.mcpServer = mcpServer; 1758 this.simulatorOptions = requireNonNull(simulatorOptions); 1759 } 1760 1761 @NonNull 1762 @Override 1763 public HttpRequestResult performHttpRequest(@NonNull Request request) { 1764 MockHttpServer server = getHttpServer().orElse(null); 1765 1766 if (server == null) 1767 throw new IllegalStateException(format("You must specify a %s in your %s to simulate requests", 1768 HttpServer.class.getSimpleName(), SokletConfig.class.getSimpleName())); 1769 1770 AtomicReference<HttpRequestResult> requestResultHolder = new AtomicReference<>(); 1771 HttpServer.RequestHandler requestHandler = server.getRequestHandler().orElse(null); 1772 1773 if (requestHandler == null) 1774 throw new IllegalStateException("You must register a request handler prior to simulating requests"); 1775 1776 requestHandler.handleRequest(request, (requestResult -> { 1777 requestResultHolder.set(requestResult); 1778 })); 1779 1780 return materializeStreamingResponse(request, requestResultHolder.get()); 1781 } 1782 1783 @NonNull 1784 private HttpRequestResult materializeStreamingResponse(@NonNull Request request, 1785 @Nullable HttpRequestResult requestResult) { 1786 requireNonNull(request); 1787 1788 if (requestResult == null) 1789 throw new IllegalStateException("No HTTP request result was produced by the simulator"); 1790 1791 StreamingResponseBody stream = requestResult.getMarshaledResponse().getStream().orElse(null); 1792 1793 if (stream == null) 1794 return requestResult; 1795 1796 byte[] bytes; 1797 Instant streamStarted = Instant.now(); 1798 1799 try { 1800 bytes = materializeStreamingResponseBody(request, requestResult, stream); 1801 notifyDidTerminateSimulatorResponseStream(request, requestResult, streamStarted, 1802 Duration.between(streamStarted, Instant.now()), null, null); 1803 } catch (StreamingResponseCanceledException e) { 1804 StreamTerminationReason cancelationReason = e.getCancelationReason(); 1805 Throwable cause = e.getCancelationCause().orElse(null); 1806 notifyDidTerminateSimulatorResponseStream(request, requestResult, streamStarted, 1807 Duration.between(streamStarted, Instant.now()), cancelationReason, cause); 1808 throw new IllegalStateException("Simulated streaming response was canceled: " + cancelationReason.name(), e); 1809 } catch (InterruptedException e) { 1810 Thread.currentThread().interrupt(); 1811 notifyDidTerminateSimulatorResponseStream(request, requestResult, streamStarted, 1812 Duration.between(streamStarted, Instant.now()), StreamTerminationReason.CLIENT_DISCONNECTED, e); 1813 throw new IllegalStateException("Simulated streaming response was canceled: CLIENT_DISCONNECTED", e); 1814 } catch (Throwable t) { 1815 notifyDidTerminateSimulatorResponseStream(request, requestResult, streamStarted, 1816 Duration.between(streamStarted, Instant.now()), StreamTerminationReason.PRODUCER_FAILED, t); 1817 1818 if (t instanceof Error error) 1819 throw error; 1820 1821 throw new IllegalStateException("Simulated streaming response failed.", t); 1822 } 1823 1824 MarshaledResponse marshaledResponse = requestResult.getMarshaledResponse().copy() 1825 .withoutStream() 1826 .body(bytes) 1827 .finish(); 1828 1829 return requestResult.copy() 1830 .marshaledResponse(marshaledResponse) 1831 .finish(); 1832 } 1833 1834 private void notifyDidTerminateSimulatorResponseStream(@NonNull Request request, 1835 @NonNull HttpRequestResult requestResult, 1836 @NonNull Instant establishedAt, 1837 @NonNull Duration streamDuration, 1838 @Nullable StreamTerminationReason cancelationReason, 1839 @Nullable Throwable throwable) { 1840 requireNonNull(request); 1841 requireNonNull(requestResult); 1842 requireNonNull(establishedAt); 1843 requireNonNull(streamDuration); 1844 1845 MockHttpServer server = getHttpServer().orElse(null); 1846 SokletConfig sokletConfig = server == null ? null : server.getSokletConfig().orElse(null); 1847 1848 if (sokletConfig == null) 1849 return; 1850 1851 MarshaledResponse marshaledResponse = requestResult.getMarshaledResponse(); 1852 ResourceMethod resourceMethod = requestResult.getResourceMethod().orElse(null); 1853 LifecycleObserver lifecycleObserver = sokletConfig.getAggregateLifecycleObserver(); 1854 StreamingResponseHandle streamingResponse = new DefaultStreamingResponseHandle(ServerType.STANDARD_HTTP, 1855 request, resourceMethod, marshaledResponse, establishedAt); 1856 StreamTermination termination = StreamTermination 1857 .with(cancelationReason == null ? StreamTerminationReason.COMPLETED : cancelationReason, streamDuration) 1858 .cause(throwable) 1859 .build(); 1860 1861 try { 1862 lifecycleObserver.willTerminateResponseStream(streamingResponse, termination); 1863 } catch (Throwable t) { 1864 try { 1865 lifecycleObserver.didReceiveLogEvent(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_WILL_TERMINATE_RESPONSE_STREAM_FAILED, 1866 format("An exception occurred while invoking %s::willTerminateResponseStream", LifecycleObserver.class.getSimpleName())) 1867 .throwable(t) 1868 .request(request) 1869 .resourceMethod(resourceMethod) 1870 .marshaledResponse(marshaledResponse) 1871 .build()); 1872 } catch (Throwable ignored) { 1873 // Keep simulator lifecycle observer failures contained. 1874 } 1875 } 1876 1877 try { 1878 lifecycleObserver.didTerminateResponseStream(streamingResponse, termination); 1879 } catch (Throwable t) { 1880 try { 1881 lifecycleObserver.didReceiveLogEvent(LogEvent.with(LogEventType.LIFECYCLE_OBSERVER_DID_TERMINATE_RESPONSE_STREAM_FAILED, 1882 format("An exception occurred while invoking %s::didTerminateResponseStream", LifecycleObserver.class.getSimpleName())) 1883 .throwable(t) 1884 .request(request) 1885 .resourceMethod(resourceMethod) 1886 .marshaledResponse(marshaledResponse) 1887 .build()); 1888 } catch (Throwable ignored) { 1889 // Keep simulator lifecycle observer failures contained. 1890 } 1891 } 1892 } 1893 1894 private void notifyDidReceiveSimulatorStreamCancelationCallbackFailure(@NonNull Request request, 1895 @NonNull HttpRequestResult requestResult, 1896 @NonNull Throwable throwable) { 1897 requireNonNull(request); 1898 requireNonNull(requestResult); 1899 requireNonNull(throwable); 1900 1901 MockHttpServer server = getHttpServer().orElse(null); 1902 SokletConfig sokletConfig = server == null ? null : server.getSokletConfig().orElse(null); 1903 1904 if (sokletConfig == null) 1905 return; 1906 1907 LifecycleObserver lifecycleObserver = sokletConfig.getAggregateLifecycleObserver(); 1908 ResourceMethod resourceMethod = requestResult.getResourceMethod().orElse(null); 1909 MarshaledResponse marshaledResponse = requestResult.getMarshaledResponse(); 1910 1911 try { 1912 lifecycleObserver.didReceiveLogEvent(LogEvent.with(LogEventType.RESPONSE_STREAM_CANCELATION_CALLBACK_FAILED, 1913 "An exception occurred while invoking a streaming response cancelation callback") 1914 .throwable(throwable) 1915 .request(request) 1916 .resourceMethod(resourceMethod) 1917 .marshaledResponse(marshaledResponse) 1918 .build()); 1919 } catch (Throwable ignored) { 1920 // Keep simulator lifecycle observer failures contained. 1921 } 1922 } 1923 1924 @NonNull 1925 private byte[] materializeStreamingResponseBody(@NonNull Request request, 1926 @NonNull HttpRequestResult requestResult, 1927 @NonNull StreamingResponseBody stream) throws Exception { 1928 requireNonNull(request); 1929 requireNonNull(requestResult); 1930 requireNonNull(stream); 1931 1932 SimulatorCancelationToken cancelationToken = new SimulatorCancelationToken(throwable -> 1933 notifyDidReceiveSimulatorStreamCancelationCallbackFailure(request, requestResult, throwable)); 1934 SimulatorStreamingResponseContext context = new SimulatorStreamingResponseContext(request, cancelationToken); 1935 SimulatorResponseStream output = new SimulatorResponseStream(getSimulatorOptions().getStreamingResponseBodyLimitInBytes(), cancelationToken); 1936 1937 try { 1938 if (stream instanceof StreamingResponseBody.WriterBody writerBody) { 1939 writerBody.getWriter().writeTo(output, context); 1940 } else if (stream instanceof StreamingResponseBody.InputStreamBody inputStreamBody) { 1941 try (java.io.InputStream inputStream = requireNonNull(inputStreamBody.getInputStreamSupplier().get()); 1942 AutoCloseable ignored = context.onCancel(() -> closeQuietly(inputStream))) { 1943 byte[] buffer = new byte[inputStreamBody.getBufferSizeInBytes()]; 1944 int read; 1945 1946 while ((read = inputStream.read(buffer)) >= 0) { 1947 context.throwIfCanceled(); 1948 if (read > 0) 1949 output.write(ByteBuffer.wrap(buffer, 0, read)); 1950 } 1951 } 1952 } else if (stream instanceof StreamingResponseBody.ReaderBody readerBody) { 1953 try (java.io.Reader reader = requireNonNull(readerBody.getReaderSupplier().get()); 1954 AutoCloseable ignored = context.onCancel(() -> closeQuietly(reader))) { 1955 materializeReader(readerBody, reader, output, context); 1956 } 1957 } else if (stream instanceof StreamingResponseBody.PublisherBody publisherBody) { 1958 materializePublisher(publisherBody, output, context); 1959 } else { 1960 throw new IllegalStateException(format("Unsupported streaming response body type: %s", stream.getClass().getName())); 1961 } 1962 } catch (StreamingResponseCanceledException e) { 1963 cancelationToken.cancel(e.getCancelationReason(), e.getCancelationCause().orElse(null)); 1964 throw e; 1965 } catch (InterruptedException e) { 1966 Thread.currentThread().interrupt(); 1967 cancelationToken.cancel(cancelationToken.getCancelationReason() 1968 .orElse(StreamTerminationReason.CLIENT_DISCONNECTED), e); 1969 throw e; 1970 } catch (Throwable t) { 1971 cancelationToken.cancel(StreamTerminationReason.PRODUCER_FAILED, t); 1972 1973 if (t instanceof Exception exception) 1974 throw exception; 1975 1976 if (t instanceof Error error) 1977 throw error; 1978 1979 throw new RuntimeException(t); 1980 } 1981 1982 return output.toByteArray(); 1983 } 1984 1985 private void materializeReader(com.soklet.StreamingResponseBody.@NonNull ReaderBody readerBody, 1986 @NonNull Reader reader, 1987 @NonNull SimulatorResponseStream output, 1988 @NonNull SimulatorStreamingResponseContext context) throws IOException, InterruptedException, StreamingResponseCanceledException, CharacterCodingException { 1989 requireNonNull(readerBody); 1990 requireNonNull(reader); 1991 requireNonNull(output); 1992 requireNonNull(context); 1993 1994 CharsetEncoder encoder = readerBody.newEncoder(); 1995 CharBuffer charBuffer = CharBuffer.allocate(readerBody.getBufferSizeInCharacters()); 1996 ByteBuffer byteBuffer = ByteBuffer.allocate(Math.max(128, (int) Math.ceil(readerBody.getBufferSizeInCharacters() * encoder.maxBytesPerChar()))); 1997 1998 while (reader.read(charBuffer) >= 0) { 1999 context.throwIfCanceled(); 2000 charBuffer.flip(); 2001 encodeCharsForSimulator(encoder, charBuffer, byteBuffer, false, output); 2002 charBuffer.compact(); 2003 } 2004 2005 charBuffer.flip(); 2006 encodeCharsForSimulator(encoder, charBuffer, byteBuffer, true, output); 2007 2008 CoderResult result; 2009 do { 2010 result = encoder.flush(byteBuffer); 2011 writeEncodedBytesForSimulator(byteBuffer, output); 2012 if (result.isError()) 2013 result.throwException(); 2014 } while (result.isOverflow()); 2015 } 2016 2017 private void encodeCharsForSimulator(@NonNull CharsetEncoder encoder, 2018 @NonNull CharBuffer charBuffer, 2019 @NonNull ByteBuffer byteBuffer, 2020 boolean endOfInput, 2021 @NonNull SimulatorResponseStream output) throws IOException, InterruptedException, StreamingResponseCanceledException, CharacterCodingException { 2022 CoderResult result; 2023 2024 do { 2025 result = encoder.encode(charBuffer, byteBuffer, endOfInput); 2026 writeEncodedBytesForSimulator(byteBuffer, output); 2027 2028 if (result.isError()) 2029 result.throwException(); 2030 } while (result.isOverflow()); 2031 } 2032 2033 private void writeEncodedBytesForSimulator(@NonNull ByteBuffer byteBuffer, 2034 @NonNull SimulatorResponseStream output) throws IOException, InterruptedException, StreamingResponseCanceledException { 2035 byteBuffer.flip(); 2036 if (byteBuffer.hasRemaining()) 2037 output.write(byteBuffer); 2038 byteBuffer.clear(); 2039 } 2040 2041 private void materializePublisher(com.soklet.StreamingResponseBody.@NonNull PublisherBody publisherBody, 2042 @NonNull SimulatorResponseStream output, 2043 @NonNull SimulatorStreamingResponseContext context) throws Exception { 2044 requireNonNull(publisherBody); 2045 requireNonNull(output); 2046 requireNonNull(context); 2047 2048 CountDownLatch completed = new CountDownLatch(1); 2049 AtomicBoolean publisherTerminated = new AtomicBoolean(false); 2050 AtomicReference<Throwable> failure = new AtomicReference<>(); 2051 AtomicReference<Flow.Subscription> subscriptionRef = new AtomicReference<>(); 2052 2053 try (AutoCloseable cancelationRegistration = context.onCancel(() -> { 2054 Flow.Subscription subscription = subscriptionRef.get(); 2055 2056 if (subscription != null) 2057 subscription.cancel(); 2058 })) { 2059 publisherBody.getPublisher().subscribe(new Flow.Subscriber<>() { 2060 @Override 2061 public void onSubscribe(Flow.Subscription subscription) { 2062 requireNonNull(subscription); 2063 2064 if (!subscriptionRef.compareAndSet(null, subscription)) { 2065 subscription.cancel(); 2066 return; 2067 } 2068 2069 subscription.request(1L); 2070 } 2071 2072 @Override 2073 public void onNext(ByteBuffer item) { 2074 Flow.Subscription subscription = subscriptionRef.get(); 2075 2076 try { 2077 context.throwIfCanceled(); 2078 output.write(requireNonNull(item)); 2079 context.throwIfCanceled(); 2080 } catch (Throwable t) { 2081 failure.compareAndSet(null, t); 2082 publisherTerminated.set(true); 2083 2084 if (subscription != null) 2085 subscription.cancel(); 2086 2087 completed.countDown(); 2088 return; 2089 } 2090 2091 if (subscription != null) 2092 subscription.request(1L); 2093 } 2094 2095 @Override 2096 public void onError(Throwable throwable) { 2097 publisherTerminated.set(true); 2098 failure.compareAndSet(null, throwable == null 2099 ? new IllegalStateException("Publisher failed without an error") 2100 : throwable); 2101 completed.countDown(); 2102 } 2103 2104 @Override 2105 public void onComplete() { 2106 publisherTerminated.set(true); 2107 completed.countDown(); 2108 } 2109 }); 2110 2111 while (!completed.await(100L, TimeUnit.MILLISECONDS)) 2112 context.throwIfCanceled(); 2113 } finally { 2114 if (!publisherTerminated.get()) { 2115 Flow.Subscription subscription = subscriptionRef.get(); 2116 2117 if (subscription != null) 2118 subscription.cancel(); 2119 } 2120 } 2121 2122 Throwable throwable = failure.get(); 2123 2124 if (throwable != null) { 2125 if (throwable instanceof Exception exception) 2126 throw exception; 2127 2128 if (throwable instanceof Error error) 2129 throw error; 2130 2131 throw new RuntimeException(throwable); 2132 } 2133 } 2134 2135 private void closeQuietly(@NonNull AutoCloseable closeable) { 2136 requireNonNull(closeable); 2137 2138 try { 2139 closeable.close(); 2140 } catch (InterruptedException e) { 2141 Thread.currentThread().interrupt(); 2142 } catch (Throwable ignored) { 2143 // Best effort only. The producer will observe cancelation separately. 2144 } 2145 } 2146 2147 @NonNull 2148 @Override 2149 public SseRequestResult performSseRequest(@NonNull Request request) { 2150 MockSseServer sseServer = getSseServer().orElse(null); 2151 2152 if (sseServer == null) 2153 throw new IllegalStateException(format("You must specify a %s in your %s to simulate Server-Sent Event requests", 2154 SseServer.class.getSimpleName(), SokletConfig.class.getSimpleName())); 2155 2156 AtomicReference<HttpRequestResult> requestResultHolder = new AtomicReference<>(); 2157 SseServer.RequestHandler requestHandler = sseServer.getRequestHandler().orElse(null); 2158 2159 if (requestHandler == null) 2160 throw new IllegalStateException("You must register a request handler prior to simulating SSE Event Source requests"); 2161 2162 requestHandler.handleRequest(request, (requestResult -> { 2163 requestResultHolder.set(requestResult); 2164 })); 2165 2166 HttpRequestResult requestResult = requestResultHolder.get(); 2167 if (requestResult == null) 2168 throw new IllegalStateException("SSE request handler did not provide a request result"); 2169 2170 SseHandshakeResult sseHandshakeResult = requestResult.getSseHandshakeResult().orElse(null); 2171 2172 if (sseHandshakeResult == null) 2173 return new SseRequestResult.RequestFailed(requestResult); 2174 2175 if (sseHandshakeResult instanceof SseHandshakeResult.Accepted acceptedHandshake) { 2176 Consumer<SseUnicaster> clientInitializer = acceptedHandshake.getClientInitializer().orElse(null); 2177 2178 // Create a synthetic logical response using values from the accepted handshake 2179 if (requestResult.getResponse().isEmpty()) 2180 requestResult = requestResult.copy() 2181 .response(Response.withStatusCode(200) 2182 .headers(acceptedHandshake.getHeaders()) 2183 .cookies(acceptedHandshake.getCookies()) 2184 .build()) 2185 .finish(); 2186 2187 HandshakeAccepted handshakeAccepted = new HandshakeAccepted(acceptedHandshake, request.getResourcePath(), requestResult, this, clientInitializer); 2188 return handshakeAccepted; 2189 } 2190 2191 if (sseHandshakeResult instanceof SseHandshakeResult.Rejected rejectedHandshake) 2192 return new HandshakeRejected(rejectedHandshake, requestResult); 2193 2194 throw new IllegalStateException(format("Encountered unexpected %s: %s", SseHandshakeResult.class.getSimpleName(), sseHandshakeResult)); 2195 } 2196 2197 @NonNull 2198 @Override 2199 public McpRequestResult performMcpRequest(@NonNull Request request) { 2200 requireNonNull(request); 2201 2202 MockMcpServer mcpServer = getMcpServer().orElse(null); 2203 2204 if (mcpServer == null) 2205 throw new IllegalStateException(format("You must specify an MCP server in your %s to simulate MCP requests", 2206 SokletConfig.class.getSimpleName())); 2207 2208 AtomicReference<HttpRequestResult> requestResultHolder = new AtomicReference<>(); 2209 McpServer.RequestHandler requestHandler = mcpServer.getRequestHandler().orElse(null); 2210 2211 if (requestHandler == null) 2212 throw new IllegalStateException("You must register a request handler prior to simulating MCP requests"); 2213 2214 requestHandler.handleRequest(request, requestResultHolder::set); 2215 2216 HttpRequestResult requestResult = requestResultHolder.get(); 2217 2218 if (requestResult == null) 2219 throw new IllegalStateException("No MCP request result was produced by the simulator"); 2220 2221 if (extractContentTypeFromHeaders(requestResult.getMarshaledResponse().getHeaders()) 2222 .filter(contentType -> contentType.equalsIgnoreCase("text/event-stream")) 2223 .isPresent()) { 2224 McpRequestResult.StreamOpened streamOpened = new McpRequestResult.StreamOpened( 2225 requestResult, 2226 mcpServer.getMcpStreamErrorHandler(), 2227 requestResult.isMcpStreamClosedAfterReplay()); 2228 2229 for (McpObject mcpStreamMessage : requestResult.getMcpStreamMessages()) 2230 streamOpened.emitMessage(mcpStreamMessage); 2231 2232 if (!requestResult.isMcpStreamClosedAfterReplay()) 2233 request.getHeader("MCP-Session-Id").ifPresent(sessionId -> mcpServer.registerOpenStream(sessionId, request, streamOpened)); 2234 2235 return streamOpened; 2236 } 2237 2238 if (request.getHttpMethod() == HttpMethod.DELETE 2239 && Objects.equals(requestResult.getMarshaledResponse().getStatusCode(), 204)) 2240 request.getHeader("MCP-Session-Id").ifPresent(mcpServer::terminateStreamsForSession); 2241 2242 return new McpRequestResult.ResponseCompleted(requestResult); 2243 } 2244 2245 @NonNull 2246 @Override 2247 public Simulator onBroadcastError(@Nullable Consumer<Throwable> onBroadcastError) { 2248 MockSseServer sseServer = getSseServer().orElse(null); 2249 2250 if (sseServer != null) 2251 sseServer.onBroadcastError(onBroadcastError); 2252 2253 return this; 2254 } 2255 2256 @NonNull 2257 @Override 2258 public Simulator onUnicastError(@Nullable Consumer<Throwable> onUnicastError) { 2259 MockSseServer sseServer = getSseServer().orElse(null); 2260 2261 if (sseServer != null) 2262 sseServer.onUnicastError(onUnicastError); 2263 2264 return this; 2265 } 2266 2267 @NonNull 2268 @Override 2269 public Simulator onMcpStreamError(@Nullable Consumer<Throwable> onMcpStreamError) { 2270 MockMcpServer mcpServer = getMcpServer().orElse(null); 2271 2272 if (mcpServer != null) 2273 mcpServer.onMcpStreamError(onMcpStreamError); 2274 2275 return this; 2276 } 2277 2278 @NonNull 2279 Optional<MockHttpServer> getHttpServer() { 2280 return Optional.ofNullable(this.server); 2281 } 2282 2283 @NonNull 2284 Optional<MockSseServer> getSseServer() { 2285 return Optional.ofNullable(this.sseServer); 2286 } 2287 2288 @NonNull 2289 Optional<MockMcpServer> getMcpServer() { 2290 return Optional.ofNullable(this.mcpServer); 2291 } 2292 2293 @NonNull 2294 SimulatorOptions getSimulatorOptions() { 2295 return this.simulatorOptions; 2296 } 2297 } 2298 2299 @NotThreadSafe 2300 private static final class SimulatorResponseStream implements ResponseStream { 2301 @NonNull 2302 private final ByteArrayOutputStream byteArrayOutputStream; 2303 @NonNull 2304 private final Integer limitInBytes; 2305 @NonNull 2306 private final SimulatorCancelationToken cancelationToken; 2307 private boolean closed; 2308 2309 private SimulatorResponseStream(@NonNull Integer limitInBytes, 2310 @NonNull SimulatorCancelationToken cancelationToken) { 2311 this.byteArrayOutputStream = new ByteArrayOutputStream(); 2312 this.limitInBytes = requireNonNull(limitInBytes); 2313 this.cancelationToken = requireNonNull(cancelationToken); 2314 } 2315 2316 @Override 2317 public void write(@NonNull byte[] bytes) throws IOException, StreamingResponseCanceledException { 2318 requireNonNull(bytes); 2319 write(ByteBuffer.wrap(bytes)); 2320 } 2321 2322 @Override 2323 public void write(@NonNull ByteBuffer byteBuffer) throws IOException, StreamingResponseCanceledException { 2324 requireNonNull(byteBuffer); 2325 this.cancelationToken.throwIfCanceled(); 2326 2327 if (this.closed) 2328 throw new StreamingResponseCanceledException(StreamTerminationReason.APPLICATION_CANCELED); 2329 2330 ByteBuffer source = byteBuffer.asReadOnlyBuffer(); 2331 int bytesToWrite = source.remaining(); 2332 2333 if ((long) this.byteArrayOutputStream.size() + bytesToWrite > this.limitInBytes) { 2334 this.cancelationToken.cancel(StreamTerminationReason.SIMULATOR_LIMIT_EXCEEDED, null); 2335 throw new StreamingResponseCanceledException(StreamTerminationReason.SIMULATOR_LIMIT_EXCEEDED); 2336 } 2337 2338 byte[] bytes = new byte[bytesToWrite]; 2339 source.get(bytes); 2340 this.byteArrayOutputStream.write(bytes); 2341 } 2342 2343 @Override 2344 public void flush() throws StreamingResponseCanceledException { 2345 this.cancelationToken.throwIfCanceled(); 2346 } 2347 2348 @Override 2349 @NonNull 2350 public Boolean isOpen() { 2351 return !this.closed && !this.cancelationToken.isCanceled(); 2352 } 2353 2354 @NonNull 2355 private byte[] toByteArray() { 2356 this.closed = true; 2357 return this.byteArrayOutputStream.toByteArray(); 2358 } 2359 } 2360 2361 @ThreadSafe 2362 private static final class SimulatorCancelationToken implements CancelationToken { 2363 @NonNull 2364 private final AtomicBoolean canceled; 2365 @NonNull 2366 private final CopyOnWriteArrayList<Runnable> callbacks; 2367 @NonNull 2368 private final Consumer<Throwable> callbackFailureConsumer; 2369 @Nullable 2370 private volatile StreamTerminationReason reason; 2371 @Nullable 2372 private volatile Throwable cause; 2373 2374 private SimulatorCancelationToken(@NonNull Consumer<Throwable> callbackFailureConsumer) { 2375 this.canceled = new AtomicBoolean(false); 2376 this.callbacks = new CopyOnWriteArrayList<>(); 2377 this.callbackFailureConsumer = requireNonNull(callbackFailureConsumer); 2378 } 2379 2380 @Override 2381 @NonNull 2382 public Boolean isCanceled() { 2383 return this.canceled.get(); 2384 } 2385 2386 @Override 2387 @NonNull 2388 public Optional<StreamTerminationReason> getCancelationReason() { 2389 return Optional.ofNullable(this.reason); 2390 } 2391 2392 @Override 2393 @NonNull 2394 public Optional<Throwable> getCancelationCause() { 2395 return Optional.ofNullable(this.cause); 2396 } 2397 2398 @Override 2399 @NonNull 2400 public AutoCloseable onCancel(@NonNull Runnable callback) { 2401 requireNonNull(callback); 2402 2403 boolean runImmediately; 2404 2405 synchronized (this) { 2406 runImmediately = this.canceled.get(); 2407 2408 if (!runImmediately) 2409 this.callbacks.add(callback); 2410 } 2411 2412 if (runImmediately) { 2413 runCallback(callback); 2414 return () -> { 2415 // No-op 2416 }; 2417 } 2418 2419 return () -> { 2420 synchronized (this) { 2421 this.callbacks.remove(callback); 2422 } 2423 }; 2424 } 2425 2426 private boolean cancel(@NonNull StreamTerminationReason reason, 2427 @Nullable Throwable cause) { 2428 requireNonNull(reason); 2429 2430 if (reason == StreamTerminationReason.COMPLETED) 2431 throw new IllegalArgumentException("Cancelation reason cannot be COMPLETED"); 2432 2433 List<Runnable> callbacksToRun; 2434 2435 synchronized (this) { 2436 if (this.canceled.get()) 2437 return false; 2438 2439 this.reason = reason; 2440 this.cause = cause; 2441 this.canceled.set(true); 2442 callbacksToRun = List.copyOf(this.callbacks); 2443 this.callbacks.clear(); 2444 } 2445 2446 for (Runnable callback : callbacksToRun) 2447 runCallback(callback); 2448 2449 return true; 2450 } 2451 2452 private void runCallback(@NonNull Runnable callback) { 2453 requireNonNull(callback); 2454 2455 try { 2456 callback.run(); 2457 } catch (Throwable t) { 2458 this.callbackFailureConsumer.accept(t); 2459 } 2460 } 2461 } 2462 2463 @ThreadSafe 2464 private static final class SimulatorStreamingResponseContext implements StreamingResponseContext { 2465 @NonNull 2466 private final Request request; 2467 @NonNull 2468 private final CancelationToken cancelationToken; 2469 2470 private SimulatorStreamingResponseContext(@NonNull Request request, 2471 @NonNull CancelationToken cancelationToken) { 2472 this.request = requireNonNull(request); 2473 this.cancelationToken = requireNonNull(cancelationToken); 2474 } 2475 2476 @Override 2477 @NonNull 2478 public CancelationToken getCancelationToken() { 2479 return this.cancelationToken; 2480 } 2481 2482 @Override 2483 @NonNull 2484 public Request getRequest() { 2485 return this.request; 2486 } 2487 2488 @Override 2489 @NonNull 2490 public Optional<Instant> getDeadline() { 2491 return Optional.empty(); 2492 } 2493 2494 @Override 2495 @NonNull 2496 public Optional<Duration> getIdleTimeout() { 2497 return Optional.empty(); 2498 } 2499 } 2500 2501 /** 2502 * Mock server that doesn't touch the network at all, useful for testing. 2503 * 2504 * @author <a href="https://www.revetkn.com">Mark Allen</a> 2505 */ 2506 @ThreadSafe 2507 static class MockHttpServer implements HttpServer { 2508 @Nullable 2509 private SokletConfig sokletConfig; 2510 private HttpServer.@Nullable RequestHandler requestHandler; 2511 2512 @Override 2513 public void start() { 2514 // No-op 2515 } 2516 2517 @Override 2518 public void stop() { 2519 // No-op 2520 } 2521 2522 @NonNull 2523 @Override 2524 public Boolean isStarted() { 2525 return true; 2526 } 2527 2528 @Override 2529 public void initialize(@NonNull SokletConfig sokletConfig, 2530 @NonNull RequestHandler requestHandler) { 2531 requireNonNull(sokletConfig); 2532 requireNonNull(requestHandler); 2533 2534 this.sokletConfig = sokletConfig; 2535 this.requestHandler = requestHandler; 2536 } 2537 2538 @NonNull 2539 protected Optional<SokletConfig> getSokletConfig() { 2540 return Optional.ofNullable(this.sokletConfig); 2541 } 2542 2543 @NonNull 2544 protected Optional<RequestHandler> getRequestHandler() { 2545 return Optional.ofNullable(this.requestHandler); 2546 } 2547 } 2548 2549 /** 2550 * Mock MCP server that doesn't touch the network at all, useful for testing. 2551 */ 2552 @ThreadSafe 2553 static class MockMcpServer implements McpServer, InternalMcpSessionMessagePublisher { 2554 @NonNull 2555 private final McpServer realImplementation; 2556 @Nullable 2557 private SokletConfig sokletConfig; 2558 private McpServer.@Nullable RequestHandler requestHandler; 2559 @NonNull 2560 private final AtomicReference<Consumer<Throwable>> mcpStreamErrorHandler; 2561 @NonNull 2562 private final ConcurrentHashMap<@NonNull String, @NonNull CopyOnWriteArrayList<McpRequestResult.StreamOpened>> openStreamsBySessionId; 2563 @NonNull 2564 private final AtomicReference<@Nullable BiConsumer<Request, String>> clientDisconnectedMcpStreamHandler; 2565 2566 public MockMcpServer(@NonNull McpServer realImplementation) { 2567 requireNonNull(realImplementation); 2568 2569 this.realImplementation = realImplementation; 2570 this.mcpStreamErrorHandler = new AtomicReference<>(); 2571 this.openStreamsBySessionId = new ConcurrentHashMap<>(); 2572 this.clientDisconnectedMcpStreamHandler = new AtomicReference<>(); 2573 } 2574 2575 @Override 2576 public void start() { 2577 // No-op 2578 } 2579 2580 @Override 2581 public void stop() { 2582 // No-op 2583 } 2584 2585 @NonNull 2586 @Override 2587 public Boolean isStarted() { 2588 return true; 2589 } 2590 2591 @Override 2592 public void initialize(@NonNull SokletConfig sokletConfig, 2593 @NonNull RequestHandler requestHandler) { 2594 requireNonNull(sokletConfig); 2595 requireNonNull(requestHandler); 2596 2597 this.sokletConfig = sokletConfig; 2598 this.requestHandler = requestHandler; 2599 } 2600 2601 @NonNull 2602 @Override 2603 public McpHandlerResolver getHandlerResolver() { 2604 return getRealImplementation().getHandlerResolver(); 2605 } 2606 2607 @NonNull 2608 @Override 2609 public McpRequestAdmissionPolicy getRequestAdmissionPolicy() { 2610 return getRealImplementation().getRequestAdmissionPolicy(); 2611 } 2612 2613 @NonNull 2614 @Override 2615 public McpRequestInterceptor getRequestInterceptor() { 2616 return getRealImplementation().getRequestInterceptor(); 2617 } 2618 2619 @NonNull 2620 @Override 2621 public McpResponseMarshaler getResponseMarshaler() { 2622 return getRealImplementation().getResponseMarshaler(); 2623 } 2624 2625 @NonNull 2626 @Override 2627 public McpCorsAuthorizer getCorsAuthorizer() { 2628 return getRealImplementation().getCorsAuthorizer(); 2629 } 2630 2631 @NonNull 2632 @Override 2633 public McpSessionStore getSessionStore() { 2634 return getRealImplementation().getSessionStore(); 2635 } 2636 2637 @NonNull 2638 protected McpServer getRealImplementation() { 2639 return this.realImplementation; 2640 } 2641 2642 @NonNull 2643 protected Optional<SokletConfig> getSokletConfig() { 2644 return Optional.ofNullable(this.sokletConfig); 2645 } 2646 2647 @NonNull 2648 protected Optional<RequestHandler> getRequestHandler() { 2649 return Optional.ofNullable(this.requestHandler); 2650 } 2651 2652 protected void onMcpStreamError(@Nullable Consumer<Throwable> onMcpStreamError) { 2653 this.mcpStreamErrorHandler.set(onMcpStreamError); 2654 } 2655 2656 @NonNull 2657 protected AtomicReference<Consumer<Throwable>> getMcpStreamErrorHandler() { 2658 return this.mcpStreamErrorHandler; 2659 } 2660 2661 protected void registerOpenStream(@NonNull String sessionId, 2662 @NonNull Request request, 2663 McpRequestResult.StreamOpened streamOpened) { 2664 requireNonNull(sessionId); 2665 requireNonNull(request); 2666 requireNonNull(streamOpened); 2667 2668 getOpenStreamsBySessionId() 2669 .computeIfAbsent(sessionId, ignored -> new CopyOnWriteArrayList<>()) 2670 .add(streamOpened); 2671 2672 streamOpened.onClose(() -> closeOpenStream(sessionId, request, streamOpened)); 2673 } 2674 2675 protected void terminateStreamsForSession(@NonNull String sessionId) { 2676 requireNonNull(sessionId); 2677 2678 CopyOnWriteArrayList<McpRequestResult.StreamOpened> streams = getOpenStreamsBySessionId().remove(sessionId); 2679 2680 if (streams == null) 2681 return; 2682 2683 for (McpRequestResult.StreamOpened streamOpened : streams) 2684 streamOpened.terminate(); 2685 } 2686 2687 @NonNull 2688 @Override 2689 public Boolean publishSessionMessage(@NonNull String sessionId, 2690 @NonNull McpObject message) { 2691 requireNonNull(sessionId); 2692 requireNonNull(message); 2693 2694 CopyOnWriteArrayList<McpRequestResult.StreamOpened> streams = getOpenStreamsBySessionId().get(sessionId); 2695 2696 if (streams == null || streams.isEmpty()) 2697 return false; 2698 2699 for (int i = streams.size() - 1; i >= 0; i--) { 2700 McpRequestResult.StreamOpened streamOpened = streams.get(i); 2701 2702 if (streamOpened.isClosed()) 2703 continue; 2704 2705 streamOpened.emitMessage(message); 2706 return true; 2707 } 2708 2709 return false; 2710 } 2711 2712 protected void onClientDisconnectedMcpStream(@Nullable BiConsumer<Request, String> clientDisconnectedMcpStreamHandler) { 2713 this.clientDisconnectedMcpStreamHandler.set(clientDisconnectedMcpStreamHandler); 2714 } 2715 2716 protected void closeOpenStream(@NonNull String sessionId, 2717 @NonNull Request request, 2718 McpRequestResult.StreamOpened streamOpened) { 2719 requireNonNull(sessionId); 2720 requireNonNull(request); 2721 requireNonNull(streamOpened); 2722 2723 CopyOnWriteArrayList<McpRequestResult.StreamOpened> streams = getOpenStreamsBySessionId().get(sessionId); 2724 2725 if (streams != null) { 2726 streams.remove(streamOpened); 2727 2728 if (streams.isEmpty()) 2729 getOpenStreamsBySessionId().remove(sessionId, streams); 2730 } 2731 2732 BiConsumer<Request, String> handler = this.clientDisconnectedMcpStreamHandler.get(); 2733 2734 if (handler != null) 2735 handler.accept(request, sessionId); 2736 } 2737 2738 @NonNull 2739 protected ConcurrentHashMap<@NonNull String, @NonNull CopyOnWriteArrayList<McpRequestResult.StreamOpened>> getOpenStreamsBySessionId() { 2740 return this.openStreamsBySessionId; 2741 } 2742 } 2743 2744 /** 2745 * Mock Server-Sent Event unicaster that doesn't touch the network at all, useful for testing. 2746 */ 2747 @ThreadSafe 2748 static class MockSseUnicaster implements SseUnicaster { 2749 @NonNull 2750 private final ResourcePath resourcePath; 2751 @NonNull 2752 private final Consumer<SseEvent> eventConsumer; 2753 @NonNull 2754 private final Consumer<SseComment> commentConsumer; 2755 @NonNull 2756 private final AtomicReference<Consumer<Throwable>> unicastErrorHandler; 2757 @NonNull 2758 private final Consumer<LogEvent> logEventConsumer; 2759 2760 public MockSseUnicaster(@NonNull ResourcePath resourcePath, 2761 @NonNull Consumer<SseEvent> eventConsumer, 2762 @NonNull Consumer<SseComment> commentConsumer, 2763 @NonNull AtomicReference<Consumer<Throwable>> unicastErrorHandler, 2764 @NonNull Consumer<LogEvent> logEventConsumer) { 2765 requireNonNull(resourcePath); 2766 requireNonNull(eventConsumer); 2767 requireNonNull(commentConsumer); 2768 requireNonNull(unicastErrorHandler); 2769 requireNonNull(logEventConsumer); 2770 2771 this.resourcePath = resourcePath; 2772 this.eventConsumer = eventConsumer; 2773 this.commentConsumer = commentConsumer; 2774 this.unicastErrorHandler = unicastErrorHandler; 2775 this.logEventConsumer = logEventConsumer; 2776 } 2777 2778 @Override 2779 public void unicastEvent(@NonNull SseEvent sseEvent) { 2780 requireNonNull(sseEvent); 2781 try { 2782 getEventConsumer().accept(sseEvent); 2783 } catch (Throwable throwable) { 2784 handleUnicastError(throwable); 2785 } 2786 } 2787 2788 @Override 2789 public void unicastComment(@NonNull SseComment sseComment) { 2790 requireNonNull(sseComment); 2791 try { 2792 getCommentConsumer().accept(sseComment); 2793 } catch (Throwable throwable) { 2794 handleUnicastError(throwable); 2795 } 2796 } 2797 2798 @NonNull 2799 @Override 2800 public ResourcePath getResourcePath() { 2801 return this.resourcePath; 2802 } 2803 2804 @NonNull 2805 protected Consumer<SseEvent> getEventConsumer() { 2806 return this.eventConsumer; 2807 } 2808 2809 @NonNull 2810 protected Consumer<SseComment> getCommentConsumer() { 2811 return this.commentConsumer; 2812 } 2813 2814 protected void handleUnicastError(@NonNull Throwable throwable) { 2815 requireNonNull(throwable); 2816 Consumer<Throwable> handler = this.unicastErrorHandler.get(); 2817 2818 if (handler != null) { 2819 try { 2820 handler.accept(throwable); 2821 return; 2822 } catch (Throwable ignored) { 2823 // Fall through to default behavior 2824 } 2825 } 2826 2827 safelyLog(LogEvent.with(LogEventType.SSE_SERVER_INTERNAL_ERROR, 2828 "SSE simulator unicast consumer failed") 2829 .throwable(throwable) 2830 .build()); 2831 } 2832 2833 protected void safelyLog(@NonNull LogEvent logEvent) { 2834 requireNonNull(logEvent); 2835 2836 try { 2837 this.logEventConsumer.accept(logEvent); 2838 } catch (Throwable ignored) { 2839 // No safe fallback sink is available here. 2840 } 2841 } 2842 } 2843 2844 /** 2845 * Mock Server-Sent Event broadcaster that doesn't touch the network at all, useful for testing. 2846 */ 2847 @ThreadSafe 2848 static class MockSseBroadcaster implements SseBroadcaster { 2849 // ConcurrentHashMap doesn't allow null values, so we use a sentinel if context is null 2850 private static final Object NULL_CONTEXT_SENTINEL; 2851 2852 static { 2853 NULL_CONTEXT_SENTINEL = new Object(); 2854 } 2855 2856 @NonNull 2857 private final ResourcePath resourcePath; 2858 // Maps the Consumer (Listener) to its Context object (e.g. Locale) 2859 @NonNull 2860 private final Map<@NonNull Consumer<SseEvent>, @NonNull Object> eventConsumers; 2861 // Same goes for comments 2862 @NonNull 2863 private final Map<@NonNull Consumer<SseComment>, @NonNull Object> commentConsumers; 2864 @NonNull 2865 private final AtomicReference<Consumer<Throwable>> broadcastErrorHandler; 2866 @NonNull 2867 private final Consumer<LogEvent> logEventConsumer; 2868 2869 public MockSseBroadcaster(@NonNull ResourcePath resourcePath, 2870 @NonNull AtomicReference<Consumer<Throwable>> broadcastErrorHandler, 2871 @NonNull Consumer<LogEvent> logEventConsumer) { 2872 requireNonNull(resourcePath); 2873 requireNonNull(broadcastErrorHandler); 2874 requireNonNull(logEventConsumer); 2875 2876 this.resourcePath = resourcePath; 2877 this.eventConsumers = new ConcurrentHashMap<>(); 2878 this.commentConsumers = new ConcurrentHashMap<>(); 2879 this.broadcastErrorHandler = broadcastErrorHandler; 2880 this.logEventConsumer = logEventConsumer; 2881 } 2882 2883 @NonNull 2884 @Override 2885 public ResourcePath getResourcePath() { 2886 return this.resourcePath; 2887 } 2888 2889 @NonNull 2890 @Override 2891 public Long getClientCount() { 2892 return Long.valueOf(getEventConsumers().size() + getCommentConsumers().size()); 2893 } 2894 2895 @Override 2896 public void broadcastEvent(@NonNull SseEvent sseEvent) { 2897 requireNonNull(sseEvent); 2898 2899 for (Consumer<SseEvent> eventConsumer : getEventConsumers().keySet()) { 2900 try { 2901 eventConsumer.accept(sseEvent); 2902 } catch (Throwable throwable) { 2903 handleBroadcastError(throwable); 2904 } 2905 } 2906 } 2907 2908 @Override 2909 public void broadcastComment(@NonNull SseComment sseComment) { 2910 requireNonNull(sseComment); 2911 2912 for (Consumer<SseComment> commentConsumer : getCommentConsumers().keySet()) { 2913 try { 2914 commentConsumer.accept(sseComment); 2915 } catch (Throwable throwable) { 2916 handleBroadcastError(throwable); 2917 } 2918 } 2919 } 2920 2921 @Override 2922 public <T> void broadcastEvent( 2923 @NonNull Function<Object, T> keySelector, 2924 @NonNull Function<T, SseEvent> eventProvider 2925 ) { 2926 requireNonNull(keySelector); 2927 requireNonNull(eventProvider); 2928 2929 // 1. Create a temporary cache for this specific broadcast operation. 2930 // This ensures we only run the expensive 'eventProvider' once per unique key. 2931 Map<T, SseEvent> payloadCache = new HashMap<>(); 2932 2933 this.getEventConsumers().forEach((consumer, context) -> { 2934 try { 2935 // 2. Derive the key from the subscriber's context 2936 T key = keySelector.apply(context); 2937 2938 // 3. Memoize: Generate the payload if we haven't seen this key yet, otherwise reuse it 2939 SseEvent event = payloadCache.computeIfAbsent(key, eventProvider); 2940 2941 // 4. Dispatch 2942 consumer.accept(event); 2943 } catch (Throwable throwable) { 2944 handleBroadcastError(throwable); 2945 } 2946 }); 2947 } 2948 2949 @Override 2950 public <T> void broadcastComment( 2951 @NonNull Function<Object, T> keySelector, 2952 @NonNull Function<T, SseComment> commentProvider 2953 ) { 2954 requireNonNull(keySelector); 2955 requireNonNull(commentProvider); 2956 2957 // 1. Create temporary cache 2958 Map<T, SseComment> commentCache = new HashMap<>(); 2959 2960 this.getCommentConsumers().forEach((consumer, context) -> { 2961 try { 2962 // 2. Derive key 2963 T key = keySelector.apply(context); 2964 2965 // 3. Memoize 2966 SseComment comment = commentCache.computeIfAbsent(key, commentProvider); 2967 2968 // 4. Dispatch 2969 consumer.accept(comment); 2970 } catch (Throwable throwable) { 2971 handleBroadcastError(throwable); 2972 } 2973 }); 2974 } 2975 2976 @NonNull 2977 public Boolean registerEventConsumer(@NonNull Consumer<SseEvent> eventConsumer) { 2978 return registerEventConsumer(eventConsumer, null); 2979 } 2980 2981 /** 2982 * Registers a consumer with an associated context, simulating a client with specific traits. 2983 */ 2984 @NonNull 2985 public Boolean registerEventConsumer(@NonNull Consumer<SseEvent> eventConsumer, @Nullable Object context) { 2986 requireNonNull(eventConsumer); 2987 // map.put returns null if the key was new, which conceptually matches "add" returning true 2988 return this.getEventConsumers().put(eventConsumer, context == null ? NULL_CONTEXT_SENTINEL : context) == null; 2989 } 2990 2991 @NonNull 2992 public Boolean unregisterEventConsumer(@NonNull Consumer<SseEvent> eventConsumer) { 2993 requireNonNull(eventConsumer); 2994 return this.getEventConsumers().remove(eventConsumer) != null; 2995 } 2996 2997 @NonNull 2998 public Boolean registerCommentConsumer(@NonNull Consumer<SseComment> commentConsumer) { 2999 return registerCommentConsumer(commentConsumer, null); 3000 } 3001 3002 /** 3003 * Registers a consumer with an associated context, simulating a client with specific traits. 3004 */ 3005 @NonNull 3006 public Boolean registerCommentConsumer(@NonNull Consumer<SseComment> commentConsumer, @Nullable Object context) { 3007 requireNonNull(commentConsumer); 3008 return this.getCommentConsumers().put(commentConsumer, context == null ? NULL_CONTEXT_SENTINEL : context) == null; 3009 } 3010 3011 @NonNull 3012 public Boolean unregisterCommentConsumer(@NonNull Consumer<SseComment> commentConsumer) { 3013 requireNonNull(commentConsumer); 3014 return this.getCommentConsumers().remove(commentConsumer) != null; 3015 } 3016 3017 @NonNull 3018 protected Map<@NonNull Consumer<SseEvent>, @NonNull Object> getEventConsumers() { 3019 return this.eventConsumers; 3020 } 3021 3022 @NonNull 3023 protected Map<@NonNull Consumer<SseComment>, @NonNull Object> getCommentConsumers() { 3024 return this.commentConsumers; 3025 } 3026 3027 protected void handleBroadcastError(@NonNull Throwable throwable) { 3028 requireNonNull(throwable); 3029 Consumer<Throwable> handler = this.broadcastErrorHandler.get(); 3030 3031 if (handler != null) { 3032 try { 3033 handler.accept(throwable); 3034 return; 3035 } catch (Throwable ignored) { 3036 // Fall through to default behavior 3037 } 3038 } 3039 3040 safelyLog(LogEvent.with(LogEventType.SSE_SERVER_INTERNAL_ERROR, 3041 "SSE simulator broadcast consumer failed") 3042 .throwable(throwable) 3043 .build()); 3044 } 3045 3046 protected void safelyLog(@NonNull LogEvent logEvent) { 3047 requireNonNull(logEvent); 3048 3049 try { 3050 this.logEventConsumer.accept(logEvent); 3051 } catch (Throwable ignored) { 3052 // No safe fallback sink is available here. 3053 } 3054 } 3055 } 3056 3057 /** 3058 * Mock Server-Sent Event server that doesn't touch the network at all, useful for testing. 3059 * 3060 * @author <a href="https://www.revetkn.com">Mark Allen</a> 3061 */ 3062 @ThreadSafe 3063 static class MockSseServer implements SseServer { 3064 @Nullable 3065 private SokletConfig sokletConfig; 3066 private SseServer.@Nullable RequestHandler requestHandler; 3067 @NonNull 3068 private final ConcurrentHashMap<@NonNull ResourcePath, @NonNull MockSseBroadcaster> broadcastersByResourcePath; 3069 @NonNull 3070 private final AtomicReference<Consumer<Throwable>> broadcastErrorHandler; 3071 @NonNull 3072 private final AtomicReference<Consumer<Throwable>> unicastErrorHandler; 3073 3074 public MockSseServer() { 3075 this.broadcastersByResourcePath = new ConcurrentHashMap<>(); 3076 this.broadcastErrorHandler = new AtomicReference<>(); 3077 this.unicastErrorHandler = new AtomicReference<>(); 3078 } 3079 3080 @Override 3081 public void start() { 3082 // No-op 3083 } 3084 3085 @Override 3086 public void stop() { 3087 // No-op 3088 } 3089 3090 @NonNull 3091 @Override 3092 public Boolean isStarted() { 3093 return true; 3094 } 3095 3096 @NonNull 3097 @Override 3098 public Optional<? extends SseBroadcaster> acquireBroadcaster(@Nullable ResourcePath resourcePath) { 3099 if (resourcePath == null) 3100 return Optional.empty(); 3101 3102 MockSseBroadcaster broadcaster = getBroadcastersByResourcePath() 3103 .computeIfAbsent(resourcePath, rp -> new MockSseBroadcaster(rp, broadcastErrorHandler, this::safelyLog)); 3104 3105 return Optional.of(broadcaster); 3106 } 3107 3108 public void registerEventConsumer(@NonNull ResourcePath resourcePath, 3109 @NonNull Consumer<SseEvent> eventConsumer) { 3110 registerEventConsumer(resourcePath, eventConsumer, null); 3111 } 3112 3113 public void registerEventConsumer(@NonNull ResourcePath resourcePath, 3114 @NonNull Consumer<SseEvent> eventConsumer, 3115 @Nullable Object context) { 3116 requireNonNull(resourcePath); 3117 requireNonNull(eventConsumer); 3118 3119 MockSseBroadcaster broadcaster = getBroadcastersByResourcePath() 3120 .computeIfAbsent(resourcePath, rp -> new MockSseBroadcaster(rp, broadcastErrorHandler, this::safelyLog)); 3121 3122 broadcaster.registerEventConsumer(eventConsumer, context); 3123 } 3124 3125 @NonNull 3126 public Boolean unregisterEventConsumer(@NonNull ResourcePath resourcePath, 3127 @NonNull Consumer<SseEvent> eventConsumer) { 3128 requireNonNull(resourcePath); 3129 requireNonNull(eventConsumer); 3130 3131 MockSseBroadcaster broadcaster = getBroadcastersByResourcePath().get(resourcePath); 3132 3133 if (broadcaster == null) 3134 return false; 3135 3136 return broadcaster.unregisterEventConsumer(eventConsumer); 3137 } 3138 3139 public void registerCommentConsumer(@NonNull ResourcePath resourcePath, 3140 @NonNull Consumer<SseComment> commentConsumer) { 3141 registerCommentConsumer(resourcePath, commentConsumer, null); 3142 } 3143 3144 public void registerCommentConsumer(@NonNull ResourcePath resourcePath, 3145 @NonNull Consumer<SseComment> commentConsumer, 3146 @Nullable Object context) { 3147 requireNonNull(resourcePath); 3148 requireNonNull(commentConsumer); 3149 3150 MockSseBroadcaster broadcaster = getBroadcastersByResourcePath() 3151 .computeIfAbsent(resourcePath, rp -> new MockSseBroadcaster(rp, broadcastErrorHandler, this::safelyLog)); 3152 3153 broadcaster.registerCommentConsumer(commentConsumer, context); 3154 } 3155 3156 @NonNull 3157 public Boolean unregisterCommentConsumer(@NonNull ResourcePath resourcePath, 3158 @NonNull Consumer<SseComment> commentConsumer) { 3159 requireNonNull(resourcePath); 3160 requireNonNull(commentConsumer); 3161 3162 MockSseBroadcaster broadcaster = getBroadcastersByResourcePath().get(resourcePath); 3163 3164 if (broadcaster == null) 3165 return false; 3166 3167 return broadcaster.unregisterCommentConsumer(commentConsumer); 3168 } 3169 3170 @Override 3171 public void initialize(@NonNull SokletConfig sokletConfig, 3172 SseServer.@NonNull RequestHandler requestHandler) { 3173 requireNonNull(sokletConfig); 3174 requireNonNull(requestHandler); 3175 3176 this.sokletConfig = sokletConfig; 3177 this.requestHandler = requestHandler; 3178 } 3179 3180 public void onBroadcastError(@Nullable Consumer<Throwable> onBroadcastError) { 3181 this.broadcastErrorHandler.set(onBroadcastError); 3182 } 3183 3184 public void onUnicastError(@Nullable Consumer<Throwable> onUnicastError) { 3185 this.unicastErrorHandler.set(onUnicastError); 3186 } 3187 3188 void safelyLog(@NonNull LogEvent logEvent) { 3189 requireNonNull(logEvent); 3190 3191 SokletConfig sokletConfig = this.sokletConfig; 3192 3193 if (sokletConfig == null) 3194 return; 3195 3196 try { 3197 sokletConfig.getAggregateLifecycleObserver().didReceiveLogEvent(logEvent); 3198 } catch (Throwable ignored) { 3199 // No safe fallback sink is available here. 3200 } 3201 } 3202 3203 @NonNull 3204 protected Optional<SokletConfig> getSokletConfig() { 3205 return Optional.ofNullable(this.sokletConfig); 3206 } 3207 3208 @NonNull 3209 protected Optional<SseServer.RequestHandler> getRequestHandler() { 3210 return Optional.ofNullable(this.requestHandler); 3211 } 3212 3213 @NonNull 3214 protected ConcurrentHashMap<@NonNull ResourcePath, @NonNull MockSseBroadcaster> getBroadcastersByResourcePath() { 3215 return this.broadcastersByResourcePath; 3216 } 3217 3218 @NonNull 3219 protected AtomicReference<Consumer<Throwable>> getUnicastErrorHandler() { 3220 return this.unicastErrorHandler; 3221 } 3222 } 3223 3224}