Safety-Vector-Studio / mcp_client.py
jwest33's picture
app migration
a70d9c5
Raw History Blame Contribute Delete
7.28 kB
"""
MCP Client for Control Vector MCP Server
Provides simplified interface to the MCP server
"""
import asyncio
from typing import Dict, Any, Optional
from latent_control_mcp.server import ControlVectorMCPServer
class MCPClient:
"""
Client for connecting to the Control Vector MCP Server.
Provides:
- Direct access to server methods for deterministic operations
- MCP tool calling interface for LLM-driven operations (only optimize_alpha)
Usage:
client = MCPClient()
await client.connect()
# Direct method calls (deterministic)
model_info = await client.server.get_model_info()
dataset = await client.server.load_safety_dataset({"sample_size": 50})
# MCP tool call (LLM-driven)
result = await client.call_tool("optimize_alpha", {...})
"""
def __init__(
self,
server_name: str = "control-vector-mcp",
local_files_only: Optional[bool] = None
):
"""
Initialize MCP client.
Args:
server_name: Name identifier for the server
local_files_only: If True, only use local models (don't download).
If None, reads from LOCAL_FILES_ONLY env var (defaults to False)
"""
self.server_name = server_name
self.server: Optional[ControlVectorMCPServer] = None
self._initialized = False
# Always download from HuggingFace Hub on Spaces (local_files_only=False)
self.local_files_only = local_files_only if local_files_only is not None else False
async def __aenter__(self):
"""Async context manager entry"""
await self.connect()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
"""Async context manager exit"""
await self.disconnect()
async def connect(self):
"""Establish connection to MCP server"""
if self._initialized:
return
# Instantiate server directly with configuration
self.server = ControlVectorMCPServer(
local_files_only=self.local_files_only
)
self._initialized = True
async def disconnect(self):
"""Close connection to MCP server"""
if not self._initialized:
return
self.server = None
self._initialized = False
def set_progress_callback(self, callback: callable):
"""
Set a callback for progress updates during long-running operations.
Args:
callback: Function to call with progress updates
"""
if self.server:
self.server.set_progress_callback(callback)
async def call_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Dict[str, Any]:
"""
Call an MCP tool.
Args:
tool_name: Name of the tool to call
arguments: Tool arguments as dictionary
Returns:
Tool result as dictionary
Raises:
RuntimeError: If client is not connected
Exception: If tool call fails
"""
# Auto-connect if not connected
if not self._initialized or not self.server:
await self.connect()
try:
# Call tool through server
result = await self.server.call_tool(tool_name, arguments)
return result
except Exception as e:
return {"success": False, "error": str(e)}
async def list_tools(self) -> list:
"""
List available MCP tools from the server.
Note: Only LLM-driven tools are exposed via MCP.
For deterministic operations, call server methods directly.
Returns:
List of MCP tool names with descriptions
"""
return [
{"name": "optimize_alpha", "description": "Find optimal steering strength (LLM-guided search)"},
{"name": "evaluate_current_vector", "description": "Quick diagnostic before deciding next action"},
{"name": "retrain_vector", "description": "Retrain vector with different layer/params (max 3 attempts)"},
{"name": "try_different_layer", "description": "Test vector at a different layer without retraining"},
]
async def list_methods(self) -> list:
"""
List available direct methods on the server.
These are deterministic operations that should be called directly
on the server instance, not via MCP call_tool.
Returns:
List of method names with descriptions
"""
return [
{"name": "get_model_info", "description": "Get model architecture information"},
{"name": "load_safety_dataset", "description": "Load harmful/harmless dataset from HuggingFace"},
{"name": "select_optimal_layer", "description": "Intelligently select optimal layer for safety vectors"},
{"name": "prepare_safety_training_data", "description": "Process safety dataset into training pairs"},
{"name": "train_control_vector", "description": "Train vector from contrastive data"},
{"name": "export_vector", "description": "Save vector and metadata"},
{"name": "apply_steering", "description": "Generate with steering applied"},
{"name": "load_vector", "description": "Load existing vector"},
{"name": "evaluate_steering", "description": "Test vector at single alpha"},
{"name": "define_behavior", "description": "Create behavior specification"},
{"name": "generate_test_prompts", "description": "Generate test prompts for behavior"},
{"name": "generate_contrastive_responses", "description": "Generate pos/neg response pairs"},
{"name": "discover_behaviors", "description": "Auto-discover behavior pairs"},
]
# Convenience function for one-off tool calls
async def call_mcp_tool(tool_name: str, arguments: Dict[str, Any]) -> Dict[str, Any]:
"""
Convenience function for calling a tool without managing client lifecycle.
Args:
tool_name: Name of the tool
arguments: Tool arguments
Returns:
Tool result
"""
async with MCPClient() as client:
return await client.call_tool(tool_name, arguments)
# Example usage
if __name__ == "__main__":
async def test_client():
"""Test MCP client connection and method calls"""
print("Testing MCP Client...")
async with MCPClient("control-vector-mcp") as client:
# List MCP tools (LLM-driven)
print("\nMCP Tools (LLM-driven):")
tools = await client.list_tools()
for tool in tools:
print(f" - {tool['name']}: {tool['description']}")
# List direct methods (deterministic)
print("\nDirect Methods (deterministic):")
methods = await client.list_methods()
for method in methods:
print(f" - {method['name']}: {method['description']}")
# Test getting model info via direct method call
print("\nTesting get_model_info (direct method call)...")
result = await client.server.get_model_info()
print(f"Result: {result}")
asyncio.run(test_client())