Repository navigation
Expand file tree
/
Copy pathsessions_python_plugin.py
More file actions
467 lines (391 loc) · 19.5 KB
/
Copy pathsessions_python_plugin.py
File metadata and controls
467 lines (391 loc) · 19.5 KB
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
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
# Copyright (c) Microsoft. All rights reserved.
import inspect
import logging
import os
import re
from collections.abc import Awaitable, Callable
from io import BytesIO
from typing import Annotated, Any
from azure.core.credentials import TokenCredential
from httpx import AsyncClient, HTTPStatusError
from pydantic import ValidationError
from semantic_kernel.const import USER_AGENT
from semantic_kernel.core_plugins.sessions_python_tool.sessions_python_settings import (
ACASessionsSettings,
SessionsPythonSettings,
)
from semantic_kernel.core_plugins.sessions_python_tool.sessions_remote_file_metadata import SessionsRemoteFileMetadata
from semantic_kernel.exceptions.function_exceptions import FunctionExecutionException, FunctionInitializationError
from semantic_kernel.functions.kernel_function_decorator import kernel_function
from semantic_kernel.kernel_pydantic import HttpsUrl, KernelBaseModel
from semantic_kernel.utils.telemetry.user_agent import HTTP_USER_AGENT, version_info
logger = logging.getLogger(__name__)
SESSIONS_USER_AGENT = f"{HTTP_USER_AGENT}/{version_info} (Language=Python)"
SESSIONS_API_VERSION = "2024-02-02-preview"
class SessionsPythonTool(KernelBaseModel):
"""A plugin for running Python code in an Azure Container Apps dynamic sessions code interpreter."""
pool_management_endpoint: HttpsUrl
settings: SessionsPythonSettings
auth_callback: Callable[..., Any | Awaitable[Any]]
http_client: AsyncClient
enable_dangerous_file_uploads: bool = False
"""Flag to enable file upload operations. Must be True along with allowed_upload_directories to enable uploads."""
allowed_upload_directories: set[str] | None = None
"""Allowed local directories for file uploads. If None, upload_file is disabled (deny-by-default)."""
allowed_download_directories: set[str] | None = None
"""Allowed local directories for file downloads. If None, all paths are allowed (permissive-by-default)."""
def __init__(
self,
auth_callback: Callable[..., Any | Awaitable[Any]] | None = None,
pool_management_endpoint: str | None = None,
settings: SessionsPythonSettings | None = None,
http_client: AsyncClient | None = None,
env_file_path: str | None = None,
token_endpoint: str | None = None,
credential: TokenCredential | None = None,
enable_dangerous_file_uploads: bool = False,
allowed_upload_directories: set[str] | list[str] | None = None,
allowed_download_directories: set[str] | list[str] | None = None,
**kwargs,
):
"""Initializes a new instance of the SessionsPythonTool class.
Args:
auth_callback: Callback to retrieve authentication token.
pool_management_endpoint: The ACA pool management endpoint URL.
settings: Python session settings.
http_client: HTTP client for making requests.
env_file_path: Path to .env file.
token_endpoint: Token endpoint for authentication.
credential: Azure credential for authentication.
enable_dangerous_file_uploads: Flag to enable file upload operations.
Must be True along with allowed_upload_directories to enable file uploads.
Default is False (file uploads disabled).
allowed_upload_directories: Set or list of allowed directories for file uploads.
If None, upload_file will be disabled (deny-by-default).
Empty set/list means no directories are allowed (all uploads denied).
allowed_download_directories: Set or list of allowed directories for file downloads.
If None, all paths are allowed (permissive-by-default).
If configured, downloads are restricted to these directories.
kwargs: Additional keyword arguments.
"""
try:
aca_settings = ACASessionsSettings(
env_file_path=env_file_path,
pool_management_endpoint=pool_management_endpoint,
token_endpoint=token_endpoint,
)
except ValidationError as e:
logger.error(f"Failed to load the ACASessionsSettings with message: {e!s}")
raise FunctionInitializationError(f"Failed to load the ACASessionsSettings with message: {e!s}") from e
if not settings:
settings = SessionsPythonSettings()
if not http_client:
http_client = AsyncClient(timeout=5)
if auth_callback is None:
auth_callback = self._default_auth_callback(aca_settings, credential)
# Convert lists to sets and filter out empty strings (which resolve to CWD)
upload_dirs = {d for d in allowed_upload_directories if d} if allowed_upload_directories is not None else None
download_dirs = (
{d for d in allowed_download_directories if d} if allowed_download_directories is not None else None
)
super().__init__(
pool_management_endpoint=aca_settings.pool_management_endpoint,
settings=settings,
auth_callback=auth_callback,
http_client=http_client,
enable_dangerous_file_uploads=enable_dangerous_file_uploads,
allowed_upload_directories=upload_dirs,
allowed_download_directories=download_dirs,
**kwargs,
)
# region Helper Methods
def _default_auth_callback(
self, aca_settings: ACASessionsSettings, credential: TokenCredential | None
) -> Callable[..., Any | Awaitable[Any]]:
"""Generates a default authentication callback using the ACA settings."""
token = aca_settings.get_sessions_auth_token(credential=credential)
if token is None:
raise FunctionInitializationError("Failed to retrieve the client auth token.")
def auth_callback() -> str:
"""Retrieve the client auth token."""
return token
return auth_callback
async def _ensure_auth_token(self) -> str:
"""Ensure the auth token is valid and handle both sync and async callbacks."""
try:
if inspect.iscoroutinefunction(self.auth_callback):
auth_token = await self.auth_callback()
else:
auth_token = self.auth_callback()
except Exception as e:
logger.error(f"Failed to retrieve the client auth token with message: {e!s}")
raise FunctionExecutionException(f"Failed to retrieve the client auth token with message: {e!s}") from e
return auth_token
def _sanitize_input(self, code: str) -> str:
"""Sanitize input to the python REPL.
Remove whitespace, backtick & python (if llm mistakes python console as terminal).
Args:
code (str): The query to sanitize
Returns:
str: The sanitized query
"""
# Removes `, whitespace & python from start
code = re.sub(r"^(\s|`)*(?i:python)?\s*", "", code)
# Removes whitespace & ` from end
return re.sub(r"(\s|`)*$", "", code)
def _construct_remote_file_path(self, remote_file_path: str) -> str:
"""Construct the remote file path.
Args:
remote_file_path (str): The remote file path.
Returns:
str: The remote file path.
"""
if not remote_file_path.startswith("/mnt/data/"):
remote_file_path = f"/mnt/data/{remote_file_path}"
return remote_file_path
def _build_url_with_version(self, base_url, endpoint, params):
"""Builds a URL with the provided base URL, endpoint, and query parameters."""
params["api-version"] = SESSIONS_API_VERSION
query_string = "&".join([f"{key}={value}" for key, value in params.items()])
if not base_url.endswith("/"):
base_url += "/"
if endpoint.endswith("/"):
endpoint = endpoint[:-1]
return f"{base_url}{endpoint}?{query_string}"
def _validate_local_path_for_upload(self, local_file_path: str) -> str:
"""Validate local path is within allowed upload directories.
Args:
local_file_path: The path to validate.
Returns:
str: The canonicalized absolute path.
Raises:
FunctionExecutionException: If file operations are disabled or path is not within allowed directories.
"""
if not self.enable_dangerous_file_uploads:
raise FunctionExecutionException(
"File upload is disabled. Set 'enable_dangerous_file_uploads=True' "
"and configure 'allowed_upload_directories' to enable."
)
if self.allowed_upload_directories is None:
raise FunctionExecutionException("File upload requires 'allowed_upload_directories' to be configured.")
canonical_path = os.path.realpath(local_file_path)
for allowed_dir in self.allowed_upload_directories:
allowed_canonical = os.path.realpath(allowed_dir)
try:
common = os.path.commonpath([allowed_canonical, canonical_path])
if common == allowed_canonical:
return canonical_path
except ValueError:
continue # Different drives on Windows
logger.warning(f"Upload denied for path: {local_file_path} (resolved: {canonical_path})")
raise FunctionExecutionException(
f"Access denied: '{local_file_path}' is not within allowed upload directories."
)
def _validate_local_path_for_download(self, local_file_path: str) -> str:
"""Validate local path is within allowed download directories (optional protection).
Args:
local_file_path: The path to validate.
Returns:
str: The canonicalized absolute path.
Raises:
FunctionExecutionException: If allowed_download_directories is set and path is not within.
"""
# Permissive by default - if no restrictions configured, allow all paths
if self.allowed_download_directories is None:
return os.path.realpath(local_file_path)
parent_dir = os.path.dirname(local_file_path) or "."
canonical_parent = os.path.realpath(parent_dir)
filename = os.path.basename(local_file_path)
canonical_path = os.path.join(canonical_parent, filename)
for allowed_dir in self.allowed_download_directories:
allowed_canonical = os.path.realpath(allowed_dir)
try:
common = os.path.commonpath([allowed_canonical, canonical_parent])
if common == allowed_canonical:
return canonical_path
except ValueError:
continue
logger.warning(f"Download denied for path: {local_file_path}")
raise FunctionExecutionException(
f"Access denied: '{local_file_path}' is not within allowed download directories."
)
# endregion
# region Kernel Functions
@kernel_function(
description="""Executes the provided Python code.
Start and end the code snippet with double quotes to define it as a string.
Insert \\n within the string wherever a new line should appear.
Add spaces directly after \\n sequences to replicate indentation.
Use \" to include double quotes within the code without ending the string.
Keep everything in a single line; the \\n sequences will represent line breaks
when the string is processed or displayed.
""",
name="execute_code",
)
async def execute_code(self, code: Annotated[str, "The valid Python code to execute"]) -> str:
"""Executes the provided Python code.
Args:
code (str): The valid Python code to execute
Returns:
str: The result of the Python code execution in the form of Result, Stdout, and Stderr
Raises:
FunctionExecutionException: If the provided code is empty.
"""
if not code:
raise FunctionExecutionException("The provided code is empty")
if self.settings.sanitize_input:
code = self._sanitize_input(code)
auth_token = await self._ensure_auth_token()
logger.info(f"Executing Python code: {code}")
self.http_client.headers.update({
"Authorization": f"Bearer {auth_token}",
"Content-Type": "application/json",
USER_AGENT: SESSIONS_USER_AGENT,
})
self.settings.python_code = code
request_body = {
"properties": self.settings.model_dump(exclude_none=True, exclude={"sanitize_input"}, by_alias=True),
}
url = self._build_url_with_version(
base_url=str(self.pool_management_endpoint),
endpoint="code/execute/",
params={"identifier": self.settings.session_id},
)
try:
response = await self.http_client.post(
url=url,
json=request_body,
)
response.raise_for_status()
result = response.json()["properties"]
return (
f"Status:\n{result['status']}\n"
f"Result:\n{result['result']}\n"
f"Stdout:\n{result['stdout']}\n"
f"Stderr:\n{result['stderr']}"
)
except HTTPStatusError as e:
error_message = e.response.text if e.response.text else e.response.reason_phrase
raise FunctionExecutionException(
f"Code execution failed with status code {e.response.status_code} and error: {error_message}"
) from e
@kernel_function(name="upload_file", description="Uploads a file for the current Session ID")
async def upload_file(
self,
*,
local_file_path: Annotated[str, "The path to the local file on the machine"],
remote_file_path: Annotated[
str | None, "The remote path to the file in the session. Defaults to /mnt/data"
] = None,
) -> Annotated[SessionsRemoteFileMetadata, "The metadata of the uploaded file"]:
"""Upload a file to the session pool.
Args:
remote_file_path (str): The path to the file in the session.
local_file_path (str): The path to the file on the local machine.
Must be within allowed_upload_directories.
Returns:
RemoteFileMetadata: The metadata of the uploaded file.
Raises:
FunctionExecutionException: If local_file_path is not provided or not in allowed directories.
"""
if not local_file_path:
raise FunctionExecutionException("Please provide a local file path to upload.")
# Validate path is in allowed directories (deny-by-default)
validated_path = self._validate_local_path_for_upload(local_file_path)
remote_file_path = self._construct_remote_file_path(remote_file_path or os.path.basename(validated_path))
auth_token = await self._ensure_auth_token()
self.http_client.headers.update({
"Authorization": f"Bearer {auth_token}",
USER_AGENT: SESSIONS_USER_AGENT,
})
url = self._build_url_with_version(
base_url=str(self.pool_management_endpoint),
endpoint="files/upload",
params={"identifier": self.settings.session_id},
)
try:
with open(validated_path, "rb") as data:
files = {"file": (remote_file_path, data, "application/octet-stream")}
response = await self.http_client.post(url=url, files=files)
response.raise_for_status()
uploaded_files = await self.list_files()
return next(
file_metadata for file_metadata in uploaded_files if file_metadata.full_path == remote_file_path
)
except HTTPStatusError as e:
error_message = e.response.text if e.response.text else e.response.reason_phrase
raise FunctionExecutionException(
f"Upload failed with status code {e.response.status_code} and error: {error_message}"
) from e
@kernel_function(name="list_files", description="Lists all files in the provided Session ID")
async def list_files(self) -> list[SessionsRemoteFileMetadata]:
"""List the files in the session pool.
Returns:
list[SessionsRemoteFileMetadata]: The metadata for the files in the session pool
"""
auth_token = await self._ensure_auth_token()
self.http_client.headers.update({
"Authorization": f"Bearer {auth_token}",
USER_AGENT: SESSIONS_USER_AGENT,
})
url = self._build_url_with_version(
base_url=str(self.pool_management_endpoint),
endpoint="files",
params={"identifier": self.settings.session_id},
)
try:
response = await self.http_client.get(
url=url,
)
response.raise_for_status()
response_json = response.json()
return [SessionsRemoteFileMetadata.from_dict(entry["properties"]) for entry in response_json["value"]]
except HTTPStatusError as e:
error_message = e.response.text if e.response.text else e.response.reason_phrase
raise FunctionExecutionException(
f"List files failed with status code {e.response.status_code} and error: {error_message}"
) from e
async def download_file(
self,
*,
remote_file_name: Annotated[str, "The name of the file to download, relative to /mnt/data"],
local_file_path: Annotated[str | None, "The local file path to save the file to, optional"] = None,
) -> Annotated[BytesIO | None, "The data of the downloaded file"]:
"""Download a file from the session pool.
Args:
remote_file_name: The name of the file to download, relative to `/mnt/data`.
local_file_path: The path to save the downloaded file to. Should include the extension.
If not provided, the file is returned as a BytesIO object.
Returns:
BytesIO | None: The file content as BytesIO if no local_file_path provided, otherwise None.
Raises:
FunctionExecutionException: If local_file_path is not in allowed directories.
"""
auth_token = await self._ensure_auth_token()
self.http_client.headers.update({
"Authorization": f"Bearer {auth_token}",
USER_AGENT: SESSIONS_USER_AGENT,
})
url = self._build_url_with_version(
base_url=str(self.pool_management_endpoint),
endpoint=f"files/content/{remote_file_name}",
params={"identifier": self.settings.session_id},
)
try:
response = await self.http_client.get(
url=url,
)
response.raise_for_status()
if local_file_path:
# Validate path is in allowed directories (optional, permissive by default)
validated_path = self._validate_local_path_for_download(local_file_path)
with open(validated_path, "wb") as f:
f.write(response.content)
return None
return BytesIO(response.content)
except HTTPStatusError as e:
error_message = e.response.text if e.response.text else e.response.reason_phrase
raise FunctionExecutionException(
f"Download failed with status code {e.response.status_code} and error: {error_message}"
) from e
# endregion