File size: 7,284 Bytes
a70d9c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
"""
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())