from __future__ import annotations import base64 import hashlib import io import json import re import shlex import ssl import tarfile import threading import time import uuid from dataclasses import dataclass from html import unescape from pathlib import Path, PurePosixPath from typing import Any from urllib.parse import quote, urlparse, urlunparse try: # 依赖在部署环境安装;这里延迟报错,方便无网络环境做静态检查。 import requests except ImportError: # pragma: no cover requests = None # type: ignore[assignment] try: import websocket except ImportError: # pragma: no cover websocket = None # type: ignore[assignment] DEFAULT_JUPYTER_WORKSPACE_ROOT = '/root/zk_agent_workspaces' REMOTE_PIP_INDEX_URL = 'https://pypi.tuna.tsinghua.edu.cn/simple' REMOTE_PIP_TRUSTED_HOST = 'pypi.tuna.tsinghua.edu.cn' class JupyterRuntimeError(RuntimeError): """Raised when a remote Jupyter runtime cannot complete a request.""" @dataclass(frozen=True) class JupyterWorkspaceBinding: account_id: str session_id: str base_url: str workspace_root: str workspace_cwd: str workspace_api_path: str jupyter_tree_url: str created_at: float def to_dict(self) -> dict[str, Any]: return { 'account_id': self.account_id, 'session_id': self.session_id, 'base_url': self.base_url, 'workspace_root': self.workspace_root, 'workspace_cwd': self.workspace_cwd, 'workspace_api_path': self.workspace_api_path, 'jupyter_tree_url': self.jupyter_tree_url, 'created_at': self.created_at, } @dataclass(frozen=True) class RemoteCommandResult: exit_code: int stdout: str stderr: str = '' timed_out: bool = False class JupyterRuntimeSession: """A small terminal-backed execution bridge for one account/session pair.""" def __init__( self, *, binding: JupyterWorkspaceBinding, http_session: Any, xsrf_token: str, ) -> None: self.binding = binding self._session = http_session self._xsrf_token = xsrf_token self._terminal_name: str | None = None self._lock = threading.RLock() account_part = sanitize_remote_path_part(binding.account_id) self.account_runtime_root = ( f'{binding.workspace_root}/.runtime/accounts/{account_part}' ) self.python_env_path = f'{self.account_runtime_root}/python/.venv' self.platform_root = '' self.skills_root = '' self.bundle_digest = '' self.synced_at: float | None = None @classmethod def from_persisted(cls, payload: dict[str, Any]) -> 'JupyterRuntimeSession': """从持久化的 Jupyter cookie 和绑定信息恢复 runtime。""" if requests is None: raise JupyterRuntimeError( 'Missing dependency: requests. Please install project dependencies.' ) binding_payload = payload.get('binding') if not isinstance(binding_payload, dict): raise JupyterRuntimeError('Invalid persisted Jupyter binding.') binding = JupyterWorkspaceBinding( account_id=str(binding_payload['account_id']), session_id=str(binding_payload['session_id']), base_url=normalize_jupyter_base_url(str(binding_payload['base_url'])), workspace_root=normalize_posix_path(str(binding_payload['workspace_root'])), workspace_cwd=normalize_posix_path(str(binding_payload['workspace_cwd'])), workspace_api_path=str(binding_payload['workspace_api_path']), jupyter_tree_url=str(binding_payload['jupyter_tree_url']), created_at=float(binding_payload.get('created_at') or time.time()), ) http_session = requests.Session() cookies = payload.get('cookies') if isinstance(cookies, list): for cookie in cookies: if not isinstance(cookie, dict): continue name = cookie.get('name') value = cookie.get('value') if not isinstance(name, str) or not isinstance(value, str): continue http_session.cookies.set( name, value, domain=( cookie.get('domain') if isinstance(cookie.get('domain'), str) else None ), path=( cookie.get('path') if isinstance(cookie.get('path'), str) else '/' ), ) xsrf_token = str( payload.get('xsrf_token') or http_session.cookies.get('_xsrf') or '' ) if not xsrf_token: raise JupyterRuntimeError('Invalid persisted Jupyter _xsrf token.') runtime = cls( binding=binding, http_session=http_session, xsrf_token=xsrf_token, ) if isinstance(payload.get('platform_root'), str): runtime.platform_root = str(payload['platform_root']) if isinstance(payload.get('skills_root'), str): runtime.skills_root = str(payload['skills_root']) if isinstance(payload.get('bundle_digest'), str): runtime.bundle_digest = str(payload['bundle_digest']) if isinstance(payload.get('synced_at'), (int, float)): runtime.synced_at = float(payload['synced_at']) return runtime @classmethod def connect( cls, *, account_id: str, session_id: str, base_url: str, password: str, workspace_root: str = DEFAULT_JUPYTER_WORKSPACE_ROOT, timeout_seconds: float = 20.0, project_root: Path | None = None, ) -> 'JupyterRuntimeSession': if requests is None: raise JupyterRuntimeError( 'Missing dependency: requests. Please install project dependencies.' ) normalized_base_url = normalize_jupyter_base_url(base_url) http_session = requests.Session() xsrf_token = _login_jupyter( http_session, normalized_base_url, password, timeout_seconds=timeout_seconds, ) workspace_root = normalize_posix_path(workspace_root) workspace_cwd = f'{workspace_root}/{sanitize_remote_path_part(session_id)}' workspace_api_path = api_path_for_absolute_path(workspace_cwd) binding = JupyterWorkspaceBinding( account_id=account_id, session_id=session_id, base_url=normalized_base_url, workspace_root=workspace_root, workspace_cwd=workspace_cwd, workspace_api_path=workspace_api_path, jupyter_tree_url=( f'{normalized_base_url}/lab/tree/{quote(workspace_api_path)}' ), created_at=time.time(), ) runtime = cls( binding=binding, http_session=http_session, xsrf_token=xsrf_token, ) runtime.bootstrap_workspace( timeout_seconds=timeout_seconds, project_root=project_root, ) return runtime def to_persisted_dict(self) -> dict[str, Any]: """导出可持久化的连接信息,用于页面刷新或服务重启后恢复。""" return { 'binding': self.binding.to_dict(), 'xsrf_token': self._xsrf_token, 'cookies': [ { 'name': cookie.name, 'value': cookie.value, 'domain': cookie.domain, 'path': cookie.path, 'secure': bool(cookie.secure), 'expires': cookie.expires, } for cookie in self._session.cookies ], 'platform_root': self.platform_root, 'skills_root': self.skills_root, 'bundle_digest': self.bundle_digest, 'synced_at': self.synced_at, 'persisted_at': time.time(), } @property def chat_workspace_root(self) -> str: """Per-user-per-chat root on shared NFS so the SFT training pod (which doesn't mount the jupyter pod's private workspace) can read and write the same artifacts. Layout: /mnt/wangsenhao/autoresearch-zk-users///""" account_part = sanitize_remote_path_part(self.binding.account_id) session_part = sanitize_remote_path_part(self.binding.session_id) return ( f'/mnt/wangsenhao/autoresearch-zk-users/' f'{account_part}/{session_part}' ) def bootstrap_workspace( self, *, timeout_seconds: float = 20.0, project_root: Path | None = None, ) -> None: chat_root = self.chat_workspace_root nfs_users_root = '/mnt/wangsenhao/autoresearch-zk-users' probe = self.run_command( ( f'mkdir -p {shlex.quote(nfs_users_root)} && ' f'touch {shlex.quote(nfs_users_root)}/.write_probe && ' f'rm -f {shlex.quote(nfs_users_root)}/.write_probe' ), timeout_seconds=timeout_seconds, max_output_chars=2000, ) if probe.exit_code != 0: raise JupyterRuntimeError( '远端 jupyter pod 未挂载 /mnt/wangsenhao 或不可写:' + (probe.stdout.strip() or probe.stderr.strip() or 'unknown error') ) ws_output = f'{shlex.quote(self.binding.workspace_cwd)}/output' chat_output_link = f'{shlex.quote(chat_root)}/output' result = self.run_command( ( 'mkdir -p ' f'{ws_output} ' f'{shlex.quote(self.binding.workspace_cwd)}/input ' f'{shlex.quote(self.binding.workspace_cwd)}/scratchpad ' f'{shlex.quote(self.binding.workspace_root)}/.runtime/uploads ' f'{shlex.quote(self.binding.workspace_root)}/runtime_uploads ' f'{shlex.quote(self.account_runtime_root)}/python ' f'{shlex.quote(chat_root)}/results ' f'{shlex.quote(chat_root)}/sft_output ' f'{shlex.quote(chat_root)}/scripts && ' f'if [ -d {chat_output_link} ] && [ ! -L {chat_output_link} ]; then ' f'cp -a {chat_output_link}/* {ws_output}/ 2>/dev/null; ' f'rm -rf {chat_output_link}; ' f'fi && ' f'ln -sfnT {ws_output} {chat_output_link}' ), timeout_seconds=timeout_seconds, max_output_chars=4000, ) if result.exit_code != 0: raise JupyterRuntimeError( result.stdout.strip() or 'Unable to bootstrap Jupyter workspace.' ) self.ensure_python_environment(timeout_seconds=max(timeout_seconds, 300.0)) if project_root is not None: self.sync_runtime_bundle(project_root, timeout_seconds=max(timeout_seconds, 180.0)) self._sync_chat_workspace_scripts(timeout_seconds=timeout_seconds) def _sync_chat_workspace_scripts(self, *, timeout_seconds: float = 20.0) -> None: if not self.skills_root: return chat_root = self.chat_workspace_root # SKILL/program §5 让 agent 直接 `bash scripts/` 调下面这些脚本, # 不同步过去 → agent 当 "脚本不存在" 处理后会肉手写 yaml,常把 # imageCommand 里的 `/scripts/prepare_and_train_sft.py` 写丢前缀, # 导致训练 pod 报 "No such file or directory"。每次 bind 全量覆盖 # 一份,跟 prepare_and_train_sft.py 同样 "刷最新版"。 script_names = ( 'prepare_and_train_sft.py', 'submit_sft.sh', 'sft_train_job.yaml.tpl', 'submit_cml_eval.sh', 'resolve_run_ids.sh', ) src_dir = f'{self.skills_root}/model-iteration/scripts' dst_dir = f'{chat_root}/scripts' copy_cmds = [] for name in script_names: src = f'{src_dir}/{name}' dst = f'{dst_dir}/{name}' copy_cmds.append( f'if [ -f {shlex.quote(src)} ]; then ' f'cp {shlex.quote(src)} {shlex.quote(dst)} && ' f'chmod +x {shlex.quote(dst)} 2>/dev/null || true; ' f'fi' ) result = self.run_command( ' && '.join(copy_cmds), timeout_seconds=timeout_seconds, max_output_chars=2000, ) if result.exit_code != 0: raise JupyterRuntimeError( result.stdout.strip() or 'Unable to sync chat workspace SFT script.' ) def to_dict(self) -> dict[str, Any]: payload = self.binding.to_dict() payload['connected'] = True payload['account_runtime_root'] = self.account_runtime_root payload['python_env_path'] = self.python_env_path payload['python_interpreter'] = self.python_interpreter payload['platform_root'] = self.platform_root payload['skills_root'] = self.skills_root payload['bundle_digest'] = self.bundle_digest payload['synced_at'] = self.synced_at return payload def render_context(self) -> str: lines = [ '[远端 Jupyter 工作区]', '当前 session 已切换到远端 Jupyter 运行时。', f'- 默认工作目录:{self.binding.workspace_cwd}', f'- 输出文件优先写入:{self.binding.workspace_cwd}/output', f'- Python 环境:{self.python_env_path}', '- Python 环境默认只初始化 venv 和 pip 清华源;缺包时使用 python_package 按需安装。', '- 线上挂载目录(例如 /mnt/...)可以通过绝对路径读取和分析。', ] if self.platform_root: lines.extend( [ f'- 平台运行包:{self.platform_root}', f'- Skill 运行包:{self.skills_root}', '- 读取 skills/...、src/...、标签定义/... 时会使用远端同步后的运行包。', ] ) lines.extend( [ '- 除非用户明确要求,不要修改平台代码目录;数据产物应留在当前 session 工作区。', f'- Jupyter 工作区入口:{self.binding.jupyter_tree_url}', ] ) return '\n'.join(lines) @property def python_interpreter(self) -> str: return f'{self.python_env_path}/bin/python' def ensure_python_environment(self, *, timeout_seconds: float = 300.0) -> None: python_env = shlex.quote(self.python_env_path) pip_conf = shlex.quote(f'{self.python_env_path}/pip.conf') marker = shlex.quote(f'{self.python_env_path}/.zk-agent-python-env-v2') command = f''' set -e if [ ! -x {python_env}/bin/python ]; then python3 -m venv {python_env} fi cat > {pip_conf} <<'EOF' [global] index-url = {REMOTE_PIP_INDEX_URL} trusted-host = {REMOTE_PIP_TRUSTED_HOST} disable-pip-version-check = true EOF touch {marker} {python_env}/bin/python -m pip --version {python_env}/bin/python - <<'PY' import sys print(sys.executable) PY ''' result = self.run_command( command, timeout_seconds=timeout_seconds, max_output_chars=12000, ) if result.exit_code != 0: raise JupyterRuntimeError( result.stdout.strip() or 'Unable to initialize remote Python env.' ) def install_python_packages( self, packages: list[str], *, timeout_seconds: float, max_output_chars: int, cancel_event: Any | None = None, ) -> RemoteCommandResult: package_args = ' '.join(shlex.quote(package) for package in packages) return self.run_command( f'{shlex.quote(self.python_interpreter)} -m pip install {package_args}', timeout_seconds=timeout_seconds, max_output_chars=max_output_chars, cancel_event=cancel_event, ) def sync_runtime_bundle( self, project_root: Path, *, timeout_seconds: float = 180.0, ) -> dict[str, Any]: bundle_bytes, digest = build_runtime_bundle(project_root) platform_root = f'{self.binding.workspace_root}/.runtime/platform/{digest[:16]}' upload_path = f'{self.binding.workspace_root}/runtime_uploads/platform-{digest}.tar.gz' digest_file = f'{platform_root}/.bundle_digest' probe = self.run_command( ( f'test -f {shlex.quote(digest_file)} ' f'&& cat {shlex.quote(digest_file)} || true' ), timeout_seconds=20.0, max_output_chars=2000, ) if digest not in probe.stdout: self.put_file_bytes( api_path_for_absolute_path(upload_path), bundle_bytes, ) tmp_root = f'{platform_root}.tmp-{uuid.uuid4().hex[:8]}' result = self.run_command( ( f'rm -rf {shlex.quote(tmp_root)} && ' f'mkdir -p {shlex.quote(tmp_root)} && ' f'tar -xzf {shlex.quote(upload_path)} -C {shlex.quote(tmp_root)} && ' f'printf %s {shlex.quote(digest)} > {shlex.quote(tmp_root)}/.bundle_digest && ' f'rm -rf {shlex.quote(platform_root)} && ' f'mv {shlex.quote(tmp_root)} {shlex.quote(platform_root)}' ), timeout_seconds=timeout_seconds, max_output_chars=12000, ) if result.exit_code != 0: raise JupyterRuntimeError( result.stdout.strip() or 'Unable to extract remote runtime bundle.' ) self.platform_root = platform_root self.skills_root = f'{platform_root}/skills' self.bundle_digest = digest self.synced_at = time.time() self.link_platform_resources() return { 'bundle_digest': digest, 'platform_root': self.platform_root, 'skills_root': self.skills_root, 'uploaded_bytes': len(bundle_bytes), } def link_platform_resources(self) -> None: if not self.platform_root: return workspace = shlex.quote(self.binding.workspace_cwd) platform = shlex.quote(self.platform_root) command = ( f'ln -sfn {platform}/skills {workspace}/skills && ' f'ln -sfn {platform}/src {workspace}/src && ' f'ln -sfn {platform}/标签定义 {workspace}/标签定义 && ' f'ln -sfn {platform}/pyproject.toml {workspace}/pyproject.toml' ) result = self.run_command( command, timeout_seconds=20.0, max_output_chars=4000, ) if result.exit_code != 0: raise JupyterRuntimeError( result.stdout.strip() or 'Unable to link runtime resources into workspace.' ) def put_file_bytes(self, api_path: str, data: bytes) -> None: encoded = base64.b64encode(data).decode('ascii') url = f'{self.binding.base_url}/api/contents/{quote(api_path, safe="/")}' response = self._session.put( url, headers={'X-XSRFToken': self._xsrf_token}, json={ 'type': 'file', 'format': 'base64', 'content': encoded, }, timeout=60, ) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to upload file to Jupyter contents API: HTTP {response.status_code} {response.text[:500]}' ) def run_command( self, command: str, *, timeout_seconds: float, max_output_chars: int, cancel_event: Any | None = None, ) -> RemoteCommandResult: if websocket is None: raise JupyterRuntimeError( 'Missing dependency: websocket-client. Please install project dependencies.' ) if not command.strip(): raise JupyterRuntimeError('Remote command is empty.') marker = f'__ZK_AGENT_EXIT_{uuid.uuid4().hex}__' workspace = shlex.quote(self.binding.workspace_cwd) env_prefix = self._remote_env_prefix() inner_command = f'{env_prefix}{command}' wrapped = ( f'mkdir -p {workspace} && ' f'cd {workspace} && ' f'bash -lc {shlex.quote(inner_command)}; ' f'printf "\\n{marker}:%s\\n" "$?"' ) with self._lock: name = self._ensure_terminal() ws_url = self._terminal_websocket_url(name) ws = websocket.create_connection( ws_url, timeout=min(timeout_seconds, 30.0), cookie=self._cookie_header(), origin=self.binding.base_url, sslopt={'cert_reqs': ssl.CERT_NONE}, header=[ f'X-XSRFToken: {self._xsrf_token}', f'Referer: {self.binding.base_url}/lab', ], ) try: self._drain_initial_messages(ws, max_wait_seconds=1.0) ws.send(json.dumps(['stdin', 'stty -echo\n'])) self._drain_initial_messages(ws, max_wait_seconds=0.3) ws.send(json.dumps(['stdin', f'{wrapped}\n'])) return self._collect_until_marker( ws, marker=marker, timeout_seconds=timeout_seconds, max_output_chars=max_output_chars, cancel_event=cancel_event, ) finally: try: ws.close() except Exception: pass def run_python( self, *, code: str | None, script_path: str | None, args: list[str], stdin: str | None, timeout_seconds: float, max_output_chars: int, cancel_event: Any | None = None, ) -> RemoteCommandResult: if bool(code) == bool(script_path): raise JupyterRuntimeError('code and script_path must specify exactly one value') temp_dir = f'{self.binding.workspace_cwd}/scratchpad/.python_exec' command_parts = [f'mkdir -p {shlex.quote(temp_dir)}'] if code is not None: script_path = f'{temp_dir}/snippet_{uuid.uuid4().hex}.py' command_parts.append( self._remote_write_text_command(script_path, code) ) assert script_path is not None quoted_args = ' '.join(shlex.quote(arg) for arg in args) python_command = f'{shlex.quote(self.python_interpreter)} {shlex.quote(script_path)}' if quoted_args: python_command = f'{python_command} {quoted_args}' if stdin is not None: stdin_path = f'{temp_dir}/stdin_{uuid.uuid4().hex}.txt' command_parts.append( self._remote_write_text_command(stdin_path, stdin) ) python_command = f'{python_command} < {shlex.quote(stdin_path)}' command_parts.append(python_command) return self.run_command( ' && '.join(command_parts), timeout_seconds=timeout_seconds, max_output_chars=max_output_chars, cancel_event=cancel_event, ) def list_dir(self, path: str, *, max_entries: int, max_output_chars: int) -> str: target = self.resolve_workspace_path(path) code = r''' import json from pathlib import Path target = Path(PATH) if not target.exists(): raise SystemExit(f"Path not found: {target}") if not target.is_dir(): raise SystemExit(f"Path is not a directory: {target}") entries = [] for item in sorted(target.iterdir(), key=lambda p: (not p.is_dir(), p.name.lower()))[:MAX_ENTRIES]: entries.append(("dir" if item.is_dir() else "file") + "\t" + str(item)) print("\n".join(entries) if entries else "(empty directory)") ''' script = ( code.replace('PATH', python_literal(target)) .replace('MAX_ENTRIES', str(max_entries)) ) result = self.run_python( code=script, script_path=None, args=[], stdin=None, timeout_seconds=20.0, max_output_chars=max_output_chars, ) if result.exit_code != 0: raise JupyterRuntimeError(result.stdout.strip() or 'remote list_dir failed') return result.stdout.strip() def list_files( self, path: str, *, kind: str, max_entries: int = 200, max_output_chars: int = 20000, ) -> list[dict[str, Any]]: target = self.resolve_workspace_path(path) code = r''' import json from datetime import datetime, timezone from pathlib import Path target = Path(PATH) entries = [] if target.exists() and target.is_dir(): for item in sorted(target.iterdir(), key=lambda p: p.stat().st_mtime, reverse=True)[:MAX_ENTRIES]: if not item.is_file(): continue stat = item.stat() entries.append({ "name": item.name, "path": str(item), "kind": KIND, "size": stat.st_size, "modified_at": datetime.fromtimestamp(stat.st_mtime, timezone.utc).isoformat().replace("+00:00", "Z"), }) print(json.dumps(entries, ensure_ascii=False)) ''' script = ( code.replace('PATH', python_literal(target)) .replace('MAX_ENTRIES', str(max_entries)) .replace('KIND', python_literal(kind)) ) result = self.run_python( code=script, script_path=None, args=[], stdin=None, timeout_seconds=20.0, max_output_chars=max_output_chars, ) if result.exit_code != 0: raise JupyterRuntimeError(result.stdout.strip() or 'remote list_files failed') try: payload = json.loads(result.stdout.strip() or '[]') except json.JSONDecodeError as exc: raise JupyterRuntimeError('remote list_files returned invalid JSON') from exc if not isinstance(payload, list): raise JupyterRuntimeError('remote list_files returned invalid payload') return [item for item in payload if isinstance(item, dict)] def read_file_bytes(self, path: str) -> tuple[bytes, str]: target = self.resolve_workspace_path(path) api_path = api_path_for_absolute_path(target) url = f'{self.binding.base_url}/api/contents/{quote(api_path, safe="/")}' response = self._session.get( url, params={'content': 1}, headers={'X-XSRFToken': self._xsrf_token}, timeout=60, ) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to read remote file: HTTP {response.status_code} {response.text[:500]}' ) payload = response.json() if payload.get('type') != 'file': raise JupyterRuntimeError('Remote path is not a file.') content = payload.get('content') file_format = payload.get('format') if not isinstance(content, str): raise JupyterRuntimeError('Remote file response did not contain content.') if file_format == 'base64': data = base64.b64decode(content) else: data = content.encode('utf-8') name = str(payload.get('name') or PurePosixPath(target).name or 'remote-file') return data, name def file_info(self, path: str) -> dict[str, Any]: target = self.resolve_workspace_path(path) api_path = api_path_for_absolute_path(target) url = f'{self.binding.base_url}/api/contents/{quote(api_path, safe="/")}' response = self._session.get( url, params={'content': 0}, headers={'X-XSRFToken': self._xsrf_token}, timeout=30, ) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to inspect remote file: HTTP {response.status_code} {response.text[:500]}' ) payload = response.json() if payload.get('type') != 'file': raise JupyterRuntimeError('Remote path is not a file.') name = str(payload.get('name') or PurePosixPath(target).name or 'remote-file') size = payload.get('size') return { 'name': name, 'path': target, 'size': int(size) if isinstance(size, int) and size >= 0 else None, } def open_file_stream(self, path: str) -> tuple[Any, str]: """打开远端文件流,由调用方负责关闭 response。""" target = self.resolve_workspace_path(path) api_path = api_path_for_absolute_path(target) url = f'{self.binding.base_url}/files/{quote(api_path, safe="/")}' response = self._session.get( url, headers={'X-XSRFToken': self._xsrf_token}, stream=True, timeout=(10, 300), ) if response.status_code >= 400: response.close() raise JupyterRuntimeError( f'Unable to stream remote file: HTTP {response.status_code} {response.text[:500]}' ) return response, str(PurePosixPath(target).name or 'remote-file') def read_text( self, path: str, *, start_line: int | None, end_line: int | None, max_output_chars: int, ) -> str: target = self.resolve_workspace_path(path) code = r''' from pathlib import Path target = Path(PATH) if not target.is_file(): raise SystemExit(f"Path is not a file: {target}") text = target.read_text(encoding="utf-8", errors="replace") start_line = START_LINE end_line = END_LINE if start_line is None and end_line is None: print(text, end="") else: lines = text.splitlines() start_idx = max((start_line or 1) - 1, 0) end_idx = end_line or len(lines) print("\n".join(f"{start_idx + idx + 1}: {line}" for idx, line in enumerate(lines[start_idx:end_idx])), end="") ''' script = ( code.replace('PATH', python_literal(target)) .replace('START_LINE', python_literal(start_line)) .replace('END_LINE', python_literal(end_line)) ) result = self.run_python( code=script, script_path=None, args=[], stdin=None, timeout_seconds=20.0, max_output_chars=max_output_chars, ) if result.exit_code != 0: raise JupyterRuntimeError(result.stdout.strip() or 'remote read_file failed') return result.stdout[:max_output_chars] def write_text( self, path: str, content: str, *, append: bool, newline_at_end: bool, max_output_chars: int, ) -> str: target = self.resolve_workspace_path(path) api_path = api_path_for_absolute_path(target) if not api_path: raise JupyterRuntimeError(f'Invalid remote path: {target}') parent_api = '/'.join(api_path.split('/')[:-1]) if parent_api: self._ensure_remote_dir(parent_api) final_content = content if append: final_content = self._read_remote_text_for_append(api_path) + content if newline_at_end and final_content and not final_content.endswith('\n'): final_content += '\n' encoded = base64.b64encode(final_content.encode('utf-8')).decode('ascii') url = f'{self.binding.base_url}/api/contents/{quote(api_path, safe="/")}' response = self._session.put( url, json={ 'type': 'file', 'format': 'base64', 'content': encoded, }, headers={ 'X-XSRFToken': self._xsrf_token, 'Content-Type': 'application/json', }, timeout=60, ) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to write remote file: HTTP {response.status_code} {response.text[:500]}' ) verb = 'appended' if append else 'wrote' return f'{verb} {target} ({len(final_content)} chars)' def _ensure_remote_dir(self, api_path: str) -> None: parts = [segment for segment in api_path.split('/') if segment] cumulative = '' for segment in parts: cumulative = f'{cumulative}/{segment}' if cumulative else segment url = f'{self.binding.base_url}/api/contents/{quote(cumulative, safe="/")}' resp = self._session.get( url, params={'content': 0}, headers={'X-XSRFToken': self._xsrf_token}, timeout=30, ) if resp.status_code == 200: kind = resp.json().get('type') if kind == 'directory': continue raise JupyterRuntimeError( f'Cannot create directory at {cumulative}: path exists as {kind}' ) if resp.status_code != 404: raise JupyterRuntimeError( f'Unable to inspect remote path {cumulative}: HTTP {resp.status_code} {resp.text[:300]}' ) create = self._session.put( url, json={'type': 'directory'}, headers={ 'X-XSRFToken': self._xsrf_token, 'Content-Type': 'application/json', }, timeout=30, ) if create.status_code >= 400: raise JupyterRuntimeError( f'Unable to create remote directory {cumulative}: HTTP {create.status_code} {create.text[:300]}' ) def _read_remote_text_for_append(self, api_path: str) -> str: url = f'{self.binding.base_url}/api/contents/{quote(api_path, safe="/")}' resp = self._session.get( url, params={'content': 1}, headers={'X-XSRFToken': self._xsrf_token}, timeout=60, ) if resp.status_code == 404: return '' if resp.status_code >= 400: raise JupyterRuntimeError( f'Unable to read existing remote file for append: HTTP {resp.status_code} {resp.text[:500]}' ) payload = resp.json() if payload.get('type') != 'file': raise JupyterRuntimeError( f'Cannot append: remote path {api_path} is not a file (type={payload.get("type")})' ) existing = payload.get('content', '') if not isinstance(existing, str): return '' if payload.get('format') == 'base64': return base64.b64decode(existing).decode('utf-8', errors='replace') return existing def resolve_workspace_path(self, raw_path: str) -> str: value = raw_path.strip() or '.' # 系统提示里给的"当前会话目录"通常是本地后端的绝对/相对路径,例如 # `.port_sessions/accounts//sessions//input` 或 # `/home/mi/.../.port_sessions/accounts//sessions//output/foo.csv`。 # 远端 Jupyter 没这套目录结构。识别这种 pattern 后,把 input/output/ # scratchpad bucket 映射到远端 workspace 的同名目录。 bucket_match = self._match_local_session_bucket(value) if bucket_match is not None: bucket, rest = bucket_match tail = f'/{rest}' if rest else '' return normalize_posix_path( f'{self.binding.workspace_cwd}/{bucket}{tail}' ) if value.startswith('/'): return normalize_posix_path(value) path = PurePosixPath(value) if self.platform_root and path.parts: platform_heads = {'skills', 'src', '标签定义'} if path.parts[0] in platform_heads or value == 'pyproject.toml': return normalize_posix_path(f'{self.platform_root}/{value}') if path.parts and path.parts[0] in {'output', 'outputs'}: tail = PurePosixPath(*path.parts[1:]) if len(path.parts) > 1 else PurePosixPath() return normalize_posix_path(f'{self.binding.workspace_cwd}/output/{tail}') if path.parts and path.parts[0] in {'scratchpad', 'scratch'}: tail = PurePosixPath(*path.parts[1:]) if len(path.parts) > 1 else PurePosixPath() return normalize_posix_path(f'{self.binding.workspace_cwd}/scratchpad/{tail}') if path.parts and path.parts[0] in {'input', 'inputs'}: tail = PurePosixPath(*path.parts[1:]) if len(path.parts) > 1 else PurePosixPath() return normalize_posix_path(f'{self.binding.workspace_cwd}/input/{tail}') return normalize_posix_path(f'{self.binding.workspace_cwd}/{value}') _LOCAL_SESSION_BUCKET_RE = re.compile( r'(?:^|/)\.port_sessions/accounts/[^/]+/sessions/[^/]+/' r'(?Pinput|inputs|output|outputs|scratchpad|scratch)' r'(?:/(?P.*))?$' ) @classmethod def _match_local_session_bucket(cls, raw: str) -> tuple[str, str] | None: m = cls._LOCAL_SESSION_BUCKET_RE.search(raw) if not m: return None bucket = m.group('bucket') normalized = { 'input': 'input', 'inputs': 'input', 'output': 'output', 'outputs': 'output', 'scratchpad': 'scratchpad', 'scratch': 'scratchpad', }[bucket] rest = m.group('rest') or '' return normalized, rest def _remote_env_prefix(self) -> str: exports = { 'ZK_AGENT_WORKSPACE': self.binding.workspace_cwd, 'ZK_AGENT_OUTPUT': f'{self.binding.workspace_cwd}/output', 'ZK_AGENT_SCRATCHPAD': f'{self.binding.workspace_cwd}/scratchpad', 'ZK_AGENT_PYTHON_ENV': self.python_env_path, 'AUTORESEARCH_CHAT_ROOT': self.chat_workspace_root, } if self.platform_root: exports['ZK_AGENT_PLATFORM_ROOT'] = self.platform_root exports['ZK_AGENT_SKILLS_ROOT'] = self.skills_root lines = [f'export {key}={shlex.quote(value)}' for key, value in exports.items()] lines.append(f'export PATH={shlex.quote(self.python_env_path + "/bin")}:$PATH') # CloudML 的 Jupyter Python 可能来自 Nix,pip 安装的科学计算轮子会依赖系统库。 # 只在库文件存在时预加载,避免 pandas/numpy/pyarrow 因 libz/libstdc++ 找不到而失败。 lines.extend( [ 'if [ -f /lib/x86_64-linux-gnu/libz.so.1 ] && ' '[ -f /lib/x86_64-linux-gnu/libstdc++.so.6 ]; then ' 'export LD_PRELOAD=/lib/x86_64-linux-gnu/libz.so.1:' '/lib/x86_64-linux-gnu/libstdc++.so.6${LD_PRELOAD:+:$LD_PRELOAD}; ' 'fi' ] ) if self.platform_root: lines.append(f'export PYTHONPATH={shlex.quote(self.platform_root)}:${{PYTHONPATH:-}}') return '; '.join(lines) + '; ' def _ensure_terminal(self) -> str: if self._terminal_name: return self._terminal_name response = self._session.post( f'{self.binding.base_url}/api/terminals', headers={'X-XSRFToken': self._xsrf_token}, timeout=20, ) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to create Jupyter terminal: HTTP {response.status_code} {response.text[:500]}' ) payload = response.json() name = str(payload.get('name') or '').strip() if not name: raise JupyterRuntimeError('Jupyter terminal response did not contain a terminal name.') self._terminal_name = name return name def _terminal_websocket_url(self, name: str) -> str: parsed = urlparse(self.binding.base_url) scheme = 'wss' if parsed.scheme == 'https' else 'ws' base_path = parsed.path.rstrip('/') path = f'{base_path}/terminals/websocket/{quote(name)}' return urlunparse((scheme, parsed.netloc, path, '', '', '')) def _cookie_header(self) -> str: return '; '.join( f'{cookie.name}={cookie.value}' for cookie in self._session.cookies ) def _drain_initial_messages(self, ws: Any, *, max_wait_seconds: float) -> None: deadline = time.monotonic() + max_wait_seconds old_timeout = ws.gettimeout() try: while time.monotonic() < deadline: ws.settimeout(max(0.05, deadline - time.monotonic())) try: ws.recv() except Exception: break finally: ws.settimeout(old_timeout) def _collect_until_marker( self, ws: Any, *, marker: str, timeout_seconds: float, max_output_chars: int, cancel_event: Any | None, ) -> RemoteCommandResult: deadline = time.monotonic() + timeout_seconds transcript: list[str] = [] exit_code: int | None = None marker_pattern = re.compile(rf'{re.escape(marker)}:(\d+)') while time.monotonic() < deadline: if cancel_event is not None and cancel_event.is_set(): try: ws.send(json.dumps(['stdin', '\x03'])) except Exception: pass raise JupyterRuntimeError('Run cancelled by user') ws.settimeout(max(0.1, min(1.0, deadline - time.monotonic()))) try: raw = ws.recv() except Exception: continue for chunk in parse_terminal_payload(raw): match = marker_pattern.search(chunk) if match: exit_code = int(match.group(1)) chunk = marker_pattern.sub('', chunk) if chunk: transcript.append(chunk) break transcript.append(chunk) if exit_code is not None: break if exit_code is None: try: ws.send(json.dumps(['stdin', '\x03'])) except Exception: pass return RemoteCommandResult( exit_code=124, stdout=truncate_text(''.join(transcript), max_output_chars), timed_out=True, ) return RemoteCommandResult( exit_code=exit_code, stdout=truncate_text(clean_terminal_output(''.join(transcript)), max_output_chars), ) def _remote_write_text_command(self, path: str, content: str) -> str: encoded = base64.b64encode(content.encode('utf-8')).decode('ascii') script = '; '.join( [ 'import base64', 'from pathlib import Path', f'path = Path({json.dumps(path)})', 'path.parent.mkdir(parents=True, exist_ok=True)', f'path.write_bytes(base64.b64decode({json.dumps(encoded)}))', ] ) return f'python3 -c {shlex.quote(script)}' class JupyterRuntimeManager: def __init__(self) -> None: self._lock = threading.RLock() self._sessions: dict[tuple[str, str], JupyterRuntimeSession] = {} def bind_session( self, *, account_id: str, session_id: str, base_url: str, password: str, workspace_root: str = DEFAULT_JUPYTER_WORKSPACE_ROOT, project_root: Path | None = None, ) -> JupyterRuntimeSession: runtime = JupyterRuntimeSession.connect( account_id=account_id, session_id=session_id, base_url=base_url, password=password, workspace_root=workspace_root, project_root=project_root, ) with self._lock: self._sessions[(account_id, session_id)] = runtime return runtime def get(self, account_id: str, session_id: str) -> JupyterRuntimeSession | None: with self._lock: return self._sessions.get((account_id, session_id)) def restore_session( self, *, account_id: str, session_id: str, payload: dict[str, Any], ) -> JupyterRuntimeSession: runtime = JupyterRuntimeSession.from_persisted(payload) if ( runtime.binding.account_id != account_id or runtime.binding.session_id != session_id ): raise JupyterRuntimeError('Persisted Jupyter binding does not match session.') with self._lock: self._sessions[(account_id, session_id)] = runtime return runtime def session_payload(self, account_id: str, session_id: str) -> dict[str, Any]: runtime = self.get(account_id, session_id) if runtime is None: return { 'connected': False, 'account_id': account_id, 'session_id': session_id, } return runtime.to_dict() def sync_account(self, account_id: str, project_root: Path) -> list[dict[str, Any]]: with self._lock: runtimes = [ runtime for (runtime_account, _), runtime in self._sessions.items() if runtime_account == account_id ] results: list[dict[str, Any]] = [] for runtime in runtimes: payload = runtime.sync_runtime_bundle(project_root) results.append( { 'session_id': runtime.binding.session_id, **payload, } ) return results def normalize_jupyter_base_url(value: str) -> str: raw = value.strip() if not raw: raise JupyterRuntimeError('Jupyter URL is empty.') parsed = urlparse(raw) if not parsed.scheme: parsed = urlparse(f'https://{raw}') if parsed.scheme not in {'http', 'https'} or not parsed.netloc: raise JupyterRuntimeError(f'Invalid Jupyter URL: {value}') path = parsed.path.rstrip('/') for marker in ('/lab', '/tree', '/notebooks', '/login'): index = path.find(marker) if index >= 0: path = path[:index] break return urlunparse((parsed.scheme, parsed.netloc, path.rstrip('/'), '', '', '')).rstrip('/') def normalize_posix_path(value: str) -> str: path = PurePosixPath(value) if not str(path).startswith('/'): path = PurePosixPath('/') / path normalized = str(path) normalized = re.sub(r'/+', '/', normalized) return normalized.rstrip('/') or '/' def sanitize_remote_path_part(value: str) -> str: return re.sub(r'[^A-Za-z0-9_.-]+', '_', value).strip('._') or uuid.uuid4().hex def api_path_for_absolute_path(path: str) -> str: normalized = normalize_posix_path(path) if normalized == '/root': return '' if normalized.startswith('/root/'): return normalized[len('/root/') :] return normalized.lstrip('/') def parse_terminal_payload(raw: str) -> list[str]: try: payload = json.loads(raw) except json.JSONDecodeError: return [raw] if isinstance(payload, list) and len(payload) >= 2: kind, content = payload[0], payload[1] if kind in {'stdout', 'stderr'} and isinstance(content, str): return [content] return [] def clean_terminal_output(text: str) -> str: # 去掉少量 ANSI 控制符,避免输出里带终端颜色和光标控制。 cleaned = re.sub(r'\x1b\[[0-?]*[ -/]*[@-~]', '', text) lines = [ line.rstrip('\r') for line in cleaned.splitlines() if line.strip() not in {'#', '>'} and not re.fullmatch(r'(>\s*)+', line.strip()) ] return '\n'.join(lines).strip('\r\n') def truncate_text(text: str, limit: int) -> str: if limit <= 0 or len(text) <= limit: return text omitted = len(text) - limit return f'{text[:limit]}\n... [truncated, {omitted} chars omitted]' def python_literal(value: object) -> str: return repr(value) def build_runtime_bundle(project_root: Path) -> tuple[bytes, str]: root = project_root.resolve() digest = calculate_runtime_bundle_digest(root) buffer = io.BytesIO() with tarfile.open(fileobj=buffer, mode='w:gz') as archive: for name in ('src', 'skills', '标签定义', 'pyproject.toml'): source = root / name if not source.exists(): continue archive.add(source, arcname=name, filter=_runtime_bundle_filter) data = buffer.getvalue() return data, digest def calculate_runtime_bundle_digest(project_root: Path) -> str: digest = hashlib.sha256() for name in ('src', 'skills', '标签定义', 'pyproject.toml'): source = project_root / name if not source.exists(): continue if source.is_file(): _update_digest_for_file(digest, source, Path(name)) continue for path in sorted(source.rglob('*')): rel = Path(name) / path.relative_to(source) if _should_skip_bundle_path(rel) or not path.is_file(): continue _update_digest_for_file(digest, path, rel) return digest.hexdigest() def _update_digest_for_file(digest: 'hashlib._Hash', path: Path, rel: Path) -> None: digest.update(rel.as_posix().encode('utf-8')) digest.update(b'\0') digest.update(path.read_bytes()) digest.update(b'\0') def _should_skip_bundle_path(path: Path | PurePosixPath) -> bool: parts = path.parts skipped_dirs = { '.git', '.venv', '.port_sessions', '.next', 'node_modules', '__pycache__', } if any(part in skipped_dirs for part in parts): return True return str(path).endswith(('.pyc', '.pyo', '.DS_Store')) def _runtime_bundle_filter(info: tarfile.TarInfo) -> tarfile.TarInfo | None: if _should_skip_bundle_path(PurePosixPath(info.name)): return None return info def _login_jupyter( http_session: Any, base_url: str, password: str, *, timeout_seconds: float, ) -> str: login_url = f'{base_url}/login?next=%2Flab' response = http_session.get(login_url, timeout=timeout_seconds) if response.status_code >= 400: raise JupyterRuntimeError( f'Unable to open Jupyter login page: HTTP {response.status_code}' ) xsrf_token = _extract_xsrf_token(response.text) or http_session.cookies.get('_xsrf') if not xsrf_token: raise JupyterRuntimeError('Unable to find Jupyter _xsrf token on login page.') post = http_session.post( f'{base_url}/login', data={'_xsrf': xsrf_token, 'password': password}, headers={'Referer': login_url}, timeout=timeout_seconds, allow_redirects=True, ) if post.status_code >= 400: raise JupyterRuntimeError( f'Jupyter login failed: HTTP {post.status_code} {post.text[:500]}' ) probe = http_session.get(f'{base_url}/api/contents', timeout=timeout_seconds) if probe.status_code >= 400: raise JupyterRuntimeError( f'Jupyter contents API probe failed after login: HTTP {probe.status_code} {probe.text[:500]}' ) return str(xsrf_token) def _extract_xsrf_token(html: str) -> str | None: patterns = [ r'name=["\']_xsrf["\']\s+value=["\']([^"\']+)["\']', r'value=["\']([^"\']+)["\']\s+name=["\']_xsrf["\']', ] for pattern in patterns: match = re.search(pattern, html) if match: return unescape(match.group(1)) return None