mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-24 06:24:18 +02:00
Compare commits
2 commits
main
...
claude/iss
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
58158e600c |
||
|
|
27dea0e3b6 |
2 changed files with 149 additions and 5 deletions
|
|
@ -96,15 +96,16 @@ class FileSystemProvider(LocalProvider):
|
|||
self._warned_files: dict[Path, float] = {}
|
||||
# Lock for serializing reload operations (created lazily)
|
||||
self._reload_lock: asyncio.Lock | None = None
|
||||
# Generation counter to deduplicate concurrent reloads
|
||||
self._reload_generation: int = 0
|
||||
|
||||
# Always load once at init to catch errors early
|
||||
self._load_components()
|
||||
|
||||
def _load_components(self) -> None:
|
||||
"""Discover and register all components from the filesystem."""
|
||||
# Clear existing components if reloading
|
||||
if self._loaded:
|
||||
self._components.clear()
|
||||
if not self._root.exists():
|
||||
logger.warning("FileSystemProvider root does not exist: %s", self._root)
|
||||
|
||||
result = discover_and_import(self._root)
|
||||
|
||||
|
|
@ -126,6 +127,12 @@ class FileSystemProvider(LocalProvider):
|
|||
for fp in successful_files:
|
||||
self._warned_files.pop(fp, None)
|
||||
|
||||
# Build new components dict and swap atomically to avoid races
|
||||
# with concurrent readers (bug fix: previously used clear-and-rebuild
|
||||
# which could expose an empty dict to concurrent requests)
|
||||
new_components: dict[str, FastMCPComponent] = {}
|
||||
self._components = new_components
|
||||
|
||||
for file_path, component in result.components:
|
||||
try:
|
||||
self._register_component(component)
|
||||
|
|
@ -159,6 +166,10 @@ class FileSystemProvider(LocalProvider):
|
|||
|
||||
Uses a lock to serialize concurrent reload operations and runs
|
||||
filesystem I/O off the event loop using asyncio.to_thread.
|
||||
|
||||
A generation counter deduplicates concurrent reload requests:
|
||||
callers that were waiting for the lock check whether the generation
|
||||
changed while they waited; if it did, the reload already happened.
|
||||
"""
|
||||
if not self._reload and self._loaded:
|
||||
return
|
||||
|
|
@ -167,10 +178,16 @@ class FileSystemProvider(LocalProvider):
|
|||
if self._reload_lock is None:
|
||||
self._reload_lock = asyncio.Lock()
|
||||
|
||||
generation_before = self._reload_generation
|
||||
|
||||
async with self._reload_lock:
|
||||
# Double-check after acquiring lock
|
||||
if self._reload or not self._loaded:
|
||||
# Double-check after acquiring lock: skip if another caller
|
||||
# already reloaded while we were waiting
|
||||
if not self._loaded or (
|
||||
self._reload and self._reload_generation == generation_before
|
||||
):
|
||||
await asyncio.to_thread(self._load_components)
|
||||
self._reload_generation += 1
|
||||
|
||||
# Override provider methods to support reload mode
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
"""Tests for FileSystemProvider."""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp.client import Client
|
||||
|
|
@ -278,6 +280,131 @@ def my_tool() -> str:
|
|||
assert "WARNING" in captured.err and "Failed to import" in captured.err
|
||||
|
||||
|
||||
class TestFileSystemProviderAtomicReload:
|
||||
"""Tests for atomic component swap during reload (bug #4)."""
|
||||
|
||||
async def test_reload_swaps_components_atomically(self, tmp_path: Path):
|
||||
"""Components dict should never be empty during a reload."""
|
||||
(tmp_path / "tool.py").write_text(
|
||||
"""\
|
||||
from fastmcp.tools import tool
|
||||
|
||||
@tool
|
||||
def my_tool() -> str:
|
||||
return "hello"
|
||||
"""
|
||||
)
|
||||
|
||||
provider = FileSystemProvider(tmp_path, reload=True)
|
||||
assert len(provider._components) == 1
|
||||
|
||||
# Reload should produce a new dict, not clear the existing one
|
||||
original_dict = provider._components
|
||||
await provider._ensure_loaded()
|
||||
# After reload, _components is a different dict object
|
||||
assert provider._components is not original_dict
|
||||
assert len(provider._components) == 1
|
||||
|
||||
async def test_reload_does_not_expose_empty_components(self, tmp_path: Path):
|
||||
"""Concurrent readers should never see an empty components dict."""
|
||||
(tmp_path / "tool.py").write_text(
|
||||
"""\
|
||||
from fastmcp.tools import tool
|
||||
|
||||
@tool
|
||||
def my_tool() -> str:
|
||||
return "hello"
|
||||
"""
|
||||
)
|
||||
|
||||
provider = FileSystemProvider(tmp_path, reload=True)
|
||||
observed_empty = False
|
||||
|
||||
original_load = provider._load_components
|
||||
|
||||
def patched_load(self: FileSystemProvider) -> None:
|
||||
nonlocal observed_empty
|
||||
original_load()
|
||||
if len(self._components) == 0:
|
||||
observed_empty = True
|
||||
|
||||
with patch.object(type(provider), "_load_components", patched_load):
|
||||
await provider._ensure_loaded()
|
||||
assert not observed_empty
|
||||
|
||||
|
||||
class TestFileSystemProviderReloadDedup:
|
||||
"""Tests for reload deduplication via generation counter (bug #5)."""
|
||||
|
||||
async def test_generation_counter_increments(self, tmp_path: Path):
|
||||
"""Generation counter should increment after each reload."""
|
||||
(tmp_path / "tool.py").write_text(
|
||||
"""\
|
||||
from fastmcp.tools import tool
|
||||
|
||||
@tool
|
||||
def my_tool() -> str:
|
||||
return "hello"
|
||||
"""
|
||||
)
|
||||
|
||||
provider = FileSystemProvider(tmp_path, reload=True)
|
||||
assert provider._reload_generation == 0
|
||||
|
||||
await provider._ensure_loaded()
|
||||
assert provider._reload_generation == 1
|
||||
|
||||
await provider._ensure_loaded()
|
||||
assert provider._reload_generation == 2
|
||||
|
||||
async def test_concurrent_reloads_deduplicated(self, tmp_path: Path):
|
||||
"""Concurrent _ensure_loaded calls should not each trigger a reload."""
|
||||
(tmp_path / "tool.py").write_text(
|
||||
"""\
|
||||
from fastmcp.tools import tool
|
||||
|
||||
@tool
|
||||
def my_tool() -> str:
|
||||
return "hello"
|
||||
"""
|
||||
)
|
||||
|
||||
provider = FileSystemProvider(tmp_path, reload=True)
|
||||
|
||||
load_count = 0
|
||||
original_load = FileSystemProvider._load_components
|
||||
|
||||
def counting_load(self: FileSystemProvider) -> None:
|
||||
nonlocal load_count
|
||||
load_count += 1
|
||||
original_load(self)
|
||||
|
||||
with patch.object(type(provider), "_load_components", counting_load):
|
||||
# Launch multiple concurrent _ensure_loaded calls
|
||||
await asyncio.gather(
|
||||
provider._ensure_loaded(),
|
||||
provider._ensure_loaded(),
|
||||
provider._ensure_loaded(),
|
||||
)
|
||||
|
||||
# Only one reload should have happened (they all read generation 0
|
||||
# before any acquired the lock, and the first one increments it)
|
||||
assert load_count == 1
|
||||
|
||||
|
||||
class TestFileSystemProviderNonExistentRoot:
|
||||
"""Tests for warning when root directory does not exist."""
|
||||
|
||||
def test_warning_on_nonexistent_root(self, tmp_path: Path, capsys):
|
||||
"""A non-existent root should log a warning."""
|
||||
nonexistent = tmp_path / "does_not_exist"
|
||||
provider = FileSystemProvider(nonexistent)
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "WARNING" in captured.err and "root does not" in captured.err
|
||||
assert len(provider._components) == 0
|
||||
|
||||
|
||||
class TestFileSystemProviderIntegration:
|
||||
"""Integration tests with FastMCP server."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue