334 lines
13 KiB
Python
334 lines
13 KiB
Python
"""OpenAPI-based tool client for MCP-compatible tool servers.
|
||
|
||
Compatible with open-webui/openapi-servers (https://github.com/open-webui/openapi-servers).
|
||
Fetches the server's OpenAPI spec, converts operations to OpenAI function-calling
|
||
definitions, and executes tool calls by making HTTP requests back to the server.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import logging
|
||
import re
|
||
from typing import Any
|
||
|
||
import httpx
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class ToolClient:
|
||
"""Client for OpenAPI-compatible tool servers.
|
||
|
||
Usage::
|
||
|
||
client = ToolClient(base_url="http://localhost:8000", api_key="…")
|
||
tools = await client.get_tools() # OpenAI tool definitions
|
||
result = await client.call("my_op", {}) # execute a tool call
|
||
"""
|
||
|
||
def __init__(self, base_url: str, api_key: str = "", timeout: float = 30.0) -> None:
|
||
self._base_url = base_url.rstrip("/")
|
||
self._api_key = api_key
|
||
self._timeout = timeout
|
||
self._spec: dict[str, Any] | None = None
|
||
self._server_specs: dict[str, dict[str, Any]] | None = None
|
||
self._tool_routes: dict[str, tuple[str, str]] = {}
|
||
self._tools: list[dict[str, Any]] | None = None
|
||
|
||
@property
|
||
def base_url(self) -> str:
|
||
return self._base_url
|
||
|
||
def _headers(self) -> dict[str, str]:
|
||
headers: dict[str, str] = {"Accept": "application/json"}
|
||
if self._api_key:
|
||
headers["Authorization"] = "Bearer " + self._api_key.strip()
|
||
return headers
|
||
|
||
async def get_spec(self) -> dict[str, Any]:
|
||
"""Fetch (and cache) the OpenAPI spec from ``/openapi.json``."""
|
||
if self._spec is not None:
|
||
return self._spec
|
||
async with httpx.AsyncClient(timeout=self._timeout) as http:
|
||
resp = await http.get(
|
||
f"{self._base_url}/openapi.json",
|
||
headers=self._headers(),
|
||
)
|
||
resp.raise_for_status()
|
||
self._spec = resp.json()
|
||
path_count = len(self._spec.get("paths", {}))
|
||
logger.info("Loaded OpenAPI spec from %s (%d paths)", self._base_url, path_count)
|
||
return self._spec
|
||
|
||
async def get_tools(self) -> list[dict[str, Any]]:
|
||
"""Return OpenAI function-calling tool definitions from the OpenAPI spec.
|
||
|
||
MCPO config-file mode exposes each configured MCP server under its own subpath
|
||
(for example ``/github/openapi.json``), while the root schema has no paths and
|
||
only links to the per-server docs. In that mode, tool names are namespaced as
|
||
``server__operation`` so similarly named tools from different MCP servers do not
|
||
collide.
|
||
"""
|
||
if self._tools is not None:
|
||
return self._tools
|
||
|
||
specs = await self._get_server_specs()
|
||
tools: list[dict[str, Any]] = []
|
||
self._tool_routes = {}
|
||
|
||
for server_name, spec in specs.items():
|
||
server_tools = _spec_to_openai_tools(spec)
|
||
for tool in server_tools:
|
||
function = tool["function"]
|
||
operation_name = function["name"]
|
||
if server_name:
|
||
namespaced_name = f"{server_name}__{operation_name}"
|
||
function["name"] = namespaced_name
|
||
function["description"] = f"[{server_name}] {function.get('description', '')}"
|
||
else:
|
||
namespaced_name = operation_name
|
||
self._tool_routes[namespaced_name] = (server_name, operation_name)
|
||
tools.append(tool)
|
||
|
||
self._tools = tools
|
||
logger.info("Registered %d tools from %s", len(self._tools), self._base_url)
|
||
return self._tools
|
||
|
||
async def call(self, tool_name: str, arguments: dict[str, Any]) -> str:
|
||
"""Execute a named tool call and return the result as a string.
|
||
|
||
Resolves *tool_name* to a path+method in the cached OpenAPI spec,
|
||
partitions *arguments* into path params, query params, and request
|
||
body, then makes the HTTP request.
|
||
"""
|
||
await self.get_tools()
|
||
if tool_name not in self._tool_routes:
|
||
raise ValueError(f"Operation '{tool_name}' not found in spec")
|
||
|
||
server_name, operation_name = self._tool_routes[tool_name]
|
||
specs = await self._get_server_specs()
|
||
spec = specs[server_name]
|
||
path, method, path_params, query_params, body = _resolve_operation(
|
||
spec, operation_name, arguments
|
||
)
|
||
|
||
base_path = f"/{server_name}" if server_name else ""
|
||
url = self._base_url + base_path + path
|
||
for key, value in path_params.items():
|
||
url = url.replace(f"{{{key}}}", str(value))
|
||
|
||
async with httpx.AsyncClient(timeout=self._timeout) as http:
|
||
resp = await http.request(
|
||
method=method.upper(),
|
||
url=url,
|
||
headers=self._headers(),
|
||
params=query_params if query_params else None,
|
||
json=body if body else None,
|
||
)
|
||
resp.raise_for_status()
|
||
content_type = resp.headers.get("content-type", "")
|
||
if "application/json" in content_type:
|
||
try:
|
||
return json.dumps(resp.json(), indent=2)
|
||
except ValueError:
|
||
pass
|
||
return resp.text
|
||
|
||
async def _get_server_specs(self) -> dict[str, dict[str, Any]]:
|
||
"""Return OpenAPI specs keyed by MCPO server name.
|
||
|
||
``""`` denotes the root OpenAPI server. Non-empty keys denote MCPO
|
||
config-file subservers such as ``github`` or ``memory``.
|
||
"""
|
||
if self._server_specs is not None:
|
||
return self._server_specs
|
||
|
||
root_spec = await self.get_spec()
|
||
if root_spec.get("paths"):
|
||
self._server_specs = {"": root_spec}
|
||
return self._server_specs
|
||
|
||
server_names = _discover_mcpo_server_names(root_spec)
|
||
if not server_names:
|
||
self._server_specs = {"": root_spec}
|
||
return self._server_specs
|
||
|
||
specs: dict[str, dict[str, Any]] = {}
|
||
async with httpx.AsyncClient(timeout=self._timeout) as http:
|
||
for server_name in server_names:
|
||
resp = await http.get(
|
||
f"{self._base_url}/{server_name}/openapi.json",
|
||
headers=self._headers(),
|
||
)
|
||
resp.raise_for_status()
|
||
spec = resp.json()
|
||
specs[server_name] = spec
|
||
logger.info(
|
||
"Loaded OpenAPI spec from %s/%s (%d paths)",
|
||
self._base_url,
|
||
server_name,
|
||
len(spec.get("paths", {})),
|
||
)
|
||
|
||
self._server_specs = specs
|
||
return self._server_specs
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# OpenAPI → OpenAI tool-definition helpers
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _discover_mcpo_server_names(spec: dict[str, Any]) -> list[str]:
|
||
"""Extract MCPO config-file server names from the root OpenAPI description."""
|
||
description = str(spec.get("info", {}).get("description", ""))
|
||
names = re.findall(r"\[([^\]]+)]\(/([^/]+)/docs\)", description)
|
||
return [name for name, path_name in names if name == path_name]
|
||
|
||
|
||
def _resolve_ref(schema: dict[str, Any], components: dict[str, Any]) -> dict[str, Any]:
|
||
"""Recursively resolve a ``$ref`` inside an OpenAPI schema."""
|
||
if "$ref" not in schema:
|
||
return schema
|
||
ref: str = schema["$ref"] # e.g. "#/components/schemas/Foo"
|
||
parts = ref.lstrip("#/").split("/")
|
||
# parts = ["components", "schemas", "Foo"]
|
||
obj: Any = {"components": components}
|
||
for part in parts:
|
||
if not isinstance(obj, dict):
|
||
return {}
|
||
obj = obj.get(part, {})
|
||
return _schema_to_json_schema(obj, components)
|
||
|
||
|
||
def _schema_to_json_schema(
|
||
schema: dict[str, Any], components: dict[str, Any] | None = None
|
||
) -> dict[str, Any]:
|
||
"""Convert an OpenAPI schema object to a JSON Schema-compatible dict."""
|
||
if not schema:
|
||
return {}
|
||
resolved = _resolve_ref(schema, components or {}) if "$ref" in schema else schema
|
||
result: dict[str, Any] = {}
|
||
for key in ("type", "description", "enum", "format", "default", "minimum", "maximum"):
|
||
if key in resolved:
|
||
result[key] = resolved[key]
|
||
if "items" in resolved:
|
||
result["items"] = _schema_to_json_schema(resolved["items"], components)
|
||
if "properties" in resolved:
|
||
result["properties"] = {
|
||
k: _schema_to_json_schema(v, components) for k, v in resolved["properties"].items()
|
||
}
|
||
if "required" in resolved:
|
||
result["required"] = resolved["required"]
|
||
return result
|
||
|
||
|
||
def _operation_id(method: str, path: str, op: dict[str, Any]) -> str:
|
||
"""Return an operation ID, generating one from method+path if absent."""
|
||
if op.get("operationId"):
|
||
return str(op["operationId"])
|
||
slug = path.strip("/").replace("/", "_").replace("{", "").replace("}", "")
|
||
return f"{method}_{slug}" if slug else method
|
||
|
||
|
||
def _spec_to_openai_tools(spec: dict[str, Any]) -> list[dict[str, Any]]:
|
||
"""Convert an OpenAPI 3.x spec to a list of OpenAI function-calling tool defs."""
|
||
tools: list[dict[str, Any]] = []
|
||
paths: dict[str, Any] = spec.get("paths", {})
|
||
components: dict[str, Any] = spec.get("components", {})
|
||
|
||
for path, path_item in paths.items():
|
||
for http_method in ("get", "post", "put", "patch", "delete"):
|
||
op: dict[str, Any] | None = path_item.get(http_method)
|
||
if op is None:
|
||
continue
|
||
|
||
op_id = _operation_id(http_method, path, op)
|
||
description: str = op.get("summary") or op.get("description") or ""
|
||
|
||
properties: dict[str, Any] = {}
|
||
required: list[str] = []
|
||
|
||
# URL / query parameters
|
||
for param in op.get("parameters", []):
|
||
name: str = param["name"]
|
||
schema = _schema_to_json_schema(param.get("schema", {"type": "string"}), components)
|
||
if param.get("description"):
|
||
schema["description"] = param["description"]
|
||
properties[name] = schema
|
||
if param.get("required"):
|
||
required.append(name)
|
||
|
||
# Request body – flatten top-level object properties
|
||
rb: dict[str, Any] = op.get("requestBody", {})
|
||
if rb:
|
||
json_content = rb.get("content", {}).get("application/json", {})
|
||
body_schema = _schema_to_json_schema(json_content.get("schema", {}), components)
|
||
if body_schema.get("type") == "object":
|
||
for prop_name, prop_schema in body_schema.get("properties", {}).items():
|
||
properties[prop_name] = prop_schema
|
||
if rb.get("required", False):
|
||
required.extend(body_schema.get("required", []))
|
||
|
||
tool: dict[str, Any] = {
|
||
"type": "function",
|
||
"function": {
|
||
"name": op_id,
|
||
"description": description,
|
||
"parameters": {
|
||
"type": "object",
|
||
"properties": properties,
|
||
},
|
||
},
|
||
}
|
||
if required:
|
||
tool["function"]["parameters"]["required"] = required
|
||
|
||
tools.append(tool)
|
||
|
||
return tools
|
||
|
||
|
||
def _resolve_operation(
|
||
spec: dict[str, Any],
|
||
operation_id: str,
|
||
arguments: dict[str, Any],
|
||
) -> tuple[str, str, dict[str, Any], dict[str, Any], dict[str, Any]]:
|
||
"""Find path+method for *operation_id* and split *arguments* into parts.
|
||
|
||
Returns ``(path, method, path_params, query_params, body)``.
|
||
|
||
* ``path_params`` – values that replace ``{placeholders}`` in the URL path.
|
||
* ``query_params`` – values appended as query string.
|
||
* ``body`` – any remaining arguments sent as the JSON request body.
|
||
"""
|
||
paths: dict[str, Any] = spec.get("paths", {})
|
||
|
||
for path, path_item in paths.items():
|
||
for http_method in ("get", "post", "put", "patch", "delete"):
|
||
op: dict[str, Any] | None = path_item.get(http_method)
|
||
if op is None:
|
||
continue
|
||
if _operation_id(http_method, path, op) != operation_id:
|
||
continue
|
||
|
||
path_params: dict[str, Any] = {}
|
||
query_params: dict[str, Any] = {}
|
||
declared_params: set[str] = set()
|
||
|
||
for param in op.get("parameters", []):
|
||
name = param["name"]
|
||
declared_params.add(name)
|
||
if name not in arguments:
|
||
continue
|
||
if param.get("in") == "path":
|
||
path_params[name] = arguments[name]
|
||
else:
|
||
query_params[name] = arguments[name]
|
||
|
||
body = {k: v for k, v in arguments.items() if k not in declared_params}
|
||
return path, http_method, path_params, query_params, body
|
||
|
||
raise ValueError(f"Operation '{operation_id}' not found in spec")
|