File size: 3,983 Bytes
0c29eb8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d4fceae
0c29eb8
 
 
d4fceae
0c29eb8
 
 
 
 
 
 
 
 
 
 
 
 
d4fceae
0c29eb8
 
 
d4fceae
0c29eb8
 
 
 
 
 
 
 
 
d4fceae
0c29eb8
 
 
 
 
d4fceae
0c29eb8
 
 
 
 
 
 
 
 
d4fceae
0c29eb8
 
 
d4fceae
0c29eb8
 
d4fceae
0c29eb8
 
 
 
 
 
 
9e4d18c
 
 
0c29eb8
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
import inspect

import pytest

from fastmcp import Client
from fastmcp.client.transports import PythonStdioTransport, StdioTransport


class TestKeepAlive:
    # https://github.com/jlowin/fastmcp/issues/581

    @pytest.fixture
    def stdio_script(self, tmp_path):
        script = inspect.cleandoc('''
            import os
            from fastmcp import FastMCP

            mcp = FastMCP()

            @mcp.tool()
            def pid() -> int:
                """Gets PID of server"""
                return os.getpid()

            if __name__ == "__main__":
                mcp.run()
            ''')
        script_file = tmp_path / "stdio.py"
        script_file.write_text(script)
        return script_file

    async def test_keep_alive_default_true(self):
        client = Client(transport=StdioTransport(command="python", args=[""]))

        assert client.transport.keep_alive is True

    async def test_keep_alive_set_false(self):
        client = Client(
            transport=StdioTransport(command="python", args=[""], keep_alive=False)
        )
        assert client.transport.keep_alive is False

    async def test_keep_alive_maintains_session_across_multiple_calls(
        self, stdio_script
    ):
        client = Client(transport=PythonStdioTransport(script_path=stdio_script))
        assert client.transport.keep_alive is True

        async with client:
            result1 = await client.call_tool("pid")
            pid1 = int(result1[0].text)  # type: ignore[attr-defined]

        async with client:
            result2 = await client.call_tool("pid")
            pid2 = int(result2[0].text)  # type: ignore[attr-defined]

        assert pid1 == pid2

    async def test_keep_alive_false_starts_new_session_across_multiple_calls(
        self, stdio_script
    ):
        client = Client(
            transport=PythonStdioTransport(script_path=stdio_script, keep_alive=False)
        )
        assert client.transport.keep_alive is False

        async with client:
            result1 = await client.call_tool("pid")
            pid1 = int(result1[0].text)  # type: ignore[attr-defined]

        async with client:
            result2 = await client.call_tool("pid")
            pid2 = int(result2[0].text)  # type: ignore[attr-defined]

        assert pid1 != pid2

    async def test_keep_alive_starts_new_session_if_manually_closed(self, stdio_script):
        client = Client(transport=PythonStdioTransport(script_path=stdio_script))
        assert client.transport.keep_alive is True

        async with client:
            result1 = await client.call_tool("pid")
            pid1 = int(result1[0].text)  # type: ignore[attr-defined]

        await client.close()

        async with client:
            result2 = await client.call_tool("pid")
            pid2 = int(result2[0].text)  # type: ignore[attr-defined]

        assert pid1 != pid2

    async def test_keep_alive_maintains_session_if_reentered(self, stdio_script):
        client = Client(transport=PythonStdioTransport(script_path=stdio_script))
        assert client.transport.keep_alive is True

        async with client:
            result1 = await client.call_tool("pid")
            pid1 = int(result1[0].text)  # type: ignore[attr-defined]

            async with client:
                result2 = await client.call_tool("pid")
                pid2 = int(result2[0].text)  # type: ignore[attr-defined]

            result3 = await client.call_tool("pid")
            pid3 = int(result3[0].text)  # type: ignore[attr-defined]

        assert pid1 == pid2 == pid3

    async def test_close_session_and_try_to_use_client_raises_error(self, stdio_script):
        client = Client(transport=PythonStdioTransport(script_path=stdio_script))
        assert client.transport.keep_alive is True

        async with client:
            await client.close()
            with pytest.raises(RuntimeError, match="Client is not connected"):
                await client.call_tool("pid")