aisolate-client 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- aisolate/__init__.py +0 -0
- aisolate/client/__init__.py +12 -0
- aisolate/client/exception.py +50 -0
- aisolate/client/gateway_client.py +384 -0
- aisolate/client/models.py +319 -0
- aisolate/client/remote_executer.py +56 -0
- aisolate/client/sandbox.py +579 -0
- aisolate/client/sandbox_manager.py +297 -0
- aisolate/client/tests/__init__.py +0 -0
- aisolate/client/tests/integration/__init__.py +0 -0
- aisolate/client/tests/integration/conftest.py +81 -0
- aisolate/client/tests/integration/test_gateway_client.py +359 -0
- aisolate/client/tests/integration/test_sandbox.py +397 -0
- aisolate/common/__init__.py +0 -0
- aisolate/common/codec.py +161 -0
- aisolate/common/errors.py +16 -0
- aisolate/common/models.py +74 -0
- aisolate/common/tests/__init__.py +0 -0
- aisolate/common/tests/fixtures.py +160 -0
- aisolate/common/tests/utils.py +117 -0
- aisolate_client-0.1.0.dist-info/METADATA +37 -0
- aisolate_client-0.1.0.dist-info/RECORD +24 -0
- aisolate_client-0.1.0.dist-info/WHEEL +5 -0
- aisolate_client-0.1.0.dist-info/top_level.txt +1 -0
aisolate/__init__.py
ADDED
|
File without changes
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
from .sandbox import Sandbox
|
|
2
|
+
from .sandbox_manager import SandboxManager
|
|
3
|
+
from .exception import SandboxError
|
|
4
|
+
from .models import Execution, KernelMessage
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"Sandbox",
|
|
8
|
+
"SandboxManager",
|
|
9
|
+
"SandboxError",
|
|
10
|
+
"Execution",
|
|
11
|
+
"KernelMessage",
|
|
12
|
+
]
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
from typing import Any, Optional
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class SandboxClientError(Exception):
|
|
5
|
+
"""Base exception for Sandbox Client errors."""
|
|
6
|
+
|
|
7
|
+
pass
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class SandboxConnectionError(SandboxClientError):
|
|
11
|
+
"""Raised for network connection-related errors."""
|
|
12
|
+
|
|
13
|
+
pass
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class SandboxTimeoutError(SandboxClientError):
|
|
17
|
+
"""Raised when a request to the sandbox API times out."""
|
|
18
|
+
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class SandboxAPIError(SandboxClientError):
|
|
23
|
+
"""
|
|
24
|
+
Raised for errors returned by the Sandbox API (e.g., HTTP 4xx or 5xx).
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
message: str,
|
|
30
|
+
status_code: Optional[int] = None,
|
|
31
|
+
response_data: Optional[Any] = None,
|
|
32
|
+
):
|
|
33
|
+
super().__init__(message)
|
|
34
|
+
self.status_code = status_code
|
|
35
|
+
self.response_data = response_data
|
|
36
|
+
|
|
37
|
+
def __str__(self):
|
|
38
|
+
return f"API Error (Status {self.status_code}): {super().__str__()} - Response: {self.response_data}"
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class SandboxWebSocketError(SandboxClientError):
|
|
42
|
+
"""Raised for errors specific to WebSocket communication."""
|
|
43
|
+
|
|
44
|
+
pass
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class SandboxError(Exception):
|
|
48
|
+
"""Raised for any errors related to the Sandbox lifecycle or execution."""
|
|
49
|
+
|
|
50
|
+
pass
|
|
@@ -0,0 +1,384 @@
|
|
|
1
|
+
import random
|
|
2
|
+
import json
|
|
3
|
+
import asyncio
|
|
4
|
+
import logging
|
|
5
|
+
import uuid
|
|
6
|
+
from asyncio import Queue
|
|
7
|
+
from typing import (
|
|
8
|
+
Any,
|
|
9
|
+
AsyncIterator,
|
|
10
|
+
Callable,
|
|
11
|
+
Dict,
|
|
12
|
+
Optional,
|
|
13
|
+
TypeVar,
|
|
14
|
+
AsyncContextManager,
|
|
15
|
+
)
|
|
16
|
+
from urllib.parse import urljoin
|
|
17
|
+
|
|
18
|
+
import httpx
|
|
19
|
+
from websockets.exceptions import ConnectionClosed
|
|
20
|
+
from websockets.asyncio.client import connect as ws_connect, ClientConnection
|
|
21
|
+
|
|
22
|
+
from aisolate.common.codec import decode_kernel_message
|
|
23
|
+
from aisolate.common.models import KernelMessage, KernelExecutionCompleted
|
|
24
|
+
from .exception import (
|
|
25
|
+
SandboxAPIError,
|
|
26
|
+
SandboxConnectionError,
|
|
27
|
+
SandboxTimeoutError,
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
logger = logging.getLogger(__name__)
|
|
31
|
+
|
|
32
|
+
T = TypeVar("T", bound=Dict[str, Any] | None)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class GatewayClient:
|
|
36
|
+
"""
|
|
37
|
+
Thin wrapper around HTTPX + WebSockets for the Gateway API.
|
|
38
|
+
|
|
39
|
+
This client maintains a single, long-lived WebSocket connection and supports
|
|
40
|
+
two patterns of use:
|
|
41
|
+
1. As an async context manager (recommended for most cases):
|
|
42
|
+
```
|
|
43
|
+
async with GatewayClient(base_url) as client:
|
|
44
|
+
await client.create_sandbox("my-sbx")
|
|
45
|
+
```
|
|
46
|
+
2. As a standalone object with manual lifecycle management:
|
|
47
|
+
```
|
|
48
|
+
client = GatewayClient(base_url)
|
|
49
|
+
try:
|
|
50
|
+
await client.start()
|
|
51
|
+
await client.create_sandbox("my-sbx")
|
|
52
|
+
finally:
|
|
53
|
+
await client.close()
|
|
54
|
+
```
|
|
55
|
+
"""
|
|
56
|
+
|
|
57
|
+
def __init__(
|
|
58
|
+
self,
|
|
59
|
+
base_url: str,
|
|
60
|
+
api_key: Optional[str] = None,
|
|
61
|
+
*,
|
|
62
|
+
timeout: float = 60.0,
|
|
63
|
+
session: Optional[httpx.AsyncClient] = None,
|
|
64
|
+
ws_connect: Optional[
|
|
65
|
+
Callable[..., AsyncContextManager[ClientConnection]]
|
|
66
|
+
] = None,
|
|
67
|
+
ping_interval: Optional[float] = 30.0,
|
|
68
|
+
ping_timeout: Optional[float] = 3600.0,
|
|
69
|
+
) -> None:
|
|
70
|
+
"""Initializes the GatewayClient configuration."""
|
|
71
|
+
self.base_url = base_url.rstrip("/") + "/"
|
|
72
|
+
self.api_key = api_key
|
|
73
|
+
self._timeout = timeout
|
|
74
|
+
self._ping_interval = ping_interval
|
|
75
|
+
self._ping_timeout = ping_timeout
|
|
76
|
+
|
|
77
|
+
# HTTPX session state
|
|
78
|
+
self._session = session
|
|
79
|
+
self._owns_session = session is None
|
|
80
|
+
# Webscocket connection
|
|
81
|
+
self._ws_connect = ws_connect
|
|
82
|
+
|
|
83
|
+
# WebSocket connection state
|
|
84
|
+
self._ws: Optional[ClientConnection] = None
|
|
85
|
+
self._listener_task: Optional[asyncio.Task] = None
|
|
86
|
+
self._response_queues: Dict[str, Queue[KernelMessage]] = {}
|
|
87
|
+
self._is_active: bool = False
|
|
88
|
+
|
|
89
|
+
async def start(self) -> "GatewayClient":
|
|
90
|
+
"""
|
|
91
|
+
Establishes the HTTP session and WebSocket connection.
|
|
92
|
+
This method is idempotent and safe to call multiple times.
|
|
93
|
+
"""
|
|
94
|
+
if self._is_active:
|
|
95
|
+
return self
|
|
96
|
+
|
|
97
|
+
# 1. Set up the HTTPX session for standard API calls
|
|
98
|
+
if self._session is None:
|
|
99
|
+
headers = {"Authorization": f"Token {self.api_key}"} if self.api_key else {}
|
|
100
|
+
self._session = httpx.AsyncClient(
|
|
101
|
+
base_url=self.base_url, timeout=self._timeout, headers=headers
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
# 2. Establish the persistent WebSocket connection
|
|
105
|
+
# The `websockets` library has more robust and configurable support for
|
|
106
|
+
# keepalive pings, which is crucial for maintaining long-lived connections
|
|
107
|
+
# through intermediaries like load balancers. We will use it directly
|
|
108
|
+
# instead of the `httpx` websocket support.
|
|
109
|
+
rel_path = "/gateway/ws"
|
|
110
|
+
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
|
111
|
+
ws_url = self._absolute_ws_url(rel_path)
|
|
112
|
+
ws_connect_func = self._ws_connect or ws_connect
|
|
113
|
+
ws_cm = ws_connect_func(
|
|
114
|
+
ws_url,
|
|
115
|
+
open_timeout=self._timeout,
|
|
116
|
+
additional_headers=headers,
|
|
117
|
+
ping_interval=self._ping_interval,
|
|
118
|
+
ping_timeout=self._ping_timeout,
|
|
119
|
+
)
|
|
120
|
+
|
|
121
|
+
self._ws = await ws_cm.__aenter__()
|
|
122
|
+
|
|
123
|
+
# 3. Start the background task to listen for and route all incoming messages
|
|
124
|
+
self._response_queues = {}
|
|
125
|
+
self._listener_task = asyncio.create_task(self._listen_for_messages())
|
|
126
|
+
|
|
127
|
+
self._is_active = True
|
|
128
|
+
logger.info("GatewayClient started and connected.")
|
|
129
|
+
return self
|
|
130
|
+
|
|
131
|
+
async def close(
|
|
132
|
+
self, exc_type: Any = None, exc: Any = None, tb: Any = None
|
|
133
|
+
) -> None:
|
|
134
|
+
"""
|
|
135
|
+
Closes the WebSocket connection and the HTTP session.
|
|
136
|
+
This method is idempotent and safe to call multiple times.
|
|
137
|
+
"""
|
|
138
|
+
if not self._is_active:
|
|
139
|
+
return
|
|
140
|
+
|
|
141
|
+
# 1. Stop the listener task
|
|
142
|
+
if self._listener_task:
|
|
143
|
+
self._listener_task.cancel()
|
|
144
|
+
try:
|
|
145
|
+
await self._listener_task
|
|
146
|
+
except asyncio.CancelledError:
|
|
147
|
+
pass
|
|
148
|
+
self._listener_task = None
|
|
149
|
+
|
|
150
|
+
# 2. Close the WebSocket connection (and wait for it to fully shut down)
|
|
151
|
+
if self._ws:
|
|
152
|
+
ws = self._ws
|
|
153
|
+
self._ws = None
|
|
154
|
+
await ws.close()
|
|
155
|
+
# if it's a websockets.WebSocketClientProtocol
|
|
156
|
+
if hasattr(ws, "wait_closed"):
|
|
157
|
+
await ws.wait_closed()
|
|
158
|
+
|
|
159
|
+
# 3. Close the HTTPX session if we own it
|
|
160
|
+
if self._owns_session and self._session:
|
|
161
|
+
await self._session.aclose()
|
|
162
|
+
self._session = None
|
|
163
|
+
|
|
164
|
+
self._is_active = False
|
|
165
|
+
logger.info("GatewayClient closed.")
|
|
166
|
+
|
|
167
|
+
async def __aenter__(self) -> "GatewayClient":
|
|
168
|
+
"""Context manager entry point. Starts the client."""
|
|
169
|
+
return await self.start()
|
|
170
|
+
|
|
171
|
+
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None:
|
|
172
|
+
"""Context manager exit point. Closes the client."""
|
|
173
|
+
await self.close(exc_type, exc, tb)
|
|
174
|
+
|
|
175
|
+
async def _listen_for_messages(self):
|
|
176
|
+
"""
|
|
177
|
+
Listens for all messages on the WebSocket and routes them to the
|
|
178
|
+
correct response queue based on their 'request_id'.
|
|
179
|
+
"""
|
|
180
|
+
if not self._ws:
|
|
181
|
+
return
|
|
182
|
+
try:
|
|
183
|
+
async for msg in self._ws:
|
|
184
|
+
try:
|
|
185
|
+
data = self._extract_message_data(msg)
|
|
186
|
+
if data is None:
|
|
187
|
+
continue
|
|
188
|
+
request_id, kernel_msg = decode_kernel_message(data)
|
|
189
|
+
if request_id and request_id in self._response_queues:
|
|
190
|
+
await self._response_queues[request_id].put(kernel_msg)
|
|
191
|
+
else:
|
|
192
|
+
logger.warning(
|
|
193
|
+
f"Received message with unknown request_id: {request_id}"
|
|
194
|
+
)
|
|
195
|
+
except asyncio.CancelledError:
|
|
196
|
+
# we were explicitly cancelled — just exit
|
|
197
|
+
return
|
|
198
|
+
except Exception:
|
|
199
|
+
logger.exception(
|
|
200
|
+
"Failed to decode or route incoming WebSocket message"
|
|
201
|
+
)
|
|
202
|
+
except ConnectionClosed as e:
|
|
203
|
+
logger.info(f"WebSocket connection closed cleanly (code={e.code}).")
|
|
204
|
+
except Exception:
|
|
205
|
+
logger.exception("Listener task terminated due to an unexpected error.")
|
|
206
|
+
|
|
207
|
+
# If the connection drops, signal an error to all pending callers
|
|
208
|
+
# by placing a sentinel value (e.g., an exception) in their queues.
|
|
209
|
+
error = SandboxConnectionError("WebSocket connection lost")
|
|
210
|
+
for queue in self._response_queues.values():
|
|
211
|
+
await queue.put(error)
|
|
212
|
+
|
|
213
|
+
async def execute_code(
|
|
214
|
+
self,
|
|
215
|
+
sandbox_id: str,
|
|
216
|
+
code: str,
|
|
217
|
+
*,
|
|
218
|
+
language: str = "python3",
|
|
219
|
+
env_vars: Optional[Dict[str, str]] = None,
|
|
220
|
+
context_id: Optional[str] = None,
|
|
221
|
+
) -> AsyncIterator[KernelMessage]:
|
|
222
|
+
"""
|
|
223
|
+
Executes code by sending a message over the persistent WebSocket.
|
|
224
|
+
It then listens on a dedicated queue for corresponding messages.
|
|
225
|
+
|
|
226
|
+
Args:
|
|
227
|
+
sandbox_id: The sandbox to execute code in
|
|
228
|
+
code: The code to execute
|
|
229
|
+
language: Programming language (only used if context_id is None)
|
|
230
|
+
env_vars: Environment variables to set
|
|
231
|
+
context_id: Optional context/notebook ID. If None, uses default context for language.
|
|
232
|
+
"""
|
|
233
|
+
if not self._is_active:
|
|
234
|
+
raise SandboxConnectionError(
|
|
235
|
+
"GatewayClient is not running. Use `async with` or call `await client.start()`."
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
request_id = str(uuid.uuid4())
|
|
239
|
+
queue: Queue[KernelMessage | Exception] = asyncio.Queue()
|
|
240
|
+
self._response_queues[request_id] = queue
|
|
241
|
+
|
|
242
|
+
try:
|
|
243
|
+
payload_data: Dict[str, Any] = {"code": code, "env_vars": env_vars}
|
|
244
|
+
|
|
245
|
+
# Use context_id if provided, otherwise use language (creates default context)
|
|
246
|
+
if context_id:
|
|
247
|
+
payload_data["context_id"] = context_id
|
|
248
|
+
else:
|
|
249
|
+
payload_data["language"] = language
|
|
250
|
+
|
|
251
|
+
payload = {
|
|
252
|
+
"request_id": request_id,
|
|
253
|
+
"type": "execute_request",
|
|
254
|
+
"sandbox_id": sandbox_id,
|
|
255
|
+
"payload": payload_data,
|
|
256
|
+
}
|
|
257
|
+
await self._ws_send(payload)
|
|
258
|
+
|
|
259
|
+
while True:
|
|
260
|
+
msg = await queue.get()
|
|
261
|
+
if isinstance(msg, Exception):
|
|
262
|
+
raise msg
|
|
263
|
+
yield msg
|
|
264
|
+
if isinstance(msg, KernelExecutionCompleted):
|
|
265
|
+
break
|
|
266
|
+
finally:
|
|
267
|
+
# Clean up the queue for this specific request
|
|
268
|
+
self._response_queues.pop(request_id, None)
|
|
269
|
+
|
|
270
|
+
async def _request(
|
|
271
|
+
self,
|
|
272
|
+
method: str,
|
|
273
|
+
endpoint: str,
|
|
274
|
+
*,
|
|
275
|
+
json: Any = None,
|
|
276
|
+
retries: int = 3,
|
|
277
|
+
backoff: float = 0.5,
|
|
278
|
+
timeout: float | httpx.Timeout | None = None,
|
|
279
|
+
) -> T:
|
|
280
|
+
"""
|
|
281
|
+
Make an HTTP request using the underlying AsyncClient.
|
|
282
|
+
"""
|
|
283
|
+
if not self._is_active or self._session is None:
|
|
284
|
+
raise RuntimeError("Client not entered – use `async with` context manager")
|
|
285
|
+
|
|
286
|
+
if not hasattr(self._session, method.lower()):
|
|
287
|
+
raise ValueError(f"Unsupported HTTP method: {method}")
|
|
288
|
+
url = urljoin(self.base_url, endpoint)
|
|
289
|
+
http_method: Callable[..., Any] = getattr(self._session, method.lower())
|
|
290
|
+
|
|
291
|
+
attempt = 0
|
|
292
|
+
while True:
|
|
293
|
+
try:
|
|
294
|
+
# Only pass json parameter if it's provided
|
|
295
|
+
if json is not None:
|
|
296
|
+
resp: httpx.Response = await http_method(url, json=json)
|
|
297
|
+
else:
|
|
298
|
+
resp: httpx.Response = await http_method(url)
|
|
299
|
+
resp.raise_for_status()
|
|
300
|
+
return resp.json() if resp.content else None
|
|
301
|
+
except (httpx.TimeoutException, httpx.RequestError) as exc:
|
|
302
|
+
if attempt == retries:
|
|
303
|
+
raise SandboxTimeoutError(f"{method} {url} failed") from exc
|
|
304
|
+
delay = (backoff * 2**attempt) + (random.uniform(0, 1) * backoff)
|
|
305
|
+
attempt += 1
|
|
306
|
+
await asyncio.sleep(delay)
|
|
307
|
+
except httpx.HTTPStatusError as exc:
|
|
308
|
+
raise SandboxAPIError(
|
|
309
|
+
f"{method} {url} returned {exc.response.status_code}"
|
|
310
|
+
) from exc
|
|
311
|
+
|
|
312
|
+
async def _ws_send(self, payload: dict) -> None:
|
|
313
|
+
if getattr(self._ws, "send", None):
|
|
314
|
+
await self._ws.send(json.dumps(payload))
|
|
315
|
+
else:
|
|
316
|
+
await self._ws.send_json(payload)
|
|
317
|
+
|
|
318
|
+
def _extract_message_data(self, msg: Any) -> Optional[bytes]:
|
|
319
|
+
"""Extracts the relevant bytes from a websocket message, handling different client types."""
|
|
320
|
+
if isinstance(msg, dict): # Starlette-style
|
|
321
|
+
if any(
|
|
322
|
+
kw in (msg.get("text") or msg.get("type") or "")
|
|
323
|
+
for kw in ("websocket.close",)
|
|
324
|
+
):
|
|
325
|
+
return None # Sentinel for close
|
|
326
|
+
return msg.get("bytes")
|
|
327
|
+
|
|
328
|
+
# websockets-style
|
|
329
|
+
return msg.encode() if isinstance(msg, str) else msg
|
|
330
|
+
|
|
331
|
+
def _absolute_ws_url(self, path: str) -> str:
|
|
332
|
+
"""Convert a relative path to a ws:// or wss:// URL based on base_url."""
|
|
333
|
+
scheme = "wss" if self.base_url.startswith("https") else "ws"
|
|
334
|
+
root = self.base_url.removeprefix("https://").removeprefix("http://")
|
|
335
|
+
return f"{scheme}://{root.rstrip('/')}{path}"
|
|
336
|
+
|
|
337
|
+
async def create_sandbox(
|
|
338
|
+
self,
|
|
339
|
+
sandbox_id: str,
|
|
340
|
+
) -> Dict[str, Any]:
|
|
341
|
+
payload = {"sandbox_id": sandbox_id}
|
|
342
|
+
return await self._request("POST", "sandboxes", json=payload)
|
|
343
|
+
|
|
344
|
+
async def delete_sandbox(self, sandbox_id: str) -> None:
|
|
345
|
+
await self._request("DELETE", f"sandboxes/{sandbox_id}")
|
|
346
|
+
|
|
347
|
+
# Context management methods
|
|
348
|
+
async def create_context(
|
|
349
|
+
self,
|
|
350
|
+
sandbox_id: str,
|
|
351
|
+
language: str = "python3",
|
|
352
|
+
*,
|
|
353
|
+
cwd: Optional[str] = None,
|
|
354
|
+
env: Optional[Dict[str, str]] = None,
|
|
355
|
+
) -> Dict[str, Any]:
|
|
356
|
+
"""Create a new context (notebook/kernel) in a sandbox."""
|
|
357
|
+
payload = {"language": language}
|
|
358
|
+
if cwd:
|
|
359
|
+
payload["cwd"] = cwd
|
|
360
|
+
if env:
|
|
361
|
+
payload["env"] = env
|
|
362
|
+
return await self._request(
|
|
363
|
+
"POST",
|
|
364
|
+
f"sandboxes/{sandbox_id}/contexts",
|
|
365
|
+
json=payload,
|
|
366
|
+
)
|
|
367
|
+
|
|
368
|
+
async def list_contexts(self, sandbox_id: str) -> list[Dict[str, Any]]:
|
|
369
|
+
"""List all contexts in a sandbox."""
|
|
370
|
+
return await self._request("GET", f"sandboxes/{sandbox_id}/contexts")
|
|
371
|
+
|
|
372
|
+
async def restart_context(self, sandbox_id: str, context_id: str) -> None:
|
|
373
|
+
"""Restart a specific context."""
|
|
374
|
+
await self._request(
|
|
375
|
+
"POST",
|
|
376
|
+
f"sandboxes/{sandbox_id}/contexts/{context_id}/restart",
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
async def delete_context(self, sandbox_id: str, context_id: str) -> None:
|
|
380
|
+
"""Delete a specific context."""
|
|
381
|
+
await self._request(
|
|
382
|
+
"DELETE",
|
|
383
|
+
f"sandboxes/{sandbox_id}/contexts/{context_id}",
|
|
384
|
+
)
|