Jeremiah Lowin commited on
Commit
e423e66
·
1 Parent(s): 54bdf2c

Add middleware for all current handlers

Browse files
src/fastmcp/server/middleware.py CHANGED
@@ -148,7 +148,7 @@ class MCPMiddleware:
148
  handler = partial(self.on_list_tools, call_next=handler)
149
  case "resources/list":
150
  handler = partial(self.on_list_resources, call_next=handler)
151
- case "resource-templates/list":
152
  handler = partial(self.on_list_resource_templates, call_next=handler)
153
  case "prompts/list":
154
  handler = partial(self.on_list_prompts, call_next=handler)
 
148
  handler = partial(self.on_list_tools, call_next=handler)
149
  case "resources/list":
150
  handler = partial(self.on_list_resources, call_next=handler)
151
+ case "resources/templates/list":
152
  handler = partial(self.on_list_resource_templates, call_next=handler)
153
  case "prompts/list":
154
  handler = partial(self.on_list_prompts, call_next=handler)
src/fastmcp/server/server.py CHANGED
@@ -352,34 +352,8 @@ class FastMCP(Generic[LifespanResultT]):
352
 
353
  async def get_resources(self) -> dict[str, Resource]:
354
  """Get all registered resources, indexed by registered key."""
355
- if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
356
- resources: dict[str, Resource] = {}
357
-
358
- # iterate such that new mounts overwrite older ones
359
- for mounted_server in self._mounted_servers:
360
- try:
361
- server_resources = await mounted_server.server.get_resources()
362
- # Apply prefix to each resource key if prefix exists
363
- if mounted_server.prefix:
364
- for resource in server_resources.values():
365
- resource = resource.with_key(
366
- add_resource_prefix(
367
- resource.key,
368
- mounted_server.prefix,
369
- self.resource_prefix_format,
370
- )
371
- )
372
- resources[resource.key] = resource
373
- else:
374
- resources.update(server_resources)
375
- except Exception as e:
376
- logger.warning(
377
- f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
378
- )
379
- continue
380
- resources.update(self._resource_manager.get_resources())
381
- self._cache.set("resources", resources)
382
- return resources
383
 
384
  async def get_resource(self, key: str) -> Resource:
385
  resources = await self.get_resources()
@@ -389,39 +363,8 @@ class FastMCP(Generic[LifespanResultT]):
389
 
390
  async def get_resource_templates(self) -> dict[str, ResourceTemplate]:
391
  """Get all registered resource templates, indexed by registered key."""
392
- if (
393
- templates := self._cache.get("resource_templates")
394
- ) is self._cache.NOT_FOUND:
395
- templates: dict[str, ResourceTemplate] = {}
396
-
397
- # iterate such that new mounts overwrite older ones
398
- for mounted_server in self._mounted_servers:
399
- try:
400
- server_templates = (
401
- await mounted_server.server.get_resource_templates()
402
- )
403
- # Apply prefix to each template key if prefix exists
404
- if mounted_server.prefix:
405
- for template in server_templates.values():
406
- template = template.with_key(
407
- add_resource_prefix(
408
- template.key,
409
- mounted_server.prefix,
410
- self.resource_prefix_format,
411
- )
412
- )
413
- templates[template.key] = template
414
- else:
415
- templates.update(server_templates)
416
- except Exception as e:
417
- logger.warning(
418
- "Failed to get resource templates from mounted server "
419
- f"'{mounted_server.prefix}': {e}"
420
- )
421
- continue
422
- templates.update(self._resource_manager.get_templates())
423
- self._cache.set("resource_templates", templates)
424
- return templates
425
 
426
  async def get_resource_template(self, key: str) -> ResourceTemplate:
427
  templates = await self.get_resource_templates()
@@ -434,30 +377,8 @@ class FastMCP(Generic[LifespanResultT]):
434
  List all available prompts.
