Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion src/arkruntime/selfhosted/envinit.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
from contextlib import suppress
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional
from typing import Any, List, Optional

from .types import Session, SkillRef

Expand All @@ -38,6 +38,7 @@ class Initializer:
def __init__(self, api: Any, options: InitializerOptions) -> None:
self.api = api
self.options = options
self._installed_skill_dirs: List[Path] = []
if not self.options.skills_dir:
self.options.skills_dir = str(Path(self.options.workdir) / "skills")

Expand All @@ -63,6 +64,20 @@ def setup(self, session: Session) -> None:
)
continue

def cleanup(self) -> None:
"""Remove only the skill directories installed by this initializer."""
errors = []
for path in self._installed_skill_dirs:
try:
shutil.rmtree(path)
except FileNotFoundError:
pass
except OSError as exc:
errors.append(f"{path}: {exc}")
self._installed_skill_dirs.clear()
if errors:
raise OSError("remove installed skills: " + "; ".join(errors))

def install_skill(self, session_id: str, skill: SkillRef) -> None:
Path(self.options.workdir).mkdir(parents=True, exist_ok=True)
Path(self.options.skills_dir).mkdir(parents=True, exist_ok=True)
Expand Down Expand Up @@ -93,6 +108,7 @@ def install_skill(self, session_id: str, skill: SkillRef) -> None:
source = _install_source_dir(Path(tmp))
target = Path(self.options.skills_dir) / name
backup = _replace_skill_dir(Path(source), target)
self._installed_skill_dirs.append(target)
if backup is not None:
try:
shutil.rmtree(backup)
Expand Down
18 changes: 15 additions & 3 deletions src/arkruntime/selfhosted/tool_result_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,11 @@
import hashlib
import json
import os
import re
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, Tuple
from typing import Dict, Optional, Tuple

from .types import ContentBlock, Event, new_user_custom_tool_result_event, new_user_tool_result_event, utc_now_iso

Expand All @@ -27,10 +28,14 @@ class ToolCallStoreDecision:
class FileToolResultStore:
"""File-backed ledger that avoids re-running side-effectful tool calls."""

def __init__(self, workdir: str) -> None:
def __init__(self, workdir: str, session_id: Optional[str] = None) -> None:
if not workdir:
raise ValueError("workdir must not be empty")
self.dir = Path(workdir) / ".ma_self_host_worker" / "tool_ledger"
self.dir = Path(workdir) / ".ma_self_hosted_worker" / "tool_ledger"
if session_id is not None:
if not session_id:
raise ValueError("session id must not be empty")
self.dir /= _session_ledger_name(session_id)
self.dir.mkdir(parents=True, exist_ok=True, mode=0o700)

def recover(self) -> Tuple[Dict[str, Event], Dict[str, bool]]:
Expand Down Expand Up @@ -144,6 +149,13 @@ def _unknown_tool_execution_result(call_id: str, event: Event) -> Event:
return new_user_tool_result_event(call_id, content, True, event.session_thread_id)


def _session_ledger_name(session_id: str) -> str:
if re.fullmatch(r"[A-Za-z0-9._-]+", session_id) and session_id not in (".", ".."):
return session_id
digest = hashlib.sha256(session_id.encode("utf-8")).hexdigest()
return f"session-{digest}"


def _sync_directory(path: Path) -> None:
try:
fd = os.open(str(path), os.O_RDONLY)
Expand Down
41 changes: 18 additions & 23 deletions src/arkruntime/selfhosted/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,9 @@

from __future__ import annotations

import hashlib
import logging
import os
import random
import re
import socket
import threading
import time
Expand Down Expand Up @@ -239,7 +237,7 @@ def run(self) -> None:
raise poller.error
return
try:
self._handle_item(_claimed_work_from_item(item), use_workdir_as_session=False)
self._handle_item(_claimed_work_from_item(item))
except (IdleTimeout, SessionTerminated):
pass
except Exception as exc: # noqa: BLE001 - continue polling after a bad work item.
Expand All @@ -250,20 +248,21 @@ def run(self) -> None:
def handle_item(self, options: HandleItemOptions) -> None:
work = self._claimed_work_from_options(options)
try:
self._handle_item(work, use_workdir_as_session=True)
self._handle_item(work)
except (IdleTimeout, SessionTerminated):
return

def _handle_item(self, work: _ClaimedWork, *, use_workdir_as_session: bool) -> None:
def _handle_item(self, work: _ClaimedWork) -> None:
if not work.environment_id:
work.environment_id = self.options.environment_id or os.environ.get("MA_ENVIRONMENT_ID", "")
heartbeat_stop = threading.Event()
work_stop = _CombinedStopEvent(self._stop, heartbeat_stop)
heartbeat_done = threading.Event()
heartbeat_cause = {"value": ""}
heartbeat = None
initializer = None
try:
workdir = self._workdir_for(work.session_id, use_workdir_as_session)
workdir = self._workdir()
heartbeat_thread = threading.Thread(
target=self._heartbeat_loop,
args=(work, heartbeat_stop, heartbeat_done, heartbeat_cause),
Expand All @@ -278,14 +277,15 @@ def _handle_item(self, work: _ClaimedWork, *, use_workdir_as_session: bool) -> N
raise ValueError("session response is empty")
if not session.id:
session.id = work.session_id
Initializer(
initializer = Initializer(
self.api,
InitializerOptions(workdir=workdir, logger=self.options.logger),
).setup(session)
)
initializer.setup(session)
if work_stop.is_set():
return
tool_context = self._tool_context(workdir, work_stop)
store = FileToolResultStore(workdir)
store = FileToolResultStore(workdir, work.session_id)
runner = SessionToolRunner(
self.api,
work.session_id,
Expand All @@ -303,6 +303,11 @@ def _handle_item(self, work: _ClaimedWork, *, use_workdir_as_session: bool) -> N
)
runner.run()
finally:
if initializer is not None:
try:
initializer.cleanup()
except OSError as exc:
self.options.logger.warning("cleanup session skills failed: %s", exc)
heartbeat_stop.set()
if heartbeat is not None:
heartbeat_done.wait(timeout=DEFAULT_HEARTBEAT_SECONDS + 1)
Expand Down Expand Up @@ -400,14 +405,11 @@ def _tool_context(self, workdir: str, cancel_event: Any) -> ToolContext:
cancel_event=cancel_event,
)

