Source code for aidputils.agents.tools.mcp.mcp_service
import logging
from typing import Optional, Dict, Any
from langchain_core.tools import StructuredTool
from aidputils.agents.tools.mcp import pagination_helper
from aidputils.agents.tools.mcp.mcp_client_manager import mcp_client_manager
from aidputils.agents.auth.util.agent_util import _run_in_event_loop
from pydantic import create_model, Field
from aidputils.agents.auth.util import auth_utils
logger = logging.getLogger(__name__)
[docs]
def get_mcp_client(server_name: str, server_url: str, auth: Optional[dict] = None, transport: str = "streamable_http", custom_headers: Optional[Dict[str, str]] = None):
"""
Build and return an MCPHTTPClient for the given server_url with optional auth handling.
Auth handling:
- NO_AUTH (default): no Authorization header is set
- BEARER_TOKEN: requires 'token' in auth dict; adds Authorization: Bearer <token>
Unsupported auth types will raise ValueError.
"""
if not server_name or not isinstance(server_name, str):
raise ValueError("server_name must be a non-empty string")
if not server_url or not isinstance(server_url, str):
raise ValueError("server_url must be a non-empty string")
headers: Dict[str, str] = {
"Accept": "application/json",
}
# Merge in any custom headers provided by caller (runtime overrides)
if isinstance(custom_headers, dict):
for k, v in custom_headers.items():
try:
if isinstance(k, str) and isinstance(v, str) and k.strip():
headers[k] = v
except Exception:
logger.warn(f"Malformed header {k} and value {v} detected")
continue
logger.debug("Constructing MCP HTTP client: url=%s transport=%s authType=%s headers=%s",
server_url, transport, (auth or {}).get("authType", "None"), list(headers.keys()))
# Register or retrieve the client from the central manager
mcp_client_manager.register_sync(
server_name=server_name,
url=server_url,
headers=headers,
auth=auth,
transport=transport,
)
return mcp_client_manager.get_mcp_client(server_name)
[docs]
async def list_objects(client, runtime_params: Dict[str, Any]):
"""
Generic dispatcher to list or fetch MCP objects based on runtime_params.operation.
Supported operations:
- list_tools: paginated listing of tools
- list_resources: paginated listing of resources
- get_prompt: fetch a prompt by name/id with optional variables
"""
op = runtime_params.get("operation")
if not op:
raise ValueError("runtime_params.operation is required")
server_name = runtime_params.get("server_name")
if not server_name or not isinstance(server_name, str):
raise ValueError("runtime_params.server_name (server name) is required")
if op == "list_tools":
return await client.get_tools(server_name=server_name)
if op == "list_resources":
uris = runtime_params.get("uris")
return await client.get_resources(server_name=server_name, uris=uris)
if op == "get_prompt":
prompt = runtime_params.get("prompt_name")
if not prompt or not isinstance(prompt, str):
raise ValueError("runtime_params prompt_name must be provided for get_prompt")
result = await client.get_prompt(server_name, prompt)
try:
return result.model_dump(mode="python")
except Exception:
return result
raise ValueError(f"Unsupported operation '{op}' for list_objects")
[docs]
async def list_objects_paginated(client, runtime_params: Dict[str, Any]):
"""
Wrapper around list_objects that normalizes pagination responses
for list_tools and list_resources. Non-paginated ops (e.g., get_prompt)
pass through the underlying result.
"""
op = runtime_params.get("operation")
page = int(runtime_params.get("page", 1) or 1)
limit = int(runtime_params.get("limit", 50) or 50)
raw = await list_objects(client, runtime_params)
if op in ("list_tools", "list_resources"):
return pagination_helper.normalize_paged_result(op, raw, page, limit)
return raw
[docs]
async def list_all_objects(client, runtime_params: Dict[str, Any]):
return await list_objects(client, runtime_params)
[docs]
def list_objects_paginated_sync(client, runtime_params: Dict[str, Any]):
"""
Synchronous wrapper for list_objects_paginated for use in sync code paths.
"""
return _run_in_event_loop(list_objects_paginated(client, runtime_params))
[docs]
def list_filtered_tools(client, runtime_params: Dict[str, Any]):
# Local helpers to avoid cross-module imports/cycles
def _py_type_from_jsonschema_type(type_name: Optional[str]):
if not isinstance(type_name, str):
return Any
t = type_name.lower()
return {
"string": str,
"integer": int,
"number": float,
"boolean": bool,
"array": list,
"object": dict,
"null": type(None),
}.get(t, Any)
def _build_args_schema_from_json_schema_with_overrides(
schema: Optional[dict],
overrides: Optional[dict],
model_name: str = "MCPToolInput",
):
if not isinstance(schema, dict):
return create_model(model_name)
props = schema.get("properties") or {}
required = set(schema.get("required") or [])
fields = {}
ov = overrides if isinstance(overrides, dict) else {}
for name, spec in props.items():
if not isinstance(spec, dict):
continue
py_type = _py_type_from_jsonschema_type(spec.get("type"))
desc = spec.get("description", "")
schema_default = spec.get("default", None)
has_override = name in ov
override_val = ov.get(name, None)
if name in required and has_override:
# Relax required and set provided override as default
fields[name] = (py_type, Field(override_val, description=desc))
elif name in required:
fields[name] = (py_type, Field(..., description=desc))
else:
default_val = override_val if has_override else (schema_default if schema_default is not None else None)
fields[name] = (py_type, Field(default_val, description=desc))
return create_model(model_name, **fields)
def _build_invoker_factory_local(client, server_name: Optional[str], allowed: list[dict]):
# Build per-name defaults from allowed entries
per_name_defaults: Dict[str, Dict[str, Any]] = {}
for it in (allowed or []):
try:
tool_obj = (it or {}).get("tool") or {}
nm = tool_obj.get("name")
if not nm:
continue
defs = (it or {}).get("argOverrides") or {}
if isinstance(defs, dict):
prev = per_name_defaults.get(nm, {})
prev.update(defs)
per_name_defaults[nm] = prev
except Exception:
continue
def _merge_defaults(tool_name: str, provided: dict) -> dict:
base: dict = {}
if tool_name in per_name_defaults:
base.update(per_name_defaults[tool_name])
base.update(provided or {})
return base
def invoker_factory(tool_name: str):
async def _invoke(**arguments: Dict[str, Any]):
# Best-effort session context for caching
try:
ctx = auth_utils.get_auth_context()
session_id = getattr(ctx, "session_id", None)
except Exception:
session_id = None
merged = _merge_defaults(tool_name, arguments)
use_cached = session_id is not None
return await client.call_tool(server_name, tool_name, merged, 10, use_cached=use_cached, session_id=session_id)
return _invoke
return invoker_factory
# New path: If allowed_tools is self-contained (AllowedToolDetails with embedded 'tool'),
# construct StructuredTools directly without fetching from the server.
allowed_tools = runtime_params.get("allowed_tools")
server_name = runtime_params.get("server_name")
if isinstance(allowed_tools, list) and any(isinstance(it, dict) and isinstance(it.get("tool"), dict) for it in allowed_tools):
invoker_factory = _build_invoker_factory_local(client, server_name, allowed_tools)
result: list[StructuredTool] = []
for it in allowed_tools:
try:
tool_obj = it.get("tool") or {}
if not isinstance(tool_obj, dict):
continue
name = tool_obj.get("name")
if not name:
continue
description = it.get("instruction") or tool_obj.get("description") or ""
input_schema = tool_obj.get("inputSchema") or {}
overrides = it.get("argOverrides") or {}
args_schema = _build_args_schema_from_json_schema_with_overrides(
input_schema, overrides, model_name=f"{name}_Input"
)
# Build coroutine via invoker factory and enforce presence
coroutine = None
if callable(invoker_factory):
try:
coroutine = invoker_factory(name)
except Exception:
coroutine = None
if coroutine is None:
raise ValueError(f"Unable to construct invoker for MCP tool '{name}' (server={server_name})")
st = StructuredTool.from_function(
coroutine=coroutine,
name=name,
description=description,
args_schema=args_schema,
infer_schema=False,
)
result.append(st)
except Exception:
logger.exception("Failed to build StructuredTool from AllowedToolDetails entry")
continue
return result
[docs]
def list_filtered_tools_sync(client, runtime_params: Dict[str, Any]):
tools = _run_in_event_loop(list_all_objects(client, runtime_params))
"""
Filter and transform a list of StructuredTool by applying 'allowed_tools' rules.
allowed_tools may be:
- list[str]: only include tools with those names
- dict[str, dict]: mapping of tool name -> {instruction, argOverrides}
- list[dict]: each dict may include 'name'/'toolName' and optional 'instruction', 'argOverrides'.
If a dict has no name, it is treated as a global override applied to all matched tools.
For matched tools:
- Override description with 'instruction' when provided
- Override argument defaults with 'argOverrides' when provided (best-effort across JSON schema or Pydantic schemas)
Returns:
list[StructuredTool]: Tools after filtering and applying overrides (new StructuredTool instances when possible).
"""
# Build allowed name set and per-name overrides
allowed_names = set()
per_name = {}
allowed_tools = runtime_params.get("allowed_tools")
if allowed_tools is None:
candidates = tools
else:
tmp_names = set()
for item in allowed_tools:
name = item.get("name")
instr = item.get("instruction")
defs = item.get("argOverrides") or item.get("arg_overrides") or {}
tmp_names.add(name)
per_name[name] = {"instruction": instr, "argOverrides": defs}
if tmp_names:
allowed_names = tmp_names
# Filter by allowed names if provided
candidates = []
for t in tools:
name = getattr(t, "name", None) if not isinstance(t, dict) else t.get("name")
if not allowed_names or (name in allowed_names):
candidates.append(t)
def _set_description(tool, text):
if not text:
return
try:
setattr(tool, "description", text)
except Exception:
if isinstance(tool, dict):
tool["description"] = text
def _get_args_schema(tool):
if isinstance(tool, dict):
# Prefer normalized args_schema, but fall back to MCP's inputSchema
return tool.get("args_schema") or tool.get("inputSchema")
# Handle StructuredTool or other objects
schema = getattr(tool, "args_schema", None)
if schema is None:
# Support mcp.types.Tool which uses 'inputSchema'
schema = getattr(tool, "inputSchema", None)
return schema
def _set_args_schema(tool, schema):
# Try to set on StructuredTool-like objects
try:
setattr(tool, "args_schema", schema)
except Exception:
pass
# Dict fallback
if isinstance(tool, dict):
tool["args_schema"] = schema
return
# Support mcp.types.Tool which uses 'inputSchema'
if hasattr(tool, "inputSchema"):
try:
setattr(tool, "inputSchema", schema)
except Exception:
pass
def _apply_defaults_to_schema(schema, defaults):
if not defaults or schema is None:
return schema
# JSON-schema style
if isinstance(schema, dict):
props = schema.get("properties") or {}
if isinstance(props, dict):
# Compare defaults provided in allowed_tools with inputSchema properties
required_list = list(schema.get("required") or [])
for k, v in defaults.items():
if k in props and isinstance(props[k], dict):
# If the schema has a property that is provided via defaults,
# remove the property definition and make it non-required.
try:
del props[k]
except Exception:
pass
if k in required_list:
try:
required_list.remove(k)
except Exception:
pass
# Write back pruned properties/required
schema["properties"] = props
if required_list:
schema["required"] = required_list
elif "required" in schema:
# If empty, drop required entirely
try:
del schema["required"]
except Exception:
pass
return schema
# Pydantic v2 class
fields = getattr(schema, "model_fields", None)
if isinstance(fields, dict):
for k, v in defaults.items():
if k in fields:
try:
field = fields[k]
field.default = v
field.default_factory = None
except Exception:
pass
return schema
# Pydantic v1 class
fields_v1 = getattr(schema, "__fields__", None)
if isinstance(fields_v1, dict):
for k, v in defaults.items():
if k in fields_v1:
try:
fld = fields_v1[k]
fld.default = v
fld.default_factory = None
except Exception:
pass
return schema
return schema
result = []
for t in candidates:
name = getattr(t, "name", None) if not isinstance(t, dict) else t.get("name")
overrides = per_name.get(name, {})
instr = overrides.get("instruction") if overrides else None
defs = overrides.get("argOverrides") if overrides else None
# Attempt to create a shallow copy StructuredTool with same behavior
tool_obj = t
try:
if isinstance(t, StructuredTool):
tool_obj = StructuredTool(
name=t.name,
description=t.description,
args_schema=getattr(t, "args_schema", None),
func=getattr(t, "func", None),
coroutine=getattr(t, "coroutine", None),
return_direct=getattr(t, "return_direct", False)
)
except Exception:
tool_obj = t
if instr:
_set_description(tool_obj, instr)
if isinstance(defs, dict):
schema = _get_args_schema(tool_obj)
new_schema = _apply_defaults_to_schema(schema, defs)
if new_schema is not None:
_set_args_schema(tool_obj, new_schema)
result.append(tool_obj if isinstance(tool_obj, StructuredTool) else t)
return result