435
  """
436
 
437
- if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
438
- prompts: dict[str, Prompt] = {}
439
-
440
- # iterate such that new mounts overwrite older ones
441
- for mounted_server in self._mounted_servers:
442
- try:
443
- server_prompts = await mounted_server.server.get_prompts()
444
- # Apply prefix to each prompt key if prefix exists
445
- if mounted_server.prefix:
446
- for prompt in server_prompts.values():
447
- prompt = prompt.with_key(
448
- f"{mounted_server.prefix}_{prompt.key}"
449
- )
450
- prompts[prompt.key] = prompt
451
- else:
452
- prompts.update(server_prompts)
453
- except Exception as e:
454
- logger.warning(
455
- f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
456
- )
457
- continue
458
- prompts.update(self._prompt_manager.get_prompts())
459
- self._cache.set("prompts", prompts)
460
- return prompts
461
 
462
  async def get_prompt(self, key: str) -> Prompt:
463
  prompts = await self.get_prompts()
@@ -582,21 +503,31 @@ class FastMCP(Generic[LifespanResultT]):
582
  return tools
583
 
584
  async def _mcp_list_resources(self) -> list[MCPResource]:
 
 
 
 
 
 
 
 
 
585
  """
586
  List all available resources, in the format expected by the low-level MCP
587
  server.
588
 
589
  """
590
- logger.debug("Handler called: list_resources")
591
 
592
- async def _final_handler(
593
  context: MiddlewareContext[dict[str, Any]],
594
- ) -> list[MCPResource]:
595
- resources = await self.get_resources()
596
- mcp_resources: list[MCPResource] = []
597
- for key, resource in resources.items():
 
598
  if self._should_enable_component(resource):
599
- mcp_resources.append(resource.to_mcp_resource(uri=key))
 
600
  return mcp_resources
601
 
602
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@@ -610,24 +541,74 @@ class FastMCP(Generic[LifespanResultT]):
610
  )
611
 
612
  # Apply the middleware chain.
613
- return await self._apply_middleware(mw_context, _final_handler)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
614
 
615
  async def _mcp_list_resource_templates(self) -> list[MCPResourceTemplate]:
 
 
 
 
 
 
 
 
 
 
616
  """
617
- List all available resource templates, in the format expected by the low-level
618
- MCP server.
619
 
620
  """
621
- logger.debug("Handler called: list_resource_templates")
622
 
623
- async def _final_handler(
624
  context: MiddlewareContext[dict[str, Any]],
625
- ) -> list[MCPResourceTemplate]:
626
- templates = await self.get_resource_templates()
627
- mcp_templates: list[MCPResourceTemplate] = []
628
- for key, template in templates.items():
 
629
  if self._should_enable_component(template):
630
- mcp_templates.append(template.to_mcp_template(uriTemplate=key))
 
631
  return mcp_templates
632
 
633
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@@ -636,29 +617,81 @@ class FastMCP(Generic[LifespanResultT]):
636
  message={}, # List resource templates doesn't have parameters
637
  source="client",
638
  type="request",
639
- method="resources/list_templates",
640
  fastmcp_context=fastmcp_ctx,
641
  )
642
 
643
  # Apply the middleware chain.
644
- return await self._apply_middleware(mw_context, _final_handler)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
645
 
646
  async def _mcp_list_prompts(self) -> list[MCPPrompt]:
 
 
 
 
 
 
 
647
  """
648
  List all available prompts, in the format expected by the low-level MCP
649
  server.
650
 
651
  """
652
- logger.debug("Handler called: list_prompts")
653
 
654
- async def _final_handler(
655
  context: MiddlewareContext[dict[str, Any]],
656
- ) -> list[MCPPrompt]:
657
- prompts = await self.get_prompts()
658
- mcp_prompts: list[MCPPrompt] = []
659
- for key, prompt in prompts.items():
 
660
  if self._should_enable_component(prompt):
661
- mcp_prompts.append(prompt.to_mcp_prompt(name=key))
 
662
  return mcp_prompts
663
 
664
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@@ -672,7 +705,42 @@ class FastMCP(Generic[LifespanResultT]):
672
  )
673
 
674
  # Apply the middleware chain.
675
- return await self._apply_middleware(mw_context, _final_handler)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
676
 
677
  async def _mcp_call_tool(
678
  self, key: str, arguments: dict[str, Any]
 
352
 
353
  async def get_resources(self) -> dict[str, Resource]:
354
  """Get all registered resources, indexed by registered key."""
355
+ resources = await self._list_resources(apply_middleware=False)
356
+ return {resource.key: resource for resource in resources}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
357
 
358
  async def get_resource(self, key: str) -> Resource:
359
  resources = await self.get_resources()
 
363
 
364
  async def get_resource_templates(self) -> dict[str, ResourceTemplate]:
365
  """Get all registered resource templates, indexed by registered key."""
366
+ templates = await self._list_resource_templates(apply_middleware=False)
367
+ return {template.key: template for template in templates}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
368
 
369
  async def get_resource_template(self, key: str) -> ResourceTemplate:
370
  templates = await self.get_resource_templates()
 
377
  List all available prompts.
378
  """
379
 
380
+ prompts = await self._list_prompts(apply_middleware=False)
381
+ return {prompt.key: prompt for prompt in prompts}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
382
 
383
  async def get_prompt(self, key: str) -> Prompt:
384
  prompts = await self.get_prompts()
 
503
  return tools
504
 
505
  async def _mcp_list_resources(self) -> list[MCPResource]:
506
+ logger.debug("Handler called: list_resources")
507
+
508
+ with fastmcp.server.context.Context(fastmcp=self):
509
+ resources = await self._middleware_list_resources()
510
+ return [
511
+ resource.to_mcp_resource(uri=resource.key) for resource in resources
512
+ ]
513
+
514
+ async def _middleware_list_resources(self) -> list[Resource]:
515
  """
516
  List all available resources, in the format expected by the low-level MCP
517
  server.
518
 
519
  """
 
520
 
521
+ async def _handler(
522
  context: MiddlewareContext[dict[str, Any]],
523
+ ) -> list[Resource]:
524
+ resources = await self._list_resources()
525
+
526
+ mcp_resources: list[Resource] = []
527
+ for resource in resources:
528
  if self._should_enable_component(resource):
529
+ mcp_resources.append(resource)
530
+
531
  return mcp_resources
532
 
533
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
 
541
  )
542
 
543
  # Apply the middleware chain.
544
+ return await self._apply_middleware(mw_context, _handler)
545
+
546
+ async def _list_resources(self, apply_middleware: bool = True) -> list[Resource]:
547
+ """
548
+ List all available resources.
549
+ """
550
+
551
+ if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
552
+ resources: list[Resource] = []
553
+
554
+ # iterate such that new mounts overwrite older ones
555
+ for mounted_server in self._mounted_servers:
556
+ try:
557
+ if apply_middleware:
558
+ server_resources = (
559
+ await mounted_server.server._middleware_list_resources()
560
+ )
561
+ else:
562
+ server_resources = await mounted_server.server._list_resources()
563
+ # Apply prefix to each resource key if prefix exists
564
+ if mounted_server.prefix:
565
+ for resource in server_resources:
566
+ resource = resource.with_key(
567
+ add_resource_prefix(
568
+ resource.key,
569
+ mounted_server.prefix,
570
+ self.resource_prefix_format,
571
+ )
572
+ )
573
+ resources.append(resource)
574
+ else:
575
+ resources.extend(server_resources)
576
+ except Exception as e:
577
+ logger.warning(
578
+ f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
579
+ )
580
+ continue
581
+ resources.extend(self._resource_manager.get_resources().values())
582
+ self._cache.set("resources", resources)
583
+ return resources
584
 
585
  async def _mcp_list_resource_templates(self) -> list[MCPResourceTemplate]:
