334 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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")