mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-10-11 07:23:19 +02:00
214 lines
7 KiB
Python
214 lines
7 KiB
Python
"""Optional capacity and expiry for the default upload store."""
|
|
|
|
import base64
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from typing import Any
|
|
from unittest.mock import Mock
|
|
|
|
import pytest
|
|
|
|
from fastmcp import FastMCP
|
|
from fastmcp.apps import file_upload
|
|
from fastmcp.apps.file_upload import FileUpload
|
|
from fastmcp.server.context import Context
|
|
from fastmcp.server.providers.addressing import hashed_backend_name
|
|
from tests.apps.test_file_upload import _make_file
|
|
|
|
|
|
@pytest.fixture
|
|
def clock(monkeypatch: pytest.MonkeyPatch) -> list[float]:
|
|
now = [100.0]
|
|
monkeypatch.setattr(file_upload, "monotonic", lambda: now[0])
|
|
return now
|
|
|
|
|
|
def scope(name: str) -> Context:
|
|
return Mock(spec=Context, session_id=name)
|
|
|
|
|
|
async def test_total_size_uses_payload_sizes():
|
|
upload = FileUpload(max_total_size=4)
|
|
server = FastMCP("test", providers=[upload])
|
|
first = _make_file("a.txt", "abcd")
|
|
first["size"] = 0
|
|
await server.call_tool(
|
|
hashed_backend_name("Files", "store_files"), {"files": [first]}
|
|
)
|
|
|
|
with pytest.raises(Exception, match="total storage size"):
|
|
await server.call_tool(
|
|
hashed_backend_name("Files", "store_files"),
|
|
{"files": [_make_file("b.txt", "b")]},
|
|
)
|
|
|
|
assert len(upload._store["__default__"]) == 1
|
|
|
|
|
|
def test_total_size_spans_scopes():
|
|
upload = FileUpload(max_total_size=4)
|
|
a, b = scope("a"), scope("b")
|
|
upload.on_store([_make_file("a.txt", "abc")], a)
|
|
|
|
with pytest.raises(ValueError, match="total storage size"):
|
|
upload.on_store([_make_file("b.txt", "ab")], b)
|
|
upload.on_store([_make_file("b.txt", "b")], b)
|
|
|
|
assert [f["name"] for f in upload.on_list(a)] == ["a.txt"]
|
|
assert [f["name"] for f in upload.on_list(b)] == ["b.txt"]
|
|
|
|
|
|
def test_rejected_batch_keeps_existing_contents():
|
|
upload = FileUpload(max_total_size=4)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "a")], ctx)
|
|
|
|
with pytest.raises(ValueError, match="total storage size"):
|
|
upload.on_store([_make_file("a.txt", "abc"), _make_file("b.txt", "bb")], ctx)
|
|
|
|
assert upload.on_read("a.txt", ctx)["content"] == "a"
|
|
assert [f["name"] for f in upload.on_list(ctx)] == ["a.txt"]
|
|
|
|
|
|
def test_replacements_return_capacity():
|
|
upload = FileUpload(max_total_size=4)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abcd")], ctx)
|
|
upload.on_store([_make_file("a.txt", "a")], ctx)
|
|
upload.on_store([_make_file("b.txt", "bcd")], ctx)
|
|
|
|
with pytest.raises(ValueError, match="total storage size"):
|
|
upload.on_store([_make_file("a.txt", "abcd")], ctx)
|
|
assert upload.on_read("a.txt", ctx)["content"] == "a"
|
|
assert upload.on_read("b.txt", ctx)["content"] == "bcd"
|
|
|
|
|
|
def test_repeated_names_use_final_contents():
|
|
upload = FileUpload(max_total_size=1)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abcd"), _make_file("a.txt", "a")], ctx)
|
|
|
|
assert upload.on_read("a.txt", ctx)["content"] == "a"
|
|
|
|
|
|
def test_zero_capacity_accepts_empty_files():
|
|
upload = FileUpload(max_total_size=0)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("empty.txt", "")], ctx)
|
|
|
|
assert upload.on_read("empty.txt", ctx)["content"] == ""
|
|
with pytest.raises(ValueError, match="total storage size"):
|
|
upload.on_store([_make_file("a.txt", "a")], ctx)
|
|
|
|
|
|
def test_default_storage_persists_across_calls(clock: list[float]):
|
|
upload = FileUpload()
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abc")], ctx)
|
|
clock[0] += 10000
|
|
upload.on_store([_make_file("b.txt", "bcd")], ctx)
|
|
|
|
assert [f["name"] for f in upload.on_list(ctx)] == ["a.txt", "b.txt"]
|
|
assert upload.on_read("a.txt", ctx)["content"] == "abc"
|
|
|
|
|
|
def test_expiry_does_not_refresh_on_read(clock: list[float]):
|
|
upload = FileUpload(file_ttl_seconds=5)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abc")], ctx)
|
|
clock[0] = 104.9
|
|
assert upload.on_read("a.txt", ctx)["content"] == "abc"
|
|
clock[0] = 105
|
|
|
|
with pytest.raises(ValueError, match="not found"):
|
|
upload.on_read("a.txt", ctx)
|
|
assert upload.on_list(ctx) == []
|
|
assert upload._store == {}
|
|
|
|
|
|
def test_replacement_starts_new_expiry(clock: list[float]):
|
|
upload = FileUpload(file_ttl_seconds=5)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abc")], ctx)
|
|
clock[0] = 104
|
|
upload.on_store([_make_file("a.txt", "def")], ctx)
|
|
clock[0] = 105
|
|
assert upload.on_read("a.txt", ctx)["content"] == "def"
|
|
clock[0] = 109
|
|
assert upload.on_list(ctx) == []
|
|
|
|
|
|
def test_expiry_returns_global_capacity(clock: list[float]):
|
|
upload = FileUpload(max_total_size=4, file_ttl_seconds=5)
|
|
a, b = scope("a"), scope("b")
|
|
upload.on_store([_make_file("a.txt", "abcd")], a)
|
|
clock[0] = 105
|
|
upload.on_store([_make_file("b.txt", "abcd")], b)
|
|
|
|
assert upload.on_list(a) == []
|
|
assert upload.on_read("b.txt", b)["content"] == "abcd"
|
|
assert set(upload._store) == {"b"}
|
|
|
|
|
|
def test_list_prunes_expired_files_in_other_scopes(clock: list[float]):
|
|
upload = FileUpload(file_ttl_seconds=5)
|
|
a, b = scope("a"), scope("b")
|
|
upload.on_store([_make_file("a.txt", "a")], a)
|
|
clock[0] = 104
|
|
upload.on_store([_make_file("b.txt", "b")], b)
|
|
clock[0] = 105
|
|
|
|
assert [f["name"] for f in upload.on_list(b)] == ["b.txt"]
|
|
assert set(upload._store) == {"b"}
|
|
|
|
|
|
@pytest.mark.parametrize("total_size", [-1, -1024])
|
|
def test_total_size_requires_nonnegative_value(total_size: int):
|
|
with pytest.raises(ValueError, match="max_total_size"):
|
|
FileUpload(max_total_size=total_size)
|
|
|
|
|
|
@pytest.mark.parametrize("ttl", [0, -1, float("inf"), float("-inf"), float("nan")])
|
|
def test_expiry_requires_positive_finite_value(ttl: float):
|
|
with pytest.raises(ValueError, match="file_ttl_seconds"):
|
|
FileUpload(file_ttl_seconds=ttl)
|
|
|
|
|
|
def test_custom_storage_controls_its_own_retention(clock: list[float]):
|
|
stored: dict[str, dict[str, Any]] = {}
|
|
|
|
class CustomUpload(FileUpload):
|
|
def on_store(self, files: list[dict[str, Any]], ctx: Context):
|
|
stored.update({f["name"]: f for f in files})
|
|
return self.on_list(ctx)
|
|
|
|
def on_list(self, ctx: Context):
|
|
return [{"name": name} for name in stored]
|
|
|
|
def on_read(self, name: str, ctx: Context):
|
|
return stored[name]
|
|
|
|
upload = CustomUpload(max_total_size=0, file_ttl_seconds=1)
|
|
ctx = scope("a")
|
|
upload.on_store([_make_file("a.txt", "abc")], ctx)
|
|
clock[0] += 10000
|
|
|
|
assert upload.on_list(ctx) == [{"name": "a.txt"}]
|
|
assert base64.b64decode(upload.on_read("a.txt", ctx)["data"]) == b"abc"
|
|
assert upload._store == {}
|
|
|
|
|
|
def test_concurrent_batches_share_capacity():
|
|
upload = FileUpload(max_total_size=4)
|
|
|
|
def store(name: str) -> bool:
|
|
try:
|
|
upload.on_store([_make_file("file.txt", "abcd")], scope(name))
|
|
except ValueError:
|
|
return False
|
|
return True
|
|
|
|
with ThreadPoolExecutor(max_workers=2) as executor:
|
|
results = list(executor.map(store, ["a", "b"]))
|
|
|
|
assert sorted(results) == [False, True]
|
|
assert sum(len(files) for files in upload._store.values()) == 1
|