586
+ logger.debug("Handler called: list_resource_templates")
587
+
588
+ with fastmcp.server.context.Context(fastmcp=self):
589
+ templates = await self._middleware_list_resource_templates()
590
+ return [
591
+ template.to_mcp_template(uriTemplate=template.key)
592
+ for template in templates
593
+ ]
594
+
595
+ async def _middleware_list_resource_templates(self) -> list[ResourceTemplate]:
596
  """
597
+ List all available resource templates, in the format expected by the low-level MCP
598
+ server.
599
 
600
  """
 
601
 
602
+ async def _handler(
603
  context: MiddlewareContext[dict[str, Any]],
604
+ ) -> list[ResourceTemplate]:
605
+ templates = await self._list_resource_templates()
606
+
607
+ mcp_templates: list[ResourceTemplate] = []
608
+ for template in templates:
609
  if self._should_enable_component(template):
610
+ mcp_templates.append(template)
611
+
612
  return mcp_templates
613
 
614
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
 
617
  message={}, # List resource templates doesn't have parameters
618
  source="client",
619
  type="request",
620
+ method="resources/templates/list",
621
  fastmcp_context=fastmcp_ctx,
622
  )
623
 
624
  # Apply the middleware chain.
625
+ return await self._apply_middleware(mw_context, _handler)
626
+
627
+ async def _list_resource_templates(
628
+ self, apply_middleware: bool = True
629
+ ) -> list[ResourceTemplate]:
630
+ """
631
+ List all available resource templates.
632
+ """
633
+
634
+ if (
635
+ templates := self._cache.get("resource_templates")
636
+ ) is self._cache.NOT_FOUND:
637
+ templates: list[ResourceTemplate] = []
638
+
639
+ # iterate such that new mounts overwrite older ones
640
+ for mounted_server in self._mounted_servers:
641
+ try:
642
+ if apply_middleware:
643
+ server_templates = await mounted_server.server._middleware_list_resource_templates()
644
+ else:
645
+ server_templates = (
646
+ await mounted_server.server._list_resource_templates()
647
+ )
648
+ # Apply prefix to each template key if prefix exists
649
+ if mounted_server.prefix:
650
+ for template in server_templates:
651
+ template = template.with_key(
652
+ add_resource_prefix(
653
+ template.key,
654
+ mounted_server.prefix,
655
+ self.resource_prefix_format,
656
+ )
657
+ )
658
+ templates.append(template)
659
+ else:
660
+ templates.extend(server_templates)
661
+ except Exception as e:
662
+ logger.warning(
663
+ "Failed to get resource templates from mounted server "
664
+ f"'{mounted_server.prefix}': {e}"
665
+ )
666
+ continue
667
+ templates.extend(self._resource_manager.get_templates().values())
668
+ self._cache.set("resource_templates", templates)
669
+ return templates
670
 
671
  async def _mcp_list_prompts(self) -> list[MCPPrompt]:
672
+ logger.debug("Handler called: list_prompts")
673
+
674
+ with fastmcp.server.context.Context(fastmcp=self):
675
+ prompts = await self._middleware_list_prompts()
676
+ return [prompt.to_mcp_prompt(name=prompt.key) for prompt in prompts]
677
+
678
+ async def _middleware_list_prompts(self) -> list[Prompt]:
679
  """
680
  List all available prompts, in the format expected by the low-level MCP
681
  server.
682
 
683
  """
 
684
 
685
+ async def _handler(
686
  context: MiddlewareContext[dict[str, Any]],
687
+ ) -> list[Prompt]:
688
+ prompts = await self._list_prompts()
689
+
690
+ mcp_prompts: list[Prompt] = []
691
+ for prompt in prompts:
692
  if self._should_enable_component(prompt):
693
+ mcp_prompts.append(prompt)
694
+
695
  return mcp_prompts
696
 
697
  with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
 
705
  )
706
 
707
  # Apply the middleware chain.
