|
1 | | -import base64 |
2 | | -import json |
3 | 1 | import uuid |
4 | 2 | from datetime import datetime, timezone |
5 | | -from typing import Any, Dict |
| 3 | +from typing import Any, Dict, Union |
6 | 4 |
|
7 | 5 | from autogen_ext.tools.mcp._config import ( |
8 | 6 | McpServerParams, |
|
32 | 30 | # Global session tracking for status endpoint |
33 | 31 | active_sessions: Dict[str, Dict[str, Any]] = {} |
34 | 32 |
|
| 33 | +# Server-side storage for pending MCP session parameters. |
| 34 | +# Params are registered via POST /ws/connect and consumed (popped) when the WebSocket connects. |
| 35 | +# This prevents attackers from injecting arbitrary server_params via the WebSocket query string. |
| 36 | +pending_session_params: Dict[str, Union[StdioServerParams, SseServerParams, StreamableHttpServerParams]] = {} |
| 37 | + |
35 | 38 |
|
36 | 39 | class CreateWebSocketConnectionRequest(BaseModel): |
37 | 40 | server_params: McpServerParams |
@@ -129,35 +132,19 @@ async def create_mcp_session(bridge: MCPWebSocketBridge, server_params: McpServe |
129 | 132 |
|
130 | 133 | @router.websocket("/ws/{session_id}") |
131 | 134 | async def mcp_websocket(websocket: WebSocket, session_id: str): |
132 | | - """Main WebSocket endpoint - now a thin layer""" |
| 135 | + """Main WebSocket endpoint - looks up server params from server-side storage""" |
| 136 | + # Look up pre-registered server params (one-time use) |
| 137 | + server_params = pending_session_params.pop(session_id, None) |
| 138 | + if server_params is None: |
| 139 | + await websocket.close(code=4004, reason="Unknown or expired session") |
| 140 | + return |
| 141 | + |
133 | 142 | await websocket.accept() |
134 | 143 | logger.info(f"MCP WebSocket connection established for session {session_id}") |
135 | 144 |
|
136 | 145 | bridge = None |
137 | 146 |
|
138 | 147 | try: |
139 | | - # Parse server parameters |
140 | | - query_params = dict(websocket.query_params) |
141 | | - server_params_encoded = query_params.get("server_params") |
142 | | - |
143 | | - if not server_params_encoded: |
144 | | - await websocket.close(code=4000, reason="Missing server_params") |
145 | | - return |
146 | | - |
147 | | - decoded_params = base64.b64decode(server_params_encoded).decode("utf-8") |
148 | | - server_params_dict = json.loads(decoded_params) |
149 | | - |
150 | | - # Create appropriate server params object |
151 | | - if server_params_dict.get("type") == "StdioServerParams": |
152 | | - server_params = StdioServerParams(**server_params_dict) |
153 | | - elif server_params_dict.get("type") == "SseServerParams": |
154 | | - server_params = SseServerParams(**server_params_dict) |
155 | | - elif server_params_dict.get("type") == "StreamableHttpServerParams": |
156 | | - server_params = StreamableHttpServerParams(**server_params_dict) |
157 | | - else: |
158 | | - await websocket.close(code=4000, reason="Invalid server parameters") |
159 | | - return |
160 | | - |
161 | 148 | # Create bridge and run MCP session |
162 | 149 | bridge = MCPWebSocketBridge(websocket, session_id) |
163 | 150 | await create_mcp_session(bridge, server_params, session_id) |
@@ -197,18 +184,18 @@ async def mcp_websocket(websocket: WebSocket, session_id: str): |
197 | 184 |
|
198 | 185 | @router.post("/ws/connect") |
199 | 186 | async def create_mcp_websocket_connection(request: CreateWebSocketConnectionRequest): |
200 | | - """Create WebSocket connection URL""" |
| 187 | + """Register server params and return a WebSocket URL with session_id only""" |
201 | 188 | try: |
202 | 189 | session_id = str(uuid.uuid4()) |
203 | 190 |
|
204 | | - server_params_json = json.dumps(serialize_for_json(request.server_params.model_dump())) |
205 | | - server_params_encoded = base64.b64encode(server_params_json.encode("utf-8")).decode("utf-8") |
| 191 | + # Store params server-side — WebSocket handler will pop them on connect |
| 192 | + pending_session_params[session_id] = request.server_params |
206 | 193 |
|
207 | 194 | return { |
208 | 195 | "status": True, |
209 | 196 | "message": "WebSocket connection URL created", |
210 | 197 | "session_id": session_id, |
211 | | - "websocket_url": f"/api/mcp/ws/{session_id}?server_params={server_params_encoded}", |
| 198 | + "websocket_url": f"/api/mcp/ws/{session_id}", |
212 | 199 | "timestamp": datetime.now(timezone.utc).isoformat(), |
213 | 200 | } |
214 | 201 |
|
|
0 commit comments