|
5 | 5 | package io.modelcontextprotocol.client; |
6 | 6 |
|
7 | 7 | import java.time.Duration; |
| 8 | +import java.util.ArrayList; |
8 | 9 | import java.util.HashMap; |
9 | 10 | import java.util.List; |
10 | 11 | import java.util.Map; |
@@ -391,6 +392,114 @@ void shouldHandleTransportSessionNotFoundException() { |
391 | 392 | verify(mockSessionSupplier, times(2)).apply(any(ContextView.class)); |
392 | 393 | } |
393 | 394 |
|
| 395 | + @Test |
| 396 | + void shouldReuseCustomInitializeRequestOnReinitialization() { |
| 397 | + List<McpSchema.InitializeRequest> capturedRequests = new ArrayList<>(); |
| 398 | + |
| 399 | + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenAnswer(invocation -> { |
| 400 | + capturedRequests.add(invocation.getArgument(1)); |
| 401 | + return Mono.just(MOCK_INIT_RESULT); |
| 402 | + }); |
| 403 | + |
| 404 | + McpSchema.InitializeRequest customRequest = McpSchema.InitializeRequest |
| 405 | + .builder("2.0.0", CLIENT_CAPABILITIES, CLIENT_INFO) |
| 406 | + .meta(Map.of("server_id", "proxy-1")) |
| 407 | + .build(); |
| 408 | + |
| 409 | + StepVerifier |
| 410 | + .create(initializer.withInitialization(customRequest, "test", init -> Mono.just(init.initializeResult()))) |
| 411 | + .expectNext(MOCK_INIT_RESULT) |
| 412 | + .verifyComplete(); |
| 413 | + |
| 414 | + initializer.handleException(new McpTransportSessionNotFoundException("Session not found")); |
| 415 | + |
| 416 | + assertThat(capturedRequests).hasSize(2); |
| 417 | + assertThat(capturedRequests.get(0).meta()).containsEntry("server_id", "proxy-1"); |
| 418 | + assertThat(capturedRequests.get(1).meta()).containsEntry("server_id", "proxy-1"); |
| 419 | + verify(mockSessionSupplier, times(2)).apply(any(ContextView.class)); |
| 420 | + } |
| 421 | + |
| 422 | + @Test |
| 423 | + void shouldUseDefaultRequestOnReinitializationWhenNoCustomRequestStored() { |
| 424 | + List<McpSchema.InitializeRequest> capturedRequests = new ArrayList<>(); |
| 425 | + |
| 426 | + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenAnswer(invocation -> { |
| 427 | + capturedRequests.add(invocation.getArgument(1)); |
| 428 | + return Mono.just(MOCK_INIT_RESULT); |
| 429 | + }); |
| 430 | + |
| 431 | + StepVerifier.create(initializer.withInitialization("test", init -> Mono.just(init.initializeResult()))) |
| 432 | + .expectNext(MOCK_INIT_RESULT) |
| 433 | + .verifyComplete(); |
| 434 | + |
| 435 | + initializer.handleException(new McpTransportSessionNotFoundException("Session not found")); |
| 436 | + |
| 437 | + assertThat(capturedRequests).hasSize(2); |
| 438 | + assertThat(capturedRequests.get(1).protocolVersion()).isEqualTo("2.0.0"); |
| 439 | + assertThat(capturedRequests.get(1).capabilities()).isEqualTo(CLIENT_CAPABILITIES); |
| 440 | + assertThat(capturedRequests.get(1).clientInfo()).isEqualTo(CLIENT_INFO); |
| 441 | + assertThat(capturedRequests.get(1).meta()).isNull(); |
| 442 | + } |
| 443 | + |
| 444 | + @Test |
| 445 | + void shouldClearStoredCustomRequestOnClose() { |
| 446 | + AtomicReference<McpSchema.InitializeRequest> capturedRequest = new AtomicReference<>(); |
| 447 | + |
| 448 | + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenAnswer(invocation -> { |
| 449 | + capturedRequest.set(invocation.getArgument(1)); |
| 450 | + return Mono.just(MOCK_INIT_RESULT); |
| 451 | + }); |
| 452 | + |
| 453 | + McpSchema.InitializeRequest customRequest = McpSchema.InitializeRequest |
| 454 | + .builder("2.0.0", CLIENT_CAPABILITIES, CLIENT_INFO) |
| 455 | + .meta(Map.of("server_id", "proxy-1")) |
| 456 | + .build(); |
| 457 | + |
| 458 | + StepVerifier |
| 459 | + .create(initializer.withInitialization(customRequest, "test", init -> Mono.just(init.initializeResult()))) |
| 460 | + .expectNext(MOCK_INIT_RESULT) |
| 461 | + .verifyComplete(); |
| 462 | + |
| 463 | + initializer.close(); |
| 464 | + |
| 465 | + StepVerifier |
| 466 | + .create(initializer.withInitialization("test after close", init -> Mono.just(init.initializeResult()))) |
| 467 | + .expectNext(MOCK_INIT_RESULT) |
| 468 | + .verifyComplete(); |
| 469 | + |
| 470 | + assertThat(capturedRequest.get().meta()).isNull(); |
| 471 | + assertThat(capturedRequest.get().protocolVersion()).isEqualTo("2.0.0"); |
| 472 | + } |
| 473 | + |
| 474 | + @Test |
| 475 | + void shouldClearStoredCustomRequestOnCloseGracefully() { |
| 476 | + AtomicReference<McpSchema.InitializeRequest> capturedRequest = new AtomicReference<>(); |
| 477 | + |
| 478 | + when(mockClientSession.sendRequest(eq(McpSchema.METHOD_INITIALIZE), any(), any())).thenAnswer(invocation -> { |
| 479 | + capturedRequest.set(invocation.getArgument(1)); |
| 480 | + return Mono.just(MOCK_INIT_RESULT); |
| 481 | + }); |
| 482 | + |
| 483 | + McpSchema.InitializeRequest customRequest = McpSchema.InitializeRequest |
| 484 | + .builder("2.0.0", CLIENT_CAPABILITIES, CLIENT_INFO) |
| 485 | + .meta(Map.of("server_id", "proxy-1")) |
| 486 | + .build(); |
| 487 | + |
| 488 | + StepVerifier |
| 489 | + .create(initializer.withInitialization(customRequest, "test", init -> Mono.just(init.initializeResult()))) |
| 490 | + .expectNext(MOCK_INIT_RESULT) |
| 491 | + .verifyComplete(); |
| 492 | + |
| 493 | + StepVerifier.create(initializer.closeGracefully()).verifyComplete(); |
| 494 | + |
| 495 | + StepVerifier |
| 496 | + .create(initializer.withInitialization("test after close", init -> Mono.just(init.initializeResult()))) |
| 497 | + .expectNext(MOCK_INIT_RESULT) |
| 498 | + .verifyComplete(); |
| 499 | + |
| 500 | + assertThat(capturedRequest.get().meta()).isNull(); |
| 501 | + } |
| 502 | + |
394 | 503 | @Test |
395 | 504 | void shouldHandleOtherExceptions() { |
396 | 505 | // Simulate a successful initialization first |
|
0 commit comments