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