def _workdir_for(self, session_id: str, use_workdir_as_session: bool) -> str:
def _workdir(self) -> str:
"""Return the shared worker workdir used for tool cwd and installed skills."""
root = str(Path(self.options.workdir or ".").resolve())
if use_workdir_as_session:
Path(root).mkdir(parents=True, exist_ok=True)
return root
workdir = str(Path(root) / _session_workdir_name(session_id))
Path(workdir).mkdir(parents=True, exist_ok=True)
return workdir
Path(root).mkdir(parents=True, exist_ok=True)
return root

def _claimed_work_from_options(self, options: HandleItemOptions) -> _ClaimedWork:
work_id = options.work_id or os.environ.get("MA_WORK_ID", "")
Expand Down Expand Up @@ -464,13 +466,6 @@ def _should_stop_item(heartbeat_cause: str) -> bool:
}


def _session_workdir_name(session_id: str) -> str:
if re.fullmatch(r"[A-Za-z0-9._-]+", session_id) and session_id not in (".", ".."):
return session_id
digest = hashlib.sha256(session_id.encode("utf-8")).hexdigest()
return f"session-{digest}"


class _CombinedStopEvent:
def __init__(self, *events: threading.Event) -> None:
self._events = events
Expand Down
7 changes: 7 additions & 0 deletions tests/selfhosted/test_envinit.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,13 @@ def test_setup_installs_skill_under_resolved_metadata_name(tmp_path) -> None:
assert (tmp_path / "skills" / "canonical-skill-name" / "SKILL.md").read_text() == "hello"
assert not (tmp_path / "skills" / "skill-1").exists()

retained = tmp_path / "skills" / "retained"
retained.mkdir()
initializer.cleanup()

assert not (tmp_path / "skills" / "canonical-skill-name").exists()
assert retained.is_dir()


def test_install_closes_skill_body_when_archive_copy_fails(tmp_path) -> None:
body = _CloseTrackingBody(b"archive-too-large")
Expand Down
19 changes: 19 additions & 0 deletions tests/selfhosted/test_tool_result_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

def test_recovery_uses_persisted_call_id(tmp_path) -> None:
store = FileToolResultStore(str(tmp_path))
assert store.dir == tmp_path / ".ma_self_hosted_worker" / "tool_ledger"
store.begin("call-1", Event(id="event-1", type="agent.tool_use", name="bash"))

pending, _ = store.recover()
Expand All @@ -21,3 +22,21 @@ def test_recovery_removes_stale_temporary_records(tmp_path) -> None:
store.recover()

assert not stale.exists()


def test_session_store_isolates_sessions(tmp_path) -> None:
first = FileToolResultStore(str(tmp_path), "session-a")
second = FileToolResultStore(str(tmp_path), "session-b")
assert first.dir == tmp_path / ".ma_self_hosted_worker" / "tool_ledger" / "session-a"

first.begin("call-1", Event(id="event-1", type="agent.tool_use", name="bash"))

assert second.recover() == ({}, {})


def test_session_store_sanitizes_session_id(tmp_path) -> None:
store = FileToolResultStore(str(tmp_path), "../../outside")
base = tmp_path / ".ma_self_hosted_worker" / "tool_ledger"

assert store.dir.parent == base
assert store.dir.name.startswith("session-")
11 changes: 4 additions & 7 deletions tests/selfhosted/test_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,22 +185,19 @@ def test_worker_options_preserve_legacy_positional_order() -> None:
custom_tools = {"custom": object()}
logger = logging.getLogger("legacy-positional-worker")

options = EnvironmentWorkerOptions(
"env-1", "worker-1", ".", False, None, None, 60, custom_tools, logger
)
options = EnvironmentWorkerOptions("env-1", "worker-1", ".", False, None, None, 60, custom_tools, logger)

assert options.custom_tools is custom_tools
assert options.logger is logger
assert options.tool_timeout_seconds is None


def test_session_id_cannot_escape_worker_root(tmp_path) -> None:
def test_worker_uses_configured_workdir(tmp_path) -> None:
worker = EnvironmentWorker(object(), EnvironmentWorkerOptions(workdir=str(tmp_path)))

workdir = worker._workdir_for("../../outside", use_workdir_as_session=False)
workdir = worker._workdir()

assert str(tmp_path.resolve()) in workdir
assert ".." not in workdir
assert workdir == str(tmp_path.resolve())


@pytest.mark.parametrize("status_code", [408, 409, 412, 429])
Expand Down
Loading