File size: 2,750 Bytes
e9b3659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from contextlib import asynccontextmanager

from mcp.client.streamable_http import streamable_http_client
from mcp.shared._httpx_utils import create_mcp_http_client
from mcp import ClientSession
from config import get_settings


class RemoteMCPClient:
    """
    Async context-manager wrapper around an MCP streamable-HTTP session.

    Usage (all work must happen inside the `async with` block so that every
    anyio cancel scope is entered and exited in the *same* task):

        async with RemoteMCPClient(url, token) as client:
            tools = await client.list_tools()
    """

    def __init__(self, server_url: str, api_key: str):
        self.server_url = server_url
        self.api_key = api_key
        self.session: ClientSession | None = None
        self._exit_stack = None

    # ── async context manager ─────────────────────────────────────────────────

    async def __aenter__(self):
        from contextlib import AsyncExitStack
        self._exit_stack = AsyncExitStack()
        await self._exit_stack.__aenter__()

        headers = {"Authorization": f"Bearer {self.api_key}"}
        http_client = create_mcp_http_client(headers=headers)
        await self._exit_stack.enter_async_context(http_client)

        # streamable_http_client yields (read_stream, write_stream) β€” 2 values
        read_stream, write_stream = await self._exit_stack.enter_async_context(
            streamable_http_client(self.server_url, http_client=http_client)
        )

        self.session = await self._exit_stack.enter_async_context(
            ClientSession(read_stream, write_stream)
        )
        await self.session.initialize()
        return self

    async def __aexit__(self, *exc_info):
        await self._exit_stack.__aexit__(*exc_info)
        self.session = None

    # ── public API ────────────────────────────────────────────────────────────

    async def list_tools(self) -> list[dict]:
        result = await self.session.list_tools()
        return [{"name": t.name, "description": t.description} for t in result.tools]

    async def call_tool(self, tool_name: str, params: dict) -> dict:
        try:
            result = await self.session.call_tool(tool_name, params)
            return {"status": "success", "data": result.content}
        except Exception as e:
            return {"status": "error", "message": str(e)}


def get_github_mcp_client() -> RemoteMCPClient:
    settings = get_settings()
    return RemoteMCPClient(settings.github_mcp_url, settings.github_mcp_token)