708
+ return await self._apply_middleware(mw_context, _handler)
709
+
710
+ async def _list_prompts(self, apply_middleware: bool = True) -> list[Prompt]:
711
+ """
712
+ List all available prompts.
713
+ """
714
+
715
+ if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
716
+ prompts: list[Prompt] = []
717
+
718
+ # iterate such that new mounts overwrite older ones
719
+ for mounted_server in self._mounted_servers:
720
+ try:
721
+ if apply_middleware:
722
+ server_prompts = (
723
+ await mounted_server.server._middleware_list_prompts()
724
+ )
725
+ else:
726
+ server_prompts = await mounted_server.server._list_prompts()
727
+ # Apply prefix to each prompt key if prefix exists
728
+ if mounted_server.prefix:
729
+ for prompt in server_prompts:
730
+ prompt = prompt.with_key(
731
+ f"{mounted_server.prefix}_{prompt.key}"
732
+ )
733
+ prompts.append(prompt)
734
+ else:
735
+ prompts.extend(server_prompts)
736
+ except Exception as e:
737
+ logger.warning(
738
+ f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
739
+ )
740
+ continue
741
+ prompts.extend(self._prompt_manager.get_prompts().values())
742
+ self._cache.set("prompts", prompts)
743
+ return prompts
744
 
