diff --git a/src/tryton_mcp/config.py b/src/tryton_mcp/config.py index ef1b1e9..8ac52d3 100644 --- a/src/tryton_mcp/config.py +++ b/src/tryton_mcp/config.py @@ -1,5 +1,7 @@ import os +from contextlib import asynccontextmanager from dataclasses import dataclass, field +from typing import Optional from sabatron_tryton_rpc_client.client import Client @@ -19,6 +21,7 @@ class TrytonSettings: port: int = field( default_factory=lambda: int(os.environ.get("TRYTON_PORT", "8000")) ) + _client: Optional[Client] = field(default=None, init=False, repr=False) def to_dict(self) -> dict: return { @@ -30,7 +33,28 @@ class TrytonSettings: } def get_client(self) -> Client: - return Client(**self.to_dict()) + if self._client is None: + self._client = Client(**self.to_dict()) + return self._client + + def connect(self): + client = self.get_client() + client.connect() + + def disconnect(self): + if self._client is not None: + close_method = getattr(self._client, "close", None) or getattr( + self._client, "disconnect", None + ) + if close_method: + close_method() + self._client = None + + @asynccontextmanager + async def lifespan(self): + self.connect() + yield + self.disconnect() def __post_init__(self): if not self.hostname: diff --git a/src/tryton_mcp/server.py b/src/tryton_mcp/server.py index aa9459f..1fc155d 100644 --- a/src/tryton_mcp/server.py +++ b/src/tryton_mcp/server.py @@ -2,15 +2,22 @@ MCP Server for Naliia Module """ -from sabatron_tryton_rpc_client.client import Client +from contextlib import asynccontextmanager from fastmcp import FastMCP -from typing import List +from typing import List, AsyncIterator import datetime from tryton_mcp.config import settings -mcp = FastMCP("Tryton MCP Server") +@asynccontextmanager +async def server_lifespan(server: FastMCP) -> AsyncIterator[None]: + settings.connect() + yield + settings.disconnect() + + +mcp = FastMCP("Tryton MCP Server", lifespan=server_lifespan) @mcp.tool() @@ -28,9 +35,7 @@ def create_schedule( } ] - client = Client(**settings.to_dict()) - - client.connect() + client = settings.get_client() name = "model.naliia.schedule.create" args = [example_schedule, {}] response = client.call(name, args) @@ -40,9 +45,7 @@ def create_schedule( @mcp.tool() def find_service_centers(): - client = Client(**settings.to_dict()) - - client.connect() + client = settings.get_client() name = "model.naliia.service_center.search_read" args = [[[]], 0, None, None, ["name", "address.street"], {}] @@ -53,9 +56,7 @@ def find_service_centers(): @mcp.tool() def find_products_and_services(): - client = Client(**settings.to_dict()) - - client.connect() + client = settings.get_client() name = "model.product.product.search_read" args = [