Files
steward_mirror/steward/tools/client.py
T

260 lines
9.7 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
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._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."""
if self._tools is not None:
return self._tools
spec = await self.get_spec()
self._tools = _spec_to_openai_tools(spec)
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.
"""
spec = await self.get_spec()
path, method, path_params, query_params, body = _resolve_operation(
spec, tool_name, arguments
)
url = self._base_url + 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
# ---------------------------------------------------------------------------
# OpenAPI → OpenAI tool-definition helpers
# ---------------------------------------------------------------------------
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")