745
  async def _mcp_call_tool(
746
  self, key: str, arguments: dict[str, Any]
tests/server/middleware/test_middleware.py CHANGED
@@ -21,9 +21,10 @@ class Recording:
21
  class RecordingMiddleware(MCPMiddleware):
22
  """A middleware that automatically records all method calls."""
23
 
24
- def __init__(self):
25
  super().__init__()
26
  self.calls: list[Recording] = []
 
27
 
28
  def __getattribute__(self, name: str) -> Callable:
29
  """Dynamically create recording methods for any on_* method."""
@@ -95,7 +96,7 @@ class RecordingMiddleware(MCPMiddleware):
95
  @pytest.fixture
96
  def recording_middleware():
97
  """Fixture that provides a recording middleware instance."""
98
- middleware = RecordingMiddleware()
99
  yield middleware
100
 
101
 
@@ -215,7 +216,7 @@ class TestMiddlewareHooks:
215
 
216
  assert recording_middleware.assert_called(times=3)
217
  assert recording_middleware.assert_called(
218
- method="resource-templates/list", times=3
219
  )
220
  assert recording_middleware.assert_called(hook="on_message", times=1)
221
  assert recording_middleware.assert_called(hook="on_request", times=1)
@@ -234,3 +235,258 @@ class TestMiddlewareHooks:
234
  assert recording_middleware.assert_called(hook="on_message", times=1)
235
  assert recording_middleware.assert_called(hook="on_request", times=1)
236
  assert recording_middleware.assert_called(hook="on_list_prompts", times=1)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
21
  class RecordingMiddleware(MCPMiddleware):
22
  """A middleware that automatically records all method calls."""
23
 
24
+ def __init__(self, name: str | None = None):
25
  super().__init__()
26
  self.calls: list[Recording] = []
27
+ self.name = name
28
 
29
  def __getattribute__(self, name: str) -> Callable:
30
  """Dynamically create recording methods for any on_* method."""
 
96
  @pytest.fixture
97
  def recording_middleware():
98
  """Fixture that provides a recording middleware instance."""
99
+ middleware = RecordingMiddleware(name="recording_middleware")
100
  yield middleware
101
 
102
 
 
216
 
217
  assert recording_middleware.assert_called(times=3)
218
  assert recording_middleware.assert_called(
219
+ method="resources/templates/list", times=3
220
  )
221
  assert recording_middleware.assert_called(hook="on_message", times=1)
222
  assert recording_middleware.assert_called(hook="on_request", times=1)
 
235
  assert recording_middleware.assert_called(hook="on_message", times=1)
236
  assert recording_middleware.assert_called(hook="on_request", times=1)
237
  assert recording_middleware.assert_called(hook="on_list_prompts", times=1)
238
+
239
+
240
+ class TestNestedMiddlewareHooks:
241
+ @pytest.fixture
242
+ @staticmethod
243
+ def nested_middleware():
244
+ return RecordingMiddleware(name="nested_middleware")
245
+
246
+ @pytest.fixture
247
+ def nested_mcp_server(self, nested_middleware: RecordingMiddleware):
248
+ mcp = FastMCP(name="Nested MCP")
249
+
250
+ @mcp.tool
251
+ def add(a: int, b: int) -> int:
252
+ return a + b
253
+
254
+ @mcp.resource("resource://test")
255
+ def test_resource() -> str:
256
+ return "test resource"
257
+
258
+ @mcp.resource("resource://test-template/{x}")
259
+ def test_resource_with_path(x: int) -> str:
260
+ return f"test resource with {x}"
261
+
262
+ @mcp.prompt
263
+ def test_prompt(x: str) -> str:
264
+ return f"test prompt with {x}"
265
+
266
+ @mcp.tool
267
+ async def progress_tool(context: Context) -> None:
268
+ await context.report_progress(progress=1, total=10, message="test")
269
+
270
+ @mcp.tool
271
+ async def log_tool(context: Context) -> None:
272
+ await context.info(message="test log")
273
+
274
+ @mcp.tool
275
+ async def sample_tool(context: Context) -> None:
276
+ await context.sample("hello")
277
+
278
+ mcp.add_middleware(nested_middleware)
279
+
280
+ return mcp
281
+
282
+ async def test_call_tool_on_parent_server(
283
+ self,
284
+ mcp_server: FastMCP,
285
+ nested_mcp_server: FastMCP,
286
+ recording_middleware: RecordingMiddleware,
287
+ nested_middleware: RecordingMiddleware,
288
+ ):
289
+ mcp_server.mount(nested_mcp_server, prefix="nested")
290
+
291
+ async with Client(mcp_server) as client:
292
+ await client.call_tool("add", {"a": 1, "b": 2})
293
+
294
+ assert recording_middleware.assert_called(times=3)
295
+ assert recording_middleware.assert_called(method="tools/call", times=3)
296
+ assert recording_middleware.assert_called(hook="on_message", times=1)
297
+ assert recording_middleware.assert_called(hook="on_request", times=1)
298
+ assert recording_middleware.assert_called(hook="on_call_tool", times=1)
299
+
300
+ assert nested_middleware.assert_called(times=0)
301
+
302
+ async def test_call_tool_on_nested_server(
303
+ self,
304
+ mcp_server: FastMCP,
305
+ nested_mcp_server: FastMCP,
306
+ recording_middleware: RecordingMiddleware,
307
+ nested_middleware: RecordingMiddleware,
308
+ ):
309
+ mcp_server.mount(nested_mcp_server, prefix="nested")
310
+
311
+ async with Client(mcp_server) as client:
312
+ await client.call_tool("nested_add", {"a": 1, "b": 2})
313
+
314
+ assert recording_middleware.assert_called(times=3)
315
+ assert recording_middleware.assert_called(method="tools/call", times=3)
316
+ assert recording_middleware.assert_called(hook="on_message", times=1)
317
+ assert recording_middleware.assert_called(hook="on_request", times=1)
318
+ assert recording_middleware.assert_called(hook="on_call_tool", times=1)
319
+
320
+ assert nested_middleware.assert_called(times=3)
321
+ assert nested_middleware.assert_called(method="tools/call", times=3)
322
+ assert nested_middleware.assert_called(hook="on_message", times=1)
323
+ assert nested_middleware.assert_called(hook="on_request", times=1)
324
+ assert nested_middleware.assert_called(hook="on_call_tool", times=1)
325
+
326
+ async def test_read_resource_on_parent_server(
327
+ self,
328
+ mcp_server: FastMCP,
329
+ nested_mcp_server: FastMCP,
330
+ recording_middleware: RecordingMiddleware,
331
+ nested_middleware: RecordingMiddleware,
332
+ ):
333
+ mcp_server.mount(nested_mcp_server, prefix="nested")
334
+
335
+ async with Client(mcp_server) as client:
336
+ await client.read_resource("resource://test")
337
+
338
+ assert recording_middleware.assert_called(times=3)
339
+ assert recording_middleware.assert_called(method="resources/read", times=3)
340
+ assert recording_middleware.assert_called(hook="on_message", times=1)
341
+ assert recording_middleware.assert_called(hook="on_request", times=1)
342
+ assert recording_middleware.assert_called(hook="on_read_resource", times=1)
343
+
344
+ assert nested_middleware.assert_called(times=0)
345
+
346
+ async def test_read_resource_on_nested_server(
347
+ self,
348
+ mcp_server: FastMCP,
349
+ nested_mcp_server: FastMCP,
350
+ recording_middleware: RecordingMiddleware,
351
+ nested_middleware: RecordingMiddleware,
352
+ ):
353
+ mcp_server.mount(nested_mcp_server, prefix="nested")
354
+
355
+ async with Client(mcp_server) as client:
356
+ await client.read_resource("resource://nested/test")
357
+
358
+ assert recording_middleware.assert_called(times=3)
359
+ assert recording_middleware.assert_called(method="resources/read", times=3)
360
+ assert recording_middleware.assert_called(hook="on_message", times=1)
361
+ assert recording_middleware.assert_called(hook="on_request", times=1)
362
+ assert recording_middleware.assert_called(hook="on_read_resource", times=1)
363
+
364
+ assert nested_middleware.assert_called(times=3)
365
+ assert nested_middleware.assert_called(method="resources/read", times=3)
366
+ assert nested_middleware.assert_called(hook="on_message", times=1)
367
+ assert nested_middleware.assert_called(hook="on_request", times=1)
368
+ assert nested_middleware.assert_called(hook="on_read_resource", times=1)
369
+
370
+ async def test_get_prompt_on_parent_server(
371
+ self,
372
+ mcp_server: FastMCP,
373
+ nested_mcp_server: FastMCP,
374
+ recording_middleware: RecordingMiddleware,
375
+ nested_middleware: RecordingMiddleware,
376
+ ):
377
+ mcp_server.mount(nested_mcp_server, prefix="nested")
378
+
379
+ async with Client(mcp_server) as client:
380
+ await client.get_prompt("test_prompt", {"x": "test"})
381
+
382
+ assert recording_middleware.assert_called(times=3)
383
+ assert recording_middleware.assert_called(method="prompts/get", times=3)
384
+ assert recording_middleware.assert_called(hook="on_message", times=1)
385
+ assert recording_middleware.assert_called(hook="on_request", times=1)
386
+ assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
387
+
388
+ assert nested_middleware.assert_called(times=0)
389
+
390
+ async def test_get_prompt_on_nested_server(
391
+ self,
392
+ mcp_server: FastMCP,
393
+ nested_mcp_server: FastMCP,
394
+ recording_middleware: RecordingMiddleware,
395
+ nested_middleware: RecordingMiddleware,
396
+ ):
397
+ mcp_server.mount(nested_mcp_server, prefix="nested")
398
+
399
+ async with Client(mcp_server) as client:
400
+ await client.get_prompt("nested_test_prompt", {"x": "test"})
401
+
402
+ assert recording_middleware.assert_called(times=3)
403
+ assert recording_middleware.assert_called(method="prompts/get", times=3)
404
+ assert recording_middleware.assert_called(hook="on_message", times=1)
405
+ assert recording_middleware.assert_called(hook="on_request", times=1)
406
+ assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
407
+
408
+ assert nested_middleware.assert_called(times=3)
409
+ assert nested_middleware.assert_called(method="prompts/get", times=3)
410
+ assert nested_middleware.assert_called(hook="on_message", times=1)
411
+ assert nested_middleware.assert_called(hook="on_request", times=1)
412
+ assert nested_middleware.assert_called(hook="on_get_prompt", times=1)
413
+
414
+ async def test_list_tools_on_nested_server(
415
+ self,
416
+ mcp_server: FastMCP,
417
+ nested_mcp_server: FastMCP,
418
+ recording_middleware: RecordingMiddleware,
419
+ nested_middleware: RecordingMiddleware,
420
+ ):
421
+ mcp_server.mount(nested_mcp_server, prefix="nested")
422
+
423
+ async with Client(mcp_server) as client:
424
+ await client.list_tools()
425
+
426
+ assert recording_middleware.assert_called(times=3)
427
+ assert recording_middleware.assert_called(method="tools/list", times=3)
428
+ assert recording_middleware.assert_called(hook="on_message", times=1)
429
+ assert recording_middleware.assert_called(hook="on_request", times=1)
430
+ assert recording_middleware.assert_called(hook="on_list_tools", times=1)
431
+
432
+ assert nested_middleware.assert_called(times=3)
433
+ assert nested_middleware.assert_called(method="tools/list", times=3)
434
+ assert nested_middleware.assert_called(hook="on_message", times=1)
435
+ assert nested_middleware.assert_called(hook="on_request", times=1)
436
+ assert nested_middleware.assert_called(hook="on_list_tools", times=1)
437
+
438
+ async def test_list_resources_on_nested_server(
439
+ self,
440
+ mcp_server: FastMCP,
441
+ nested_mcp_server: FastMCP,
442
+ recording_middleware: RecordingMiddleware,
443
+ nested_middleware: RecordingMiddleware,
444
+ ):
445
+ mcp_server.mount(nested_mcp_server, prefix="nested")
446
+
447
+ async with Client(mcp_server) as client:
448
+ await client.list_resources()
449
+
450
+ assert recording_middleware.assert_called(times=3)
451
+ assert recording_middleware.assert_called(method="resources/list", times=3)
452
+ assert recording_middleware.assert_called(hook="on_message", times=1)
453
+ assert recording_middleware.assert_called(hook="on_request", times=1)
454
+ assert recording_middleware.assert_called(hook="on_list_resources", times=1)
455
+
456
+ assert nested_middleware.assert_called(times=3)
457
+ assert nested_middleware.assert_called(method="resources/list", times=3)
458
+ assert nested_middleware.assert_called(hook="on_message", times=1)
459
+ assert nested_middleware.assert_called(hook="on_request", times=1)
460
+ assert nested_middleware.assert_called(hook="on_list_resources", times=1)
461
+
462
+ async def test_list_resource_templates_on_nested_server(
463
+ self,
464
+ mcp_server: FastMCP,
465
+ nested_mcp_server: FastMCP,
466
+ recording_middleware: RecordingMiddleware,
467
+ nested_middleware: RecordingMiddleware,
468
+ ):
469
+ mcp_server.mount(nested_mcp_server, prefix="nested")
470
+
471
+ async with Client(mcp_server) as client:
472
+ await client.list_resource_templates()
473
+
474
+ assert recording_middleware.assert_called(times=3)
475
+ assert recording_middleware.assert_called(
476
+ method="resources/templates/list", times=3
477
+ )
478
+ assert recording_middleware.assert_called(hook="on_message", times=1)
479
+ assert recording_middleware.assert_called(hook="on_request", times=1)
480
+ assert recording_middleware.assert_called(
481
+ hook="on_list_resource_templates", times=1
482
+ )
483
+
484
+ assert nested_middleware.assert_called(times=3)
485
+ assert nested_middleware.assert_called(
486
+ method="resources/templates/list", times=3
487
+ )
488
+ assert nested_middleware.assert_called(hook="on_message", times=1)
489
+ assert nested_middleware.assert_called(hook="on_request", times=1)
490
+ assert nested_middleware.assert_called(
491
+ hook="on_list_resource_templates", times=1
492
+ )