diff --git a/MANIFEST.in b/MANIFEST.in new file mode 100644 index 0000000..1eeef06 --- /dev/null +++ b/MANIFEST.in @@ -0,0 +1 @@ +prune tests diff --git a/README.md b/README.md index 129b16a..c92baba 100644 --- a/README.md +++ b/README.md @@ -326,3 +326,9 @@ uv run pre-commit install # one-time setup - httpx >= 0.23.0 - pydantic >= 2.0 - typing-extensions >= 4.7 + +## License + +This project is licensed under the Apache License 2.0. See [LICENSE](./LICENSE). +For third-party open-source software notices, see +[THIRD_PARTY_NOTICES.md](./THIRD_PARTY_NOTICES.md). diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md new file mode 100644 index 0000000..00d30fb --- /dev/null +++ b/THIRD_PARTY_NOTICES.md @@ -0,0 +1,38 @@ +# Third-Party Notices + +This repository contains code that is derived from or structurally adapted +from third-party open-source projects. + +## Anthropic self-hosted worker SDK + +Portions of the self-hosted worker lifecycle and local agent tool +implementations under `src/arkruntime/selfhosted` are structurally adapted +from Anthropic's self-hosted worker SDK implementations: + +- https://github.com/anthropics/anthropic-sdk-python +- https://github.com/anthropics/anthropic-sdk-go + +The upstream projects are licensed under the MIT License. The MIT copyright +and permission notice is preserved below as required by that license. + +```text +Copyright 2023 Anthropic, PBC. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +``` diff --git a/examples/README.md b/examples/README.md index a238b60..b9f5fb5 100644 --- a/examples/README.md +++ b/examples/README.md @@ -19,6 +19,9 @@ python examples/async_responses_create.py | `environments.py` | Managed-Agents: Environment lifecycle — Create/Get/List/Update/Delete (cloud + unrestricted networking) | | `sessions_loop.py` | Managed-Agents: end-to-end agent loop — Agent + Env + Session, send user.message, stream events until idle | | `memory_stores.py` | Managed-Agents: MemoryStore + nested Memory CRUD | +| `self_hosted_worker.py` | Managed-Agents: self-hosted worker poll / handle loop | + +`self_hosted_worker.py` uses the client's production default `https://ark.cn-beijing.volces.com/api/v3`. The Managed-Agents examples additionally accept `ARK_MODEL_ID` for the model id (falls back to a `${YOUR_MODEL_ID}` placeholder that will 400 at runtime). diff --git a/examples/self_hosted_worker.py b/examples/self_hosted_worker.py new file mode 100644 index 0000000..c8ae820 --- /dev/null +++ b/examples/self_hosted_worker.py @@ -0,0 +1,88 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +# Prepare the Python environment from the repository root before running: +# +# python3 -m venv .venv +# source .venv/bin/activate +# +# python -m pip install -U pip +# python -m pip install -e . +# python examples/self_hosted_worker.py + +from __future__ import annotations + +import logging +import os +import signal +import sys + +_SETUP_INSTRUCTIONS = """Run these commands from the repository root: + +python3 -m venv .venv +source .venv/bin/activate + +python -m pip install -U pip +python -m pip install -e . +python examples/self_hosted_worker.py""" + +try: + from arkruntime import Ark + from arkruntime.selfhosted import ClientAPI, EnvironmentWorker, EnvironmentWorkerOptions +except ModuleNotFoundError as exc: + print(f"failed to import arkruntime: {exc}", file=sys.stderr) + print(_SETUP_INSTRUCTIONS, file=sys.stderr) + raise SystemExit(1) from exc + +logger = logging.getLogger("arkruntime.selfhosted.example") + + +def configure_logging() -> None: + level_name = os.environ.get("ARK_LOG", "info").upper() + level = getattr(logging, level_name, logging.INFO) + logging.basicConfig( + level=level, + format="%(asctime)s %(levelname)s %(name)s %(message)s", + force=True, + ) + + +def required_env(name: str) -> str: + value = os.environ.get(name, "") + if not value: + raise RuntimeError(f"{name} is required") + return value + + +def main() -> None: + configure_logging() + base_url = os.environ.get("ARK_BASE_URL", "") + environment_id = required_env("MA_ENVIRONMENT_ID") + options = EnvironmentWorkerOptions( + environment_id=environment_id, + worker_id=os.environ.get("MA_WORKER_ID", ""), + workdir=os.environ.get("MA_WORKDIR", "."), + ) + client_options = {"api_key": required_env("ARK_API_KEY")} + if base_url: + client_options["base_url"] = base_url + client = Ark(**client_options) + worker = EnvironmentWorker(ClientAPI(client), options) + for sig in (signal.SIGINT, signal.SIGTERM): + signal.signal(sig, lambda _signum, _frame: worker.close()) + logger.info( + "starting self-hosted worker base_url=%s environment_id=%s worker_id=%s workdir=%s", + base_url or "default", + environment_id, + options.worker_id, + options.workdir, + ) + try: + worker.run() + finally: + worker.close() + client.close() + + +if __name__ == "__main__": + main() diff --git a/pyproject.toml b/pyproject.toml index 92e83db..6ee7c2f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -44,6 +44,9 @@ dev = [ "respx>=0.21", ] +[tool.setuptools] +license-files = ["LICENSE", "THIRD_PARTY_NOTICES.md"] + [tool.setuptools.packages.find] where = ["src"] include = ["arkruntime*"] diff --git a/src/arkruntime/resources/__init__.py b/src/arkruntime/resources/__init__.py index 8202b7c..e31e22f 100644 --- a/src/arkruntime/resources/__init__.py +++ b/src/arkruntime/resources/__init__.py @@ -7,7 +7,7 @@ from .chat import AsyncChat, Chat from .content_generation import AsyncContentGeneration, ContentGeneration from .embeddings import AsyncEmbeddings, Embeddings -from .environments import AsyncEnvironments, Environments +from .environments import AsyncEnvironments, AsyncEnvironmentWork, Environments, EnvironmentWork from .files import AsyncFiles, Files from .images import AsyncImages, Images from .memory_stores import ( @@ -56,6 +56,8 @@ "AsyncAgents", "Environments", "AsyncEnvironments", + "EnvironmentWork", + "AsyncEnvironmentWork", "MemoryStores", "AsyncMemoryStores", "Memories", diff --git a/src/arkruntime/resources/environments/__init__.py b/src/arkruntime/resources/environments/__init__.py index 8181e7f..add92af 100644 --- a/src/arkruntime/resources/environments/__init__.py +++ b/src/arkruntime/resources/environments/__init__.py @@ -1,3 +1,4 @@ from arkruntime.resources.environments.environments import AsyncEnvironments, Environments +from arkruntime.resources.environments.work import AsyncEnvironmentWork, EnvironmentWork -__all__ = ["Environments", "AsyncEnvironments"] +__all__ = ["Environments", "AsyncEnvironments", "EnvironmentWork", "AsyncEnvironmentWork"] diff --git a/src/arkruntime/resources/environments/environments.py b/src/arkruntime/resources/environments/environments.py index 3d66d2b..ad184ca 100644 --- a/src/arkruntime/resources/environments/environments.py +++ b/src/arkruntime/resources/environments/environments.py @@ -5,6 +5,7 @@ import httpx from ..._base_client import make_request_options +from ..._compat import cached_property from ..._managed_agents_serialize import dump_body from ..._resource import AsyncAPIResource, SyncAPIResource from ..._types import NOT_GIVEN, Body, Headers, NotGiven, Query @@ -13,6 +14,7 @@ from ...types.environment.environment import Environment from ...types.environment.environment_scope import EnvironmentScope from ...types.environment.list_environments_response import ListEnvironmentsResponse +from .work import AsyncEnvironmentWork, EnvironmentWork __all__ = ["Environments", "AsyncEnvironments"] @@ -29,6 +31,10 @@ def _list_query(*, limit=NOT_GIVEN, page=NOT_GIVEN) -> dict: class Environments(SyncAPIResource): + @cached_property + def work(self) -> EnvironmentWork: + return EnvironmentWork(self._client) + def create( self, *, @@ -123,6 +129,10 @@ def delete( class AsyncEnvironments(AsyncAPIResource): + @cached_property + def work(self) -> AsyncEnvironmentWork: + return AsyncEnvironmentWork(self._client) + async def create( self, *, diff --git a/src/arkruntime/resources/environments/work.py b/src/arkruntime/resources/environments/work.py new file mode 100644 index 0000000..ee342b9 --- /dev/null +++ b/src/arkruntime/resources/environments/work.py @@ -0,0 +1,234 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional +from urllib.parse import quote + +import httpx + +from ..._base_client import make_request_options +from ..._managed_agents_serialize import dump_body +from ..._resource import AsyncAPIResource, SyncAPIResource +from ...types.environment.heartbeat_work_response import HeartbeatWorkResponse +from ...types.environment.stop_work_body import StopWorkBody +from ...types.environment.work_item import WorkItem + +__all__ = ["EnvironmentWork", "AsyncEnvironmentWork"] + +_WORKER_ID_HEADER = "Ark-Worker-ID" + + +def _path(environment_id: str, suffix: str) -> str: + if not environment_id: + raise ValueError("environment_id is required") + return f"/environments/{quote(environment_id, safe='')}/work/{suffix.lstrip('/')}" + + +def _query(*, block_ms: int = 0, reclaim_older_than_ms: int = 0) -> dict: + query = {} + if block_ms > 0: + query["block_ms"] = block_ms + if reclaim_older_than_ms > 0: + query["reclaim_older_than_ms"] = reclaim_older_than_ms + return query + + +def _worker_headers(worker_id: str, extra_headers: Optional[dict]) -> dict: + headers = dict(extra_headers or {}) + if worker_id: + headers[_WORKER_ID_HEADER] = worker_id + return headers + + +class EnvironmentWork(SyncAPIResource): + def poll( + self, + environment_id: str, + *, + worker_id: str = "", + block_ms: int = 999, + reclaim_older_than_ms: int = 0, + extra_headers=None, + extra_query=None, + timeout: float | httpx.Timeout | None = None, + ) -> Optional[WorkItem]: + raw = self._get( + _path(environment_id, "poll"), + options=make_request_options( + query=_query(block_ms=block_ms, reclaim_older_than_ms=reclaim_older_than_ms), + extra_headers=_worker_headers(worker_id, extra_headers), + extra_query=extra_query, + timeout=timeout, + ), + cast_to=object, + ) + if not isinstance(raw, dict) or not raw.get("id"): + return None + try: + return WorkItem.model_validate(raw) + except Exception as exc: + raise ValueError(f"invalid WorkItem response: {exc}") from exc + + def ack( + self, + environment_id: str, + work_id: str, + *, + worker_id: str = "", + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> WorkItem: + if not work_id: + raise ValueError("work_id is required") + return self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/ack"), + options=make_request_options( + extra_headers=_worker_headers(worker_id, extra_headers), + timeout=timeout, + ), + cast_to=WorkItem, + ) + + def heartbeat( + self, + environment_id: str, + work_id: str, + *, + expected_last_heartbeat: str = "", + desired_ttl_seconds: int = 0, + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> HeartbeatWorkResponse: + if not work_id: + raise ValueError("work_id is required") + query = {} + if expected_last_heartbeat: + query["expected_last_heartbeat"] = expected_last_heartbeat + if desired_ttl_seconds > 0: + query["desired_ttl_seconds"] = desired_ttl_seconds + return self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/heartbeat"), + options=make_request_options( + query=query, + extra_headers=extra_headers, + timeout=timeout, + ), + cast_to=HeartbeatWorkResponse, + ) + + def stop( + self, + environment_id: str, + work_id: str, + *, + force: bool = False, + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> WorkItem: + if not work_id: + raise ValueError("work_id is required") + body = StopWorkBody(force=True) if force else StopWorkBody() + return self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/stop"), + body=dump_body(body.model_dump(exclude_none=True, by_alias=True)), + options=make_request_options(extra_headers=extra_headers, timeout=timeout), + cast_to=WorkItem, + ) + + +class AsyncEnvironmentWork(AsyncAPIResource): + async def poll( + self, + environment_id: str, + *, + worker_id: str = "", + block_ms: int = 999, + reclaim_older_than_ms: int = 0, + extra_headers=None, + extra_query=None, + timeout: float | httpx.Timeout | None = None, + ) -> Optional[WorkItem]: + raw = await self._get( + _path(environment_id, "poll"), + options=make_request_options( + query=_query(block_ms=block_ms, reclaim_older_than_ms=reclaim_older_than_ms), + extra_headers=_worker_headers(worker_id, extra_headers), + extra_query=extra_query, + timeout=timeout, + ), + cast_to=object, + ) + if not isinstance(raw, dict) or not raw.get("id"): + return None + try: + return WorkItem.model_validate(raw) + except Exception as exc: + raise ValueError(f"invalid WorkItem response: {exc}") from exc + + async def ack( + self, + environment_id: str, + work_id: str, + *, + worker_id: str = "", + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> WorkItem: + if not work_id: + raise ValueError("work_id is required") + return await self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/ack"), + options=make_request_options( + extra_headers=_worker_headers(worker_id, extra_headers), + timeout=timeout, + ), + cast_to=WorkItem, + ) + + async def heartbeat( + self, + environment_id: str, + work_id: str, + *, + expected_last_heartbeat: str = "", + desired_ttl_seconds: int = 0, + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> HeartbeatWorkResponse: + if not work_id: + raise ValueError("work_id is required") + query = {} + if expected_last_heartbeat: + query["expected_last_heartbeat"] = expected_last_heartbeat + if desired_ttl_seconds > 0: + query["desired_ttl_seconds"] = desired_ttl_seconds + return await self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/heartbeat"), + options=make_request_options( + query=query, + extra_headers=extra_headers, + timeout=timeout, + ), + cast_to=HeartbeatWorkResponse, + ) + + async def stop( + self, + environment_id: str, + work_id: str, + *, + force: bool = False, + extra_headers=None, + timeout: float | httpx.Timeout | None = None, + ) -> WorkItem: + if not work_id: + raise ValueError("work_id is required") + body = StopWorkBody(force=True) if force else StopWorkBody() + return await self._post_without_retry( + _path(environment_id, f"{quote(work_id, safe='')}/stop"), + body=dump_body(body.model_dump(exclude_none=True, by_alias=True)), + options=make_request_options(extra_headers=extra_headers, timeout=timeout), + cast_to=WorkItem, + ) diff --git a/src/arkruntime/selfhosted/__init__.py b/src/arkruntime/selfhosted/__init__.py new file mode 100644 index 0000000..476cdaf --- /dev/null +++ b/src/arkruntime/selfhosted/__init__.py @@ -0,0 +1,50 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from .client import ClientAPI +from .envinit import Initializer, InitializerOptions +from .session_tool_runner import SessionToolRunner, SessionToolRunnerOptions +from .tool_result_store import FileToolResultStore +from .tools import Tool, ToolContext, ToolResult, ToolSet, default_toolset +from .types import ( + EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT, + APIError, + Event, + HeartbeatResponse, + IdleTimeout, + ListEventsResponse, + Session, + SessionTerminated, + SkillRef, + WorkItem, +) +from .worker import EnvironmentWorker, EnvironmentWorkerOptions, HandleItemOptions, WorkPoller, WorkPollerOptions + +__all__ = [ + "APIError", + "ClientAPI", + "EnvironmentWorker", + "EnvironmentWorkerOptions", + "Event", + "EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT", + "FileToolResultStore", + "HandleItemOptions", + "HeartbeatResponse", + "IdleTimeout", + "Initializer", + "InitializerOptions", + "ListEventsResponse", + "Session", + "SessionToolRunner", + "SessionToolRunnerOptions", + "SessionTerminated", + "SkillRef", + "Tool", + "ToolContext", + "ToolResult", + "ToolSet", + "WorkItem", + "WorkPoller", + "WorkPollerOptions", + "default_toolset", +] diff --git a/src/arkruntime/selfhosted/client.py b/src/arkruntime/selfhosted/client.py new file mode 100644 index 0000000..f15cc3e --- /dev/null +++ b/src/arkruntime/selfhosted/client.py @@ -0,0 +1,420 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import contextlib +import json +import os +import random +import time +from dataclasses import replace +from typing import Any, Dict, Iterator, List, Optional +from urllib.parse import quote + +import httpx + +from arkruntime import Ark +from arkruntime._constants import CLIENT_REQUEST_HEADER, SERVER_REQUEST_HEADER +from arkruntime._types import NOT_GIVEN + +from .types import ( + EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT, + APIError, + Event, + HeartbeatResponse, + ListEventsResponse, + Session, + SkillContent, + SkillRef, + WorkItem, +) + +RETRYABLE_STATUS_CODES = {408, 429} +SKILL_TYPE_SKILL_HUB = "skill_hub" +SKILL_HUB_BASE_URL = "https://skills.volces.com/v1/skills" +MAX_SKILL_HUB_METADATA_BYTES = 1 << 20 +SENSITIVE_EXTERNAL_HEADERS = ( + "authorization", + "proxy-authorization", + "x-api-key", + "x-ark-api-key", +) + + +class ClientAPI: + """Adapter from the public Ark client to the self-hosted worker API.""" + + def __init__(self, client: Ark) -> None: + if client is None: + raise ValueError("ark client is required") + self.client = client + + def poll_work( + self, + environment_id: str, + *, + worker_id: str = "", + block_ms: int = 999, + reclaim_older_than_ms: int = 0, + ) -> Optional[WorkItem]: + if not environment_id: + raise ValueError("environment_id is required") + item = self.client.environments.work.poll( + environment_id, + worker_id=worker_id, + block_ms=block_ms, + reclaim_older_than_ms=reclaim_older_than_ms, + timeout=max(5.0, block_ms / 1000.0 + 5.0), + ) + if item is None: + return None + return item + + def ack_work(self, environment_id: str, work_id: str, *, worker_id: str = "") -> None: + if not environment_id: + raise ValueError("environment_id is required") + if not work_id: + raise ValueError("work_id is required") + self.client.environments.work.ack( + environment_id, + work_id, + worker_id=worker_id, + timeout=10.0, + ) + + def heartbeat_work( + self, + environment_id: str, + work_id: str, + *, + expected_last_heartbeat: str, + desired_ttl_seconds: int = 30, + ) -> HeartbeatResponse: + if not environment_id: + raise ValueError("environment_id is required") + if not work_id: + raise ValueError("work_id is required") + response = self.client.environments.work.heartbeat( + environment_id, + work_id, + expected_last_heartbeat=(expected_last_heartbeat or EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT), + desired_ttl_seconds=desired_ttl_seconds, + timeout=max(1.0, min(float(desired_ttl_seconds or 30) / 2, 30.0)), + ) + return response + + def stop_work(self, environment_id: str, work_id: str, *, force: bool = False) -> None: + if not environment_id: + raise ValueError("environment_id is required") + if not work_id: + raise ValueError("work_id is required") + self.client.environments.work.stop( + environment_id, + work_id, + force=force, + timeout=10.0, + ) + + def get_session(self, session_id: str) -> Session: + if not session_id: + raise ValueError("session_id is required") + raw = _model_to_dict(self.client.sessions.retrieve(session_id, timeout=30.0)) + return Session.from_mapping(raw) + + def list_events( + self, + session_id: str, + *, + created_at_gt: str = "", + page: str = "", + limit: int = 100, + order: str = "asc", + types: Optional[List[str]] = None, + ) -> ListEventsResponse: + if not session_id: + raise ValueError("session_id is required") + resp = self.client.sessions.events.list( + session_id, + created_at_gt=created_at_gt if created_at_gt else NOT_GIVEN, + page=page if page else NOT_GIVEN, + limit=limit if limit > 0 else NOT_GIVEN, + order=order if order else NOT_GIVEN, + types=types if types else NOT_GIVEN, + timeout=30.0, + ) + return ListEventsResponse( + events=[_event_from_model(event) for event in getattr(resp, "events", []) or []], + next_page=getattr(resp, "next_page", "") or "", + ) + + def stream_events(self, session_id: str, *, timeout: Optional[float] = 30.0) -> Iterator[Event]: + if not session_id: + raise ValueError("session_id is required") + for frame in self.client.sessions.events.stream(session_id, timeout=timeout): + yield _event_from_model(getattr(frame, "data", None)) + + def send_event(self, session_id: str, event: Event) -> None: + if not session_id: + raise ValueError("session_id is required") + self._request_json( + "POST", + f"/sessions/{_escape(session_id)}/events", + json={"events": [event.to_dict()]}, + timeout=15.0, + max_retries=0, + ) + + def resolve_skill(self, skill: SkillRef) -> SkillRef: + """Enrich a session skill reference with control-plane metadata.""" + skill_id = skill.id_value().strip() + if not skill_id: + raise ValueError("skill id is required") + metadata = self.client.skills.retrieve(skill_id) + name = str(getattr(metadata, "name", "") or "").strip() + if not name: + raise ValueError(f"skill name is empty: {skill_id}") + return replace( + skill, + name=name, + version=skill.version or str(getattr(metadata, "latest_version", "") or "").strip(), + ) + + def open_skill(self, session_id: str, skill: SkillRef) -> SkillContent: + if skill.download_url: + return self._open_external_skill(skill.download_url) + skill_id = skill.id_value() + if not skill_id: + raise ValueError("skill id is required") + if not skill.version: + raise ValueError("skill version is required") + if skill.type.strip().lower() == SKILL_TYPE_SKILL_HUB: + slug = self._lookup_skill_hub_slug(skill_id) + return self._open_external_skill(self._skill_hub_download_url(slug, skill.version)) + response = self.client._client.stream( + "GET", + self._url(f"/skills/{_escape(skill_id)}/versions/{_escape(skill.version)}/content"), + headers=self._headers(), + ) + resp = response.__enter__() + try: + self._raise_for_response(resp) + except BaseException: + response.__exit__(*os.sys.exc_info()) + raise + return SkillContent( + body=_ClosingStream(response, resp), + content_length=int(resp.headers.get("content-length") or -1), + file_name=os.path.basename(str(resp.request.url.path)), + content_type=resp.headers.get("content-type", ""), + ) + + def _lookup_skill_hub_slug(self, skill_id: str) -> str: + request = self._external_request( + SKILL_HUB_BASE_URL, + params={"skillIds": skill_id}, + ) + try: + response = self.client._client.send(request, stream=True, follow_redirects=True) + except (httpx.TimeoutException, httpx.TransportError) as exc: + raise APIError(0, f"lookup skill hub metadata: {exc}", "") from exc + try: + self._raise_for_response(response) + body = bytearray() + for chunk in response.iter_bytes(chunk_size=65536): + body.extend(chunk) + if len(body) > MAX_SKILL_HUB_METADATA_BYTES: + raise APIError(response.status_code, "skill hub metadata response is too large", "") + try: + payload = json.loads(bytes(body)) + except (TypeError, ValueError) as exc: + raise APIError(response.status_code, f"decode skill hub metadata: {exc}", "") from exc + finally: + response.close() + skills = payload.get("Skills") if isinstance(payload, dict) else None + for candidate in skills or []: + if not isinstance(candidate, dict) or str(candidate.get("Id") or "").strip() != skill_id: + continue + slug = str(candidate.get("Slug") or "").strip().strip("/") + if not slug: + raise APIError(500, f"skill hub slug is empty: {skill_id}", "") + return slug + raise APIError(404, f"skill hub skill not found: {skill_id}", "") + + def _skill_hub_download_url(self, slug: str, version: str) -> str: + segments = [] + for segment in slug.strip("/").split("/"): + value = segment.strip() + if not value or value in (".", ".."): + raise ValueError(f"invalid skill hub slug: {slug!r}") + segments.append(quote(value, safe="")) + return str( + httpx.URL( + f"{SKILL_HUB_BASE_URL}/download/{'/'.join(segments)}", + params={"version": version}, + ) + ) + + def _open_external_skill(self, url: str) -> SkillContent: + request = self._external_request(url) + try: + response = self.client._client.send(request, stream=True, follow_redirects=True) + except (httpx.TimeoutException, httpx.TransportError) as exc: + raise APIError(0, str(exc), "") from exc + try: + self._raise_for_response(response) + except BaseException: + response.close() + raise + return SkillContent( + body=_ClosingStream(None, response), + content_length=int(response.headers.get("content-length") or -1), + file_name=os.path.basename(str(response.request.url.path)), + content_type=response.headers.get("content-type", ""), + ) + + def _external_request( + self, + url: str, + *, + params: Optional[Dict[str, Any]] = None, + ) -> httpx.Request: + request = self.client._client.build_request("GET", url, params=params) + for header in SENSITIVE_EXTERNAL_HEADERS: + request.headers.pop(header, None) + return request + + def _request_json( + self, + method: str, + path: str, + *, + params: Optional[Dict[str, Any]] = None, + headers: Optional[Dict[str, str]] = None, + json: Optional[Dict[str, Any]] = None, + timeout: Optional[float] = None, + max_retries: Optional[int] = None, + ) -> Dict[str, Any]: + retry_count = getattr(self.client, "max_retries", 0) if max_retries is None else max_retries + retry_count = max(0, int(retry_count or 0)) + for attempt in range(retry_count + 1): + try: + request_options: Dict[str, Any] = {} + if timeout is not None: + request_options["timeout"] = timeout + resp = self.client._client.request( + method, + self._url(path), + params=params or None, + headers=self._headers(headers), + json=json, + **request_options, + ) + except (httpx.TimeoutException, httpx.TransportError) as exc: + if attempt < retry_count: + self._sleep_retry(attempt) + continue + raise APIError(0, str(exc), "") from exc + if _should_retry(resp.status_code) and attempt < retry_count: + if not resp.is_closed: + with contextlib.suppress(Exception): + resp.read() + resp.close() + self._sleep_retry(attempt) + continue + break + self._raise_for_response(resp) + if not resp.content: + return {} + try: + data = resp.json() + except ValueError as exc: + raise APIError(resp.status_code, str(exc), resp.headers.get(SERVER_REQUEST_HEADER, "")) from exc + return data if isinstance(data, dict) else {} + + def _headers(self, extra: Optional[Dict[str, str]] = None) -> Dict[str, str]: + headers = {"Accept": "application/json", "Content-Type": "application/json"} + headers.update(self.client.auth_headers or {}) + headers.update(extra or {}) + return headers + + def _url(self, path: str) -> str: + return f"{str(self.client._base_url).rstrip('/')}/{path.lstrip('/')}" + + def _raise_for_response(self, resp: httpx.Response) -> None: + if resp.status_code < 400: + return + if not resp.is_closed: + with contextlib.suppress(Exception): + resp.read() + request_id = resp.headers.get(SERVER_REQUEST_HEADER) or resp.headers.get(CLIENT_REQUEST_HEADER, "") + try: + err = self.client._make_status_error_from_response(resp, request_id=request_id) + except Exception as exc: # noqa: BLE001 - preserve status for worker classification. + raise APIError(resp.status_code, str(exc), request_id) from exc + raise APIError(resp.status_code, str(err), request_id) from err + + def _sleep_retry(self, attempt: int) -> None: + delay = min(8.0, 0.5 * (2**attempt)) + time.sleep(delay * random.uniform(0.75, 1.25)) + + +class _ClosingStream: + def __init__(self, manager: Any, response: httpx.Response) -> None: + self._manager = manager + self._response = response + self._iterator = response.iter_bytes(chunk_size=65536) + + def read(self, size: int = -1) -> bytes: + if size is None or size < 0: + return self._response.read() + try: + return next(self._iterator) + except StopIteration: + return b"" + + def iter_bytes(self, chunk_size: int = 65536) -> Iterator[bytes]: + yield from self._response.iter_bytes(chunk_size=chunk_size) + + def close(self) -> None: + if self._manager is None: + self._response.close() + return + self._manager.__exit__(None, None, None) + + +def _escape(value: str) -> str: + return quote(value, safe="") + + +def _should_retry(status_code: int) -> bool: + return status_code in RETRYABLE_STATUS_CODES or status_code >= 500 + + +def _model_to_dict(value: Any) -> Dict[str, Any]: + if value is None: + return {} + if hasattr(value, "model_dump"): + return value.model_dump(mode="json", by_alias=True) + if hasattr(value, "dict"): + return value.dict(by_alias=True) + if isinstance(value, dict): + return value + return {} + + +def _event_from_model(value: Any) -> Event: + if value is None: + return Event(type="") + raw = _model_to_dict(value) + if not raw and isinstance(value, dict): + raw = value + raw_payload = raw.get("raw_payload") + if raw_payload and isinstance(raw_payload, str): + try: + parsed = json.loads(raw_payload) + if isinstance(parsed, dict): + raw = parsed + except ValueError: + pass + return Event.from_mapping(raw) diff --git a/src/arkruntime/selfhosted/envinit.py b/src/arkruntime/selfhosted/envinit.py new file mode 100644 index 0000000..303af05 --- /dev/null +++ b/src/arkruntime/selfhosted/envinit.py @@ -0,0 +1,248 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import logging +import os +import re +import shutil +import stat +import tarfile +import tempfile +import zipfile +from contextlib import suppress +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Optional + +from .types import Session, SkillRef + +DEFAULT_MAX_ARCHIVE_BYTES = 128 << 20 +DEFAULT_MAX_EXTRACTED_BYTES = 512 << 20 +DEFAULT_MAX_ARCHIVE_ENTRIES = 10000 +_SAFE_NAME = re.compile(r"^[A-Za-z0-9._-]+$") + + +@dataclass +class InitializerOptions: + workdir: str + skills_dir: str = "" + max_archive_bytes: int = DEFAULT_MAX_ARCHIVE_BYTES + max_extracted_bytes: int = DEFAULT_MAX_EXTRACTED_BYTES + max_archive_entries: int = DEFAULT_MAX_ARCHIVE_ENTRIES + logger: logging.Logger = logging.getLogger("arkruntime.selfhosted.envinit") + + +class Initializer: + def __init__(self, api: Any, options: InitializerOptions) -> None: + self.api = api + self.options = options + if not self.options.skills_dir: + self.options.skills_dir = str(Path(self.options.workdir) / "skills") + + def setup(self, session: Session) -> None: + if not session: + raise ValueError("session must not be empty") + if not self.options.workdir: + raise ValueError("workdir must not be empty") + Path(self.options.workdir).mkdir(parents=True, exist_ok=True) + Path(self.options.skills_dir).mkdir(parents=True, exist_ok=True) + for skill in session.skill_refs(): + try: + self.install_skill(session.id, skill) + except Exception as exc: # noqa: BLE001 - one broken skill must not fail session setup. + # Follow the Go SDK: skill install errors are observable but do + # not fail the whole session setup. + self.options.logger.warning( + "failed to install skill session_id=%s skill=%s version=%s err=%s", + session.id, + skill.name_value(), + skill.version, + exc, + ) + continue + + 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) + resolver = getattr(self.api, "resolve_skill", None) + if callable(resolver) and skill.id_value().strip(): + skill = resolver(skill) + name = _safe_skill_dir_name(skill) + self.options.logger.info( + "install skill session_id=%s skill=%s version=%s", + session_id, + name, + skill.version, + ) + try: + content = self.api.open_skill(session_id, skill) + except Exception as exc: + raise RuntimeError(f"download skill {name}: {exc}") from exc + if content is None or content.body is None: + raise RuntimeError(f"download skill {name}: empty content") + try: + archive_path = self._copy_archive(name, content.body) + finally: + if hasattr(content.body, "close"): + content.body.close() + tmp = tempfile.mkdtemp(prefix=f".{name}-", dir=self.options.skills_dir) + try: + self._extract_archive(archive_path, tmp) + source = _install_source_dir(Path(tmp)) + target = Path(self.options.skills_dir) / name + backup = _replace_skill_dir(Path(source), target) + if backup is not None: + try: + shutil.rmtree(backup) + except OSError as exc: + self.options.logger.warning( + "remove old skill backup failed session_id=%s skill=%s path=%s err=%s", + session_id, + name, + backup, + exc, + ) + finally: + with suppress(OSError): + os.remove(archive_path) + with suppress(OSError): + if os.path.exists(tmp): + shutil.rmtree(tmp) + + def _copy_archive(self, name: str, body: Any) -> str: + fd, path = tempfile.mkstemp(prefix=f"ark-skill-{name}-") + copied = 0 + try: + with os.fdopen(fd, "wb") as out: + while True: + chunk = body.read(65536) + if not chunk: + break + copied += len(chunk) + if copied > self.options.max_archive_bytes: + raise ValueError(f"skill archive too large: {copied} bytes") + out.write(chunk) + if copied == 0: + raise ValueError("skill archive is empty") + return path + except Exception: + with suppress(OSError): + os.remove(path) + raise + + def _extract_archive(self, archive_path: str, dst: str) -> None: + with open(archive_path, "rb") as f: + magic = f.read(4) + if magic.startswith(b"PK"): + self._extract_zip(archive_path, dst) + return + if magic.startswith(b"\x1f\x8b"): + self._extract_tar_gz(archive_path, dst) + return + raise ValueError("unsupported skill archive format") + + def _extract_zip(self, archive_path: str, dst: str) -> None: + total = 0 + with zipfile.ZipFile(archive_path) as zf: + entries = zf.infolist() + if len(entries) > self.options.max_archive_entries: + raise ValueError(f"skill archive contains too many entries: {len(entries)}") + for info in entries: + target = _safe_join(dst, info.filename) + if info.is_dir(): + Path(target).mkdir(parents=True, exist_ok=True) + continue + entry_type = stat.S_IFMT(info.external_attr >> 16) + if entry_type and not stat.S_ISREG(entry_type): + raise ValueError(f"unsupported zip entry type: {info.filename}") + remaining = self.options.max_extracted_bytes - total + if remaining < 0 or info.file_size > remaining: + raise ValueError( + f"skill extracted content too large: more than {self.options.max_extracted_bytes} bytes" + ) + Path(target).parent.mkdir(parents=True, exist_ok=True) + with zf.open(info) as src, open(target, "wb") as out: + total += _copy_limited(src, out, remaining, self.options.max_extracted_bytes) + + def _extract_tar_gz(self, archive_path: str, dst: str) -> None: + total = 0 + entries = 0 + with tarfile.open(archive_path, "r:gz") as tf: + for member in tf: + entries += 1 + if entries > self.options.max_archive_entries: + raise ValueError(f"skill archive contains too many entries: {entries}") + target = _safe_join(dst, member.name) + if member.isdir(): + Path(target).mkdir(parents=True, exist_ok=True) + continue + if not member.isfile(): + raise ValueError(f"unsupported tar entry type: {member.name}") + remaining = self.options.max_extracted_bytes - total + if remaining < 0 or member.size < 0 or member.size > remaining: + raise ValueError( + f"skill extracted content too large: more than {self.options.max_extracted_bytes} bytes" + ) + src = tf.extractfile(member) + if src is None: + continue + Path(target).parent.mkdir(parents=True, exist_ok=True) + with src, open(target, "wb") as out: + total += _copy_limited(src, out, remaining, self.options.max_extracted_bytes) + + +def _safe_skill_dir_name(skill: SkillRef) -> str: + for candidate in (skill.name, skill.display_name, skill.id_value()): + name = candidate.strip() + if name and name not in (".", "..") and _SAFE_NAME.match(name): + return name + raise ValueError(f"invalid skill name: {skill.name!r}") + + +def _safe_join(root: str, name: str) -> str: + if not name or os.path.isabs(name): + raise ValueError(f"invalid archive path: {name}") + root_path = Path(root).resolve() + target = (root_path / name).resolve() + if target != root_path and root_path not in target.parents: + raise ValueError(f"archive path escapes skill dir: {name}") + return str(target) + + +def _install_source_dir(tmp: Path) -> str: + entries = list(tmp.iterdir()) + if len(entries) == 1 and entries[0].is_dir(): + return str(entries[0]) + return str(tmp) + + +def _replace_skill_dir(source: Path, target: Path) -> Optional[Path]: + if not os.path.lexists(target): + os.replace(source, target) + return None + backup = Path(tempfile.mkdtemp(prefix=f".{target.name}-backup-", dir=target.parent)) + backup.rmdir() + os.replace(target, backup) + try: + os.replace(source, target) + except OSError as exc: + try: + os.replace(backup, target) + except OSError as rollback_exc: + raise OSError(f"replace skill: {exc}; rollback: {rollback_exc}") from exc + raise + return backup + + +def _copy_limited(src: Any, out: Any, limit: int, configured_limit: int) -> int: + written = 0 + while True: + chunk = src.read(min(65536, limit - written + 1)) + if not chunk: + return written + written += len(chunk) + if written > limit: + raise ValueError(f"skill extracted content too large: more than {configured_limit} bytes") + out.write(chunk) diff --git a/src/arkruntime/selfhosted/session_tool_runner.py b/src/arkruntime/selfhosted/session_tool_runner.py new file mode 100644 index 0000000..948dec4 --- /dev/null +++ b/src/arkruntime/selfhosted/session_tool_runner.py @@ -0,0 +1,552 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import logging +import queue +import random +import threading +import time +from dataclasses import dataclass, field, replace +from typing import Any, Dict, Iterable, List, Optional + +from .tool_result_store import FileToolResultStore +from .tools import Tool, ToolContext, ToolResult, ToolSet, error_result +from .types import ( + CONFIRMATION_ALLOW, + CONFIRMATION_DENY, + DEFAULT_MAX_IDLE_SECONDS, + EVENT_LIST_ORDER_ASC, + EVENT_TYPE_AGENT_CUSTOM_TOOL_USE, + EVENT_TYPE_AGENT_TOOL_USE, + EVENT_TYPE_SESSION_DELETED, + EVENT_TYPE_SESSION_STATUS_IDLE, + EVENT_TYPE_SESSION_STATUS_TERMINATED, + EVENT_TYPE_USER_CUSTOM_TOOL_RESULT, + EVENT_TYPE_USER_TOOL_CONFIRMATION, + EVENT_TYPE_USER_TOOL_RESULT, + PERMISSION_ALLOW, + PERMISSION_DENY, + SESSION_STOP_REASON_END_TURN, + ContentBlock, + Event, + EventStreamUnsupported, + IdleTimeout, + SessionTerminated, + ToolCallResult, + is_fatal_4xx, + new_user_custom_tool_result_event, + new_user_tool_result_event, + tool_confirmation_call_id, + tool_result_call_id, + tool_use_call_id, +) + +STREAM_BACKOFF_START = 0.5 +STREAM_BACKOFF_CAP = 10.0 +STREAM_HEALTHY_AFTER = 30.0 +SEND_RETRIES = 3 +STREAM_QUEUE_SIZE = 256 + + +@dataclass +class SessionToolRunnerOptions: + work_id: str = "" + tools: ToolSet = None # type: ignore[assignment] + tool_context: ToolContext = None # type: ignore[assignment] + custom_tools: Dict[str, Tool] = field(default_factory=dict) + result_store: Optional[FileToolResultStore] = None + event_page: str = "" + event_poll_interval_seconds: float = 0.5 + event_limit: int = 100 + max_idle_seconds: Optional[float] = DEFAULT_MAX_IDLE_SECONDS + tool_timeout_seconds: Optional[float] = None + send_timeout_seconds: float = 15.0 + prefer_stream: bool = True + stream_timeout_seconds: float = 30.0 + stop_event: Any = None + logger: logging.Logger = logging.getLogger("arkruntime.selfhosted.session_tool_runner") + on_tool_error: Any = None + + +class SessionToolRunner: + def __init__(self, api: Any, session_id: str, options: SessionToolRunnerOptions) -> None: + if not session_id: + raise ValueError("session id must not be empty") + if api is None: + raise ValueError("session tool runner api must not be empty") + if options is None: + raise ValueError("session tool runner options must not be empty") + if options.tools is None: + raise ValueError("session tool runner tools must not be empty") + if options.tool_context is None: + raise ValueError("session tool runner tool context must not be empty") + self.api = api + self.session_id = session_id + self.options = options + self._stop = threading.Event() + self._results: List[ToolCallResult] = [] + self._state = _RunnerState(self) + + @property + def results(self) -> List[ToolCallResult]: + return self._results + + def close(self) -> None: + self._stop.set() + + def _is_stopped(self) -> bool: + return self._stop.is_set() or bool(self.options.stop_event and self.options.stop_event.is_set()) + + def run(self) -> List[ToolCallResult]: + if self.options.result_store is not None: + pending, processed = self.options.result_store.recover() + self._state.pending_results.update(pending) + self._state.processed.update(processed) + self._state.answered.update(processed) + if self.options.prefer_stream and hasattr(self.api, "stream_events"): + try: + self._consume_stream_loop() + return self._results + except EventStreamUnsupported: + pass + except (IdleTimeout, SessionTerminated): + raise + except Exception: + if self._is_stopped(): + return self._results + raise + self._consume_list() + return self._results + + def _consume_stream_loop(self) -> None: + backoff = STREAM_BACKOFF_START + while not self._is_stopped(): + event_queue: "queue.Queue[object]" = queue.Queue(maxsize=STREAM_QUEUE_SIZE) + opened_at = time.monotonic() + pump = threading.Thread(target=self._pump_stream, args=(event_queue,), daemon=True) + pump.start() + self._state.reconcile() + while not self._is_stopped() and (pump.is_alive() or not event_queue.empty()): + self._state.flush_results() + self._raise_if_idle_expired() + try: + item = event_queue.get(timeout=self._state.next_wait_seconds(0.5)) + except queue.Empty: + continue + if isinstance(item, BaseException): + if time.monotonic() - opened_at > STREAM_HEALTHY_AFTER: + backoff = STREAM_BACKOFF_START + if isinstance(item, EventStreamUnsupported): + raise item + if is_fatal_4xx(item): + raise item + break + self._state.handle_stream_event(item) # type: ignore[arg-type] + self._sleep_or_idle(_jitter(backoff)) + backoff = min(backoff * 2, STREAM_BACKOFF_CAP) + + def _pump_stream(self, event_queue: "queue.Queue[object]") -> None: + stream = None + try: + stream = self.api.stream_events(self.session_id, timeout=self.options.stream_timeout_seconds) + for event in stream: + if self._is_stopped() or not self._put_stream_item(event_queue, event): + return + except Exception as exc: # noqa: BLE001 - forwarded to the owner loop. + self._put_stream_item(event_queue, exc) + finally: + close = getattr(stream, "close", None) + if callable(close): + try: + close() + except Exception: # noqa: BLE001 - stream may already be closing in its pump thread. + pass + + def _put_stream_item(self, event_queue: "queue.Queue[object]", item: object) -> bool: + while not self._is_stopped(): + try: + event_queue.put(item, timeout=0.1) + return True + except queue.Full: + continue + return False + + def _consume_list(self) -> None: + while not self._is_stopped(): + self._state.flush_results() + self._state.reconcile(reconcile=False) + self._raise_if_idle_expired() + self._sleep_or_idle(self.options.event_poll_interval_seconds) + + def _sleep_or_idle(self, seconds: float) -> None: + deadline = time.monotonic() + max(seconds, 0) + while not self._is_stopped(): + self._raise_if_idle_expired() + remaining = deadline - time.monotonic() + if remaining <= 0: + return + wait_for = self._state.next_wait_seconds(min(remaining, 0.5)) + if self.options.stop_event is not None: + self.options.stop_event.wait(wait_for) + else: + self._stop.wait(wait_for) + + def _raise_if_idle_expired(self) -> None: + if self._state.idle_expired(): + raise IdleTimeout("session idle after end_turn") + + +class _RunnerState: + def __init__(self, runner: SessionToolRunner) -> None: + self.runner = runner + self.page = runner.options.event_page + self.processed: Dict[str, bool] = {} + self.seen: Dict[str, bool] = {} + self.answered: Dict[str, bool] = {} + self.pending_results: Dict[str, Event] = {} + self.pending_ask: Dict[str, Event] = {} + self.confirmations: Dict[str, Event] = {} + self.external_tools: Dict[str, Event] = {} + self.idle_armed_at = 0.0 + self.idle_arm_pending = False + + def reconcile(self, *, reconcile: bool = True) -> None: + backoff = STREAM_BACKOFF_START + while not self.runner._is_stopped(): + try: + self._reconcile_once(reconcile=reconcile) + return + except Exception as exc: # noqa: BLE001 - runner owns retry classification. + if is_fatal_4xx(exc): + raise + self.runner.options.logger.warning( + "reconcile list events failed err=%s sleep=%.3fs", + exc, + backoff, + ) + self.runner._sleep_or_idle(_jitter(backoff)) + backoff = min(backoff * 2, STREAM_BACKOFF_CAP) + + def _reconcile_once(self, *, reconcile: bool) -> None: + events: List[Event] = [] + page = "" + while not self.runner._is_stopped(): + resp = self.runner.api.list_events( + self.runner.session_id, + page=page, + limit=min(max(self.runner.options.event_limit, 1), 1000), + order=EVENT_LIST_ORDER_ASC, + ) + if resp is None: + break + events.extend(resp.events) + if not resp.next_page: + break + page = resp.next_page + self.process_listed_events(events, reconcile=reconcile) + + def process_listed_events(self, events: Iterable[Event], reconcile: bool = False) -> None: + pending: List[Event] = [] + pending_ids: Dict[str, bool] = {} + touched_idle = False + last_was_end_turn = False + for event in events: + seen_now = self.mark_event_seen(event) + if not reconcile and not seen_now: + continue + if seen_now and event.type != EVENT_TYPE_USER_TOOL_CONFIRMATION: + touched_idle = True + last_was_end_turn = ( + event.type == EVENT_TYPE_SESSION_STATUS_IDLE + and event.stop_reason_type() == SESSION_STOP_REASON_END_TURN + ) + if event.type == EVENT_TYPE_USER_TOOL_CONFIRMATION: + self.record_confirmation(event) + elif event.type in (EVENT_TYPE_USER_TOOL_RESULT, EVENT_TYPE_USER_CUSTOM_TOOL_RESULT): + self.mark_answered(tool_result_call_id(event)) + elif event.type in (EVENT_TYPE_AGENT_TOOL_USE, EVENT_TYPE_AGENT_CUSTOM_TOOL_USE): + call_id = tool_use_call_id(event) + if call_id and not pending_ids.get(call_id): + pending.append(event) + pending_ids[call_id] = True + elif event.type in (EVENT_TYPE_SESSION_STATUS_TERMINATED, EVENT_TYPE_SESSION_DELETED): + raise SessionTerminated("session terminated") + if touched_idle: + self.disarm_idle() + for event in pending: + if self.is_answered(tool_use_call_id(event)): + continue + self.handle_tool_use(event, event.type == EVENT_TYPE_AGENT_CUSTOM_TOOL_USE) + self.release_confirmed_tool_uses() + if touched_idle and last_was_end_turn: + if self.has_unblocked_outstanding_tool(pending): + self.disarm_idle() + else: + self.arm_idle() + + def note_idle_event(self, event: Event) -> None: + if event.type == EVENT_TYPE_USER_TOOL_CONFIRMATION: + return + if event.type == EVENT_TYPE_SESSION_STATUS_IDLE and event.stop_reason_type() == SESSION_STOP_REASON_END_TURN: + self.arm_idle() + return + self.disarm_idle() + + def handle_stream_event(self, event: Event) -> None: + if not self.mark_event_seen(event): + return + self.note_idle_event(event) + self.handle_event(event) + + def handle_event(self, event: Event) -> None: + if event.type == EVENT_TYPE_USER_TOOL_CONFIRMATION: + self.record_confirmation(event) + self.release_confirmed_tool_uses() + elif event.type in (EVENT_TYPE_USER_TOOL_RESULT, EVENT_TYPE_USER_CUSTOM_TOOL_RESULT): + self.mark_answered(tool_result_call_id(event)) + elif event.type in (EVENT_TYPE_AGENT_TOOL_USE, EVENT_TYPE_AGENT_CUSTOM_TOOL_USE): + self.handle_tool_use(event, event.type == EVENT_TYPE_AGENT_CUSTOM_TOOL_USE) + elif event.type in (EVENT_TYPE_SESSION_STATUS_TERMINATED, EVENT_TYPE_SESSION_DELETED): + raise SessionTerminated("session terminated") + + def mark_event_seen(self, event: Event) -> bool: + key = event.id or tool_use_call_id(event) + if not key: + return True + if self.seen.get(key): + return False + self.seen[key] = True + return True + + def mark_answered(self, call_id: str) -> None: + if not call_id: + return + self.answered[call_id] = True + self.processed[call_id] = True + self.pending_results.pop(call_id, None) + self.pending_ask.pop(call_id, None) + self.external_tools.pop(call_id, None) + self.maybe_arm_pending_idle() + + def is_answered(self, call_id: str) -> bool: + return bool(call_id and self.answered.get(call_id)) + + def record_confirmation(self, event: Event) -> None: + call_id = tool_confirmation_call_id(event) + if call_id and not self.is_answered(call_id): + self.confirmations[call_id] = event + + def release_confirmed_tool_uses(self) -> None: + ready = [event for call_id, event in self.pending_ask.items() if call_id in self.confirmations] + for event in ready: + self.pending_ask.pop(tool_use_call_id(event), None) + self.handle_tool_use(event, event.type == EVENT_TYPE_AGENT_CUSTOM_TOOL_USE) + + def has_unblocked_outstanding_tool(self, pending: Iterable[Event]) -> bool: + for event in pending: + call_id = tool_use_call_id(event) + if not call_id or self.is_answered(call_id): + continue + if call_id in self.pending_ask or call_id in self.pending_results: + continue + return True + return False + + def handle_tool_use(self, event: Event, custom: bool) -> None: + call_id = tool_use_call_id(event) + if not call_id or self.is_answered(call_id): + return + pending = self.pending_results.get(call_id) + if pending is not None: + self.send_result(call_id, event, custom, "", pending) + return + if not self.owns_tool(event, custom): + self.external_tools[call_id] = event + self.maybe_arm_pending_idle() + self.runner._results.append(ToolCallResult(call_id, event.name, custom, posted=False, event=event)) + return + confirmation, allowed = self.permission_allows(event, custom, call_id) + if not allowed: + self.runner._results.append( + ToolCallResult(call_id, event.name, custom, confirmation=confirmation, posted=False, event=event) + ) + return + if self.runner.options.result_store is not None: + decision = self.runner.options.result_store.begin(call_id, event) + if decision.sent: + self.mark_answered(call_id) + return + if decision.result is not None: + self.pending_results[call_id] = decision.result + self.send_result(call_id, event, custom, "", decision.result) + return + result = self.execute_tool(event, custom) + self.post_result(event, custom, call_id, result, confirmation) + + def owns_tool(self, event: Event, custom: bool) -> bool: + if custom: + return event.name in self.runner.options.custom_tools + return self.runner.options.tools.has(event.name) + + def permission_allows(self, event: Event, custom: bool, call_id: str) -> tuple: + if custom: + return "", True + permission = event.evaluated_permission + if permission in ("", PERMISSION_ALLOW): + return "", True + if permission == "ask": + confirmation = self.confirmations.get(call_id) + if confirmation is None: + self.pending_ask[call_id] = event + return "", False + if confirmation.result == CONFIRMATION_ALLOW: + return CONFIRMATION_ALLOW, True + self.mark_answered(call_id) + return CONFIRMATION_DENY, False + if permission == PERMISSION_DENY: + self.mark_answered(call_id) + return CONFIRMATION_DENY, False + self.pending_ask[call_id] = event + return "", False + + def execute_tool(self, event: Event, custom: bool) -> ToolResult: + context = replace(self.runner.options.tool_context) + if self.runner.options.tool_timeout_seconds is not None: + context.tool_timeout_seconds = self.runner.options.tool_timeout_seconds + if custom: + tool = self.runner.options.custom_tools[event.name] + try: + return tool.execute(event.input, context) + except Exception as exc: # noqa: BLE001 - custom tool failures become tool results. + return error_result(str(exc)) + return self.runner.options.tools.execute(event.name, event.input, context) + + def post_result(self, event: Event, custom: bool, call_id: str, result: ToolResult, confirmation: str) -> None: + blocks = list(result.content) + if custom: + out = new_user_custom_tool_result_event(call_id, blocks, result.is_error, event.session_thread_id) + else: + out = new_user_tool_result_event(call_id, blocks, result.is_error, event.session_thread_id) + if self.runner.options.result_store is not None: + try: + self.runner.options.result_store.save_result(call_id, out) + except Exception as exc: # noqa: BLE001 - delivery must continue after local ledger failure. + self.runner.options.logger.warning( + "persist tool result failed tool_use_id=%s err=%s", + call_id, + exc, + ) + self.pending_results[call_id] = out + self.send_result(call_id, event, custom, confirmation, out) + + def send_result(self, call_id: str, event: Event, custom: bool, confirmation: str, out: Event) -> None: + posted = self.retry_send_event(out, call_id) + if posted: + self.mark_answered(call_id) + if self.runner.options.result_store is not None: + try: + self.runner.options.result_store.mark_sent(call_id) + except Exception as exc: # noqa: BLE001 - MA already accepted the result. + self.runner.options.logger.warning( + "mark tool result sent failed tool_use_id=%s event_id=%s err=%s", + call_id, + out.id, + exc, + ) + elif self.runner.options.result_store is not None: + self.pending_results[call_id] = out + self.runner._results.append( + ToolCallResult( + tool_use_id=call_id, + name=event.name, + custom=custom, + confirmation=confirmation, + posted=posted, + event=event, + result=out, + ) + ) + + def retry_send_event(self, event: Event, call_id: str) -> bool: + last_exc: Optional[BaseException] = None + for attempt in range(SEND_RETRIES): + try: + self.runner.api.send_event(self.runner.session_id, event) + return True + except Exception as exc: # noqa: BLE001 - mirrors Go retry classification. + last_exc = exc + if is_fatal_4xx(exc) or self.runner._is_stopped(): + break + if attempt < SEND_RETRIES - 1: + self.runner._sleep_or_idle(attempt + 1) + if self.runner.options.on_tool_error and last_exc is not None: + self.runner.options.on_tool_error(event, last_exc) + return False + + def flush_results(self) -> None: + for call_id, event in list(self.pending_results.items()): + if self.retry_send_event(event, call_id): + self.mark_answered(call_id) + if self.runner.options.result_store is not None: + try: + self.runner.options.result_store.mark_sent(call_id) + except Exception as exc: # noqa: BLE001 - MA already accepted the result. + self.runner.options.logger.warning( + "mark pending tool result sent failed tool_use_id=%s event_id=%s err=%s", + call_id, + event.id, + exc, + ) + self.maybe_arm_pending_idle() + + def arm_idle(self) -> None: + if not self.max_idle_seconds(): + return + if self.has_idle_blockers(): + self.idle_arm_pending = True + self.idle_armed_at = 0.0 + return + self.idle_arm_pending = False + self.idle_armed_at = time.monotonic() + + def disarm_idle(self) -> None: + self.idle_arm_pending = False + self.idle_armed_at = 0.0 + + def maybe_arm_pending_idle(self) -> None: + if self.idle_arm_pending and not self.has_idle_blockers(): + self.idle_arm_pending = False + self.idle_armed_at = time.monotonic() + + def has_idle_blockers(self) -> bool: + return bool(self.pending_ask or self.pending_results or self.external_tools) + + def idle_expired(self) -> bool: + max_idle = self.max_idle_seconds() + return bool(max_idle and self.idle_armed_at and time.monotonic() - self.idle_armed_at >= max_idle) + + def max_idle_seconds(self) -> float: + if self.runner.options.max_idle_seconds is None: + return 0.0 + return self.runner.options.max_idle_seconds + + def next_wait_seconds(self, fallback: float) -> float: + if not self.idle_armed_at: + return fallback + remaining = self.max_idle_seconds() - (time.monotonic() - self.idle_armed_at) + if remaining <= 0: + return 0.001 + return max(0.001, min(fallback, remaining)) + + +def _jitter(seconds: float) -> float: + if seconds <= 0: + return 0 + half = seconds / 2 + return half + random.random() * half + + +def result_content_blocks(result: ToolResult) -> List[ContentBlock]: + return list(result.content) diff --git a/src/arkruntime/selfhosted/tool_result_store.py b/src/arkruntime/selfhosted/tool_result_store.py new file mode 100644 index 0000000..0d11b00 --- /dev/null +++ b/src/arkruntime/selfhosted/tool_result_store.py @@ -0,0 +1,158 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import hashlib +import json +import os +import tempfile +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Tuple + +from .types import ContentBlock, Event, new_user_custom_tool_result_event, new_user_tool_result_event, utc_now_iso + +STATE_STARTED = "started" +STATE_RESULT = "result" +STATE_SENT = "sent" + + +@dataclass +class ToolCallStoreDecision: + sent: bool = False + result: Event = None # type: ignore[assignment] + + +class FileToolResultStore: + """File-backed ledger that avoids re-running side-effectful tool calls.""" + + def __init__(self, workdir: str) -> None: + if not workdir: + raise ValueError("workdir must not be empty") + self.dir = Path(workdir) / ".ma_self_host_worker" / "tool_ledger" + self.dir.mkdir(parents=True, exist_ok=True, mode=0o700) + + def recover(self) -> Tuple[Dict[str, Event], Dict[str, bool]]: + pending: Dict[str, Event] = {} + processed: Dict[str, bool] = {} + for path in self.dir.glob(".tool-result-*.tmp"): + try: + path.unlink() + except FileNotFoundError: + pass + for path in self.dir.glob("*.json"): + record = self._read_path(path) + call_id = str(record.get("call_id") or "") + state = str(record.get("state") or "") + if not call_id: + raise ValueError(f"tool result record {path} missing call_id") + if state == STATE_SENT: + processed[call_id] = True + elif state == STATE_RESULT: + pending[call_id] = Event.from_mapping(record.get("result") or {}) + elif state == STATE_STARTED: + event = Event.from_mapping(record.get("event") or {}) + result = _unknown_tool_execution_result(call_id, event) + record["result"] = result.to_dict() + record["state"] = STATE_RESULT + self._write_record(record) + pending[call_id] = result + else: + raise ValueError(f"unknown tool result state {state!r} for call {call_id}") + return pending, processed + + def begin(self, call_id: str, event: Event) -> ToolCallStoreDecision: + if not call_id: + raise ValueError("call id must not be empty") + try: + record = self._read(call_id) + except FileNotFoundError: + self._write_record({"call_id": call_id, "state": STATE_STARTED, "event": event.to_dict()}) + return ToolCallStoreDecision() + state = str(record.get("state") or "") + if state == STATE_SENT: + return ToolCallStoreDecision(sent=True) + if state == STATE_RESULT: + return ToolCallStoreDecision(result=Event.from_mapping(record.get("result") or {})) + if state == STATE_STARTED: + result = _unknown_tool_execution_result(call_id, Event.from_mapping(record.get("event") or {})) + record["result"] = result.to_dict() + record["state"] = STATE_RESULT + self._write_record(record) + return ToolCallStoreDecision(result=result) + raise ValueError(f"unknown tool result state {state!r} for call {call_id}") + + def save_result(self, call_id: str, result: Event) -> None: + record = self._read(call_id) + record["state"] = STATE_RESULT + record["result"] = result.to_dict() + self._write_record(record) + + def mark_sent(self, call_id: str) -> None: + record = self._read(call_id) + record["state"] = STATE_SENT + self._write_record(record) + + def _read(self, call_id: str) -> dict: + return self._read_path(self._path(call_id)) + + def _read_path(self, path: Path) -> dict: + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + + def _write_record(self, record: dict) -> None: + call_id = str(record.get("call_id") or "") + if not call_id: + raise ValueError("call id must not be empty") + record["updated_at"] = utc_now_iso() + target = self._path(call_id) + fd, tmp = tempfile.mkstemp(prefix=".tool-result-", suffix=".tmp", dir=self.dir) + try: + with os.fdopen(fd, "w", encoding="utf-8") as output: + fd = -1 + json.dump(record, output, ensure_ascii=False, indent=2) + output.flush() + os.fsync(output.fileno()) + os.replace(tmp, target) + _sync_directory(self.dir) + finally: + if fd >= 0: + os.close(fd) + try: + os.remove(tmp) + except FileNotFoundError: + pass + + def _path(self, call_id: str) -> Path: + digest = hashlib.sha256(call_id.encode("utf-8")).hexdigest() + return self.dir / f"{digest}.json" + + +def _unknown_tool_execution_result(call_id: str, event: Event) -> Event: + content = [ + ContentBlock( + type="text", + text=( + "tool execution state is unknown after worker restart; refusing to re-execute this tool_use " + "to avoid duplicate side effects" + ), + ) + ] + if event.type == "agent.custom_tool_use": + return new_user_custom_tool_result_event(call_id, content, True, event.session_thread_id) + return new_user_tool_result_event(call_id, content, True, event.session_thread_id) + + +def _sync_directory(path: Path) -> None: + try: + fd = os.open(str(path), os.O_RDONLY) + except OSError: + return + try: + os.fsync(fd) + except OSError: + # Some non-POSIX filesystems do not support syncing directories. + pass + finally: + os.close(fd) diff --git a/src/arkruntime/selfhosted/tools.py b/src/arkruntime/selfhosted/tools.py new file mode 100644 index 0000000..ce64b1d --- /dev/null +++ b/src/arkruntime/selfhosted/tools.py @@ -0,0 +1,400 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import glob as globlib +import json +import os +import re +import signal +import stat +import subprocess +import tempfile +import threading +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable, Dict, Iterable, Mapping, Optional + +from .types import DEFAULT_TOOL_TIMEOUT_SECONDS, ContentBlock + +MAX_OUTPUT_BYTES = 100000 +MAX_SEARCH_MATCHES = 1000 +_SENSITIVE_ENV_PREFIXES = ( + "AIME_", + "ARK_", + "MA_", + "X_CODE_", + "ANTHROPIC_", + "OPENAI_", + "AWS_", + "AZURE_", + "GOOGLE_", +) +_SENSITIVE_ENV_NAMES = { + "VOLC_ACCESSKEY", + "VOLC_SECRETKEY", + "BYTEPLUS_ACCESSKEY", + "BYTEPLUS_SECRETKEY", + "TOKEN", + "SECRET", + "PASSWORD", + "PASSWD", + "PRIVATE_KEY", + "API_KEY", + "ACCESS_KEY", + "SECRET_KEY", + "JWT", + "PAT", +} +_SENSITIVE_ENV_SUFFIXES = ( + "_TOKEN", + "_SECRET", + "_PASSWORD", + "_PASSWD", + "_PRIVATE_KEY", + "_API_KEY", + "_ACCESS_KEY", + "_SECRET_KEY", + "_JWT", + "_PAT", +) + + +@dataclass +class ToolResult: + content: Iterable[ContentBlock] + is_error: bool = False + + +class Tool: + name: str + + def execute(self, tool_input: Any, context: "ToolContext") -> ToolResult: + raise NotImplementedError + + +@dataclass +class ToolContext: + workdir: str + env: Optional[Dict[str, str]] = None + unrestricted_paths: bool = False + tool_timeout_seconds: float = DEFAULT_TOOL_TIMEOUT_SECONDS + cancel_event: Any = None + + +class FunctionTool(Tool): + def __init__(self, name: str, fn: Callable[[Any, ToolContext], ToolResult]) -> None: + self.name = name + self._fn = fn + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + return self._fn(tool_input, context) + + +class ToolSet: + def __init__(self, tools: Optional[Iterable[Tool]] = None) -> None: + self._tools: Dict[str, Tool] = {} + for tool in tools or []: + self.add(tool) + + def add(self, tool: Tool) -> None: + if not tool.name: + raise ValueError("tool name must not be empty") + self._tools[tool.name] = tool + + def has(self, name: str) -> bool: + return name in self._tools + + def execute(self, name: str, tool_input: Any, context: ToolContext) -> ToolResult: + tool = self._tools.get(name) + if tool is None: + return error_result(f"tool {name!r} is not registered") + try: + return tool.execute(tool_input, context) + except Exception as exc: # noqa: BLE001 - tool errors must be reported, not raised. + return error_result(str(exc)) + + +class BashTool(Tool): + name = "bash" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + command = str(args.get("command") or args.get("cmd") or "") + if not command: + return error_result("bash command is required") + env = _scrubbed_env(context.env) + proc = subprocess.Popen( # noqa: S602 - this is the explicit bash tool contract. + command, + shell=True, + cwd=context.workdir, + env=env, + text=False, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + output = _BoundedOutput(MAX_OUTPUT_BYTES) + reader = threading.Thread(target=_drain_output, args=(proc.stdout, output), daemon=True) + reader.start() + deadline = time.monotonic() + max(context.tool_timeout_seconds, 0) + failure = "" + while proc.poll() is None: + if _is_cancelled(context): + failure = "tool execution canceled" + _kill_process(proc) + break + if context.tool_timeout_seconds > 0 and time.monotonic() >= deadline: + failure = f"tool execution timed out after {context.tool_timeout_seconds:g}s" + _kill_process(proc) + break + time.sleep(0.05) + proc.wait() + reader.join(timeout=1) + text = output.text() + if failure: + return error_result(f"{failure}\n{text}".rstrip()) + if proc.returncode != 0: + text = f"exit code {proc.returncode}\n{text}" if text else f"exit code {proc.returncode}" + return text_result(text, is_error=proc.returncode != 0) + + +class ReadFileTool(Tool): + name = "read" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + path = _safe_path(context, str(args.get("path") or args.get("file") or "")) + limit = min(max(int(args.get("limit") or 20000), 0), MAX_OUTPUT_BYTES) + offset = int(args.get("offset") or 0) + with open(path, "r", encoding="utf-8", errors="replace") as f: + if offset > 0: + f.seek(offset) + return text_result(f.read(limit)) + + +class WriteFileTool(Tool): + name = "write" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + raw_path = str(args.get("path") or args.get("file") or "") + path = _safe_path(context, raw_path) + content = str(args.get("content") or "") + Path(path).parent.mkdir(parents=True, exist_ok=True) + verified = _safe_path(context, raw_path) + if verified != path: + raise ValueError("path resolution changed while writing") + _write_file_atomically(path, content.encode("utf-8"), 0o600) + return text_result(f"wrote {len(content)} bytes") + + +class EditFileTool(Tool): + name = "edit" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + raw_path = str(args.get("path") or args.get("file") or "") + path = _safe_path(context, raw_path) + old = str(args.get("old_string") or args.get("old") or "") + new = str(args.get("new_string") or args.get("new") or "") + if not old: + return error_result("old_string is required") + data = Path(path).read_text(encoding="utf-8") + if old not in data: + return error_result("old_string was not found") + mode = stat.S_IMODE(os.stat(path).st_mode) + verified = _safe_path(context, raw_path) + if verified != path: + raise ValueError("path resolution changed while editing") + _write_file_atomically(path, data.replace(old, new, 1).encode("utf-8"), mode) + return text_result("edited") + + +class GlobTool(Tool): + name = "glob" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + pattern = str(args.get("pattern") or "") + if not pattern: + return error_result("pattern is required") + root = Path(context.workdir).resolve() + matches = [] + for match in globlib.glob(str(root / pattern), recursive=True): + if _is_cancelled(context): + return error_result("tool execution canceled") + try: + resolved = Path(match).resolve() + if context.unrestricted_paths or resolved == root or root in resolved.parents: + matches.append(str(resolved.relative_to(root)) if root in resolved.parents else str(resolved)) + if len(matches) >= MAX_SEARCH_MATCHES: + break + except (OSError, ValueError): + continue + return text_result(_truncate("\n".join(sorted(matches)))) + + +class GrepTool(Tool): + name = "grep" + + def execute(self, tool_input: Any, context: ToolContext) -> ToolResult: + args = _as_mapping(tool_input) + pattern = str(args.get("pattern") or args.get("query") or "") + target = str(args.get("path") or ".") + if not pattern: + return error_result("pattern is required") + root = _safe_path(context, target) + rx = re.compile(pattern) + lines = [] + paths: Iterable[str] = [root] + if os.path.isdir(root): + paths = (os.path.join(base, name) for base, _, names in os.walk(root) for name in names) + for path in paths: + if _is_cancelled(context): + return error_result("tool execution canceled") + try: + verified = _safe_path(context, path) + with open(verified, "r", encoding="utf-8", errors="replace") as f: + for number, line in enumerate(f, 1): + if _is_cancelled(context): + return error_result("tool execution canceled") + if rx.search(line): + rel = os.path.relpath(path, context.workdir) + lines.append(f"{rel}:{number}:{line.rstrip()}") + if len(lines) >= 200 or sum(len(value) + 1 for value in lines) >= MAX_OUTPUT_BYTES: + return text_result(_truncate("\n".join(lines))) + except (OSError, ValueError): + continue + return text_result("\n".join(lines)) + + +def default_toolset() -> ToolSet: + return ToolSet([BashTool(), ReadFileTool(), WriteFileTool(), EditFileTool(), GlobTool(), GrepTool()]) + + +def text_result(text: str, *, is_error: bool = False) -> ToolResult: + return ToolResult([ContentBlock(type="text", text=text)], is_error=is_error) + + +def error_result(text: str) -> ToolResult: + return text_result(text, is_error=True) + + +def _as_mapping(value: Any) -> Mapping[str, Any]: + if isinstance(value, Mapping): + return value + if isinstance(value, (bytes, bytearray)): + value = value.decode("utf-8", errors="replace") + if isinstance(value, str): + try: + parsed = json.loads(value) + except ValueError: + return {"command": value} + if isinstance(parsed, Mapping): + return parsed + return {} + + +def _safe_path(context: ToolContext, path: str) -> str: + if not path: + raise ValueError("path is required") + root = Path(context.workdir).resolve() + target = Path(path) + if not target.is_absolute(): + target = root / target + resolved = target.resolve() + if context.unrestricted_paths: + return str(resolved) + if resolved != root and root not in resolved.parents: + raise ValueError(f"path escapes workdir: {path}") + return str(resolved) + + +def _truncate(text: str, limit: int = 100000) -> str: + if len(text) <= limit: + return text + return text[:limit] + "\n... truncated ..." + + +def _is_cancelled(context: ToolContext) -> bool: + return bool(context.cancel_event and context.cancel_event.is_set()) + + +def _is_sensitive_env_key(key: str) -> bool: + upper = key.strip().upper() + return ( + upper.startswith(_SENSITIVE_ENV_PREFIXES) + or upper in _SENSITIVE_ENV_NAMES + or upper.endswith(_SENSITIVE_ENV_SUFFIXES) + ) + + +def _scrubbed_env(extra: Optional[Mapping[str, str]]) -> Dict[str, str]: + source = dict(os.environ) if extra is None else dict(extra) + return {key: value for key, value in source.items() if not _is_sensitive_env_key(key)} + + +def _write_file_atomically(path: str, data: bytes, mode: int) -> None: + parent = str(Path(path).parent) + fd, tmp = tempfile.mkstemp(prefix=".ark-write-", dir=parent) + try: + os.fchmod(fd, mode) + with os.fdopen(fd, "wb") as output: + fd = -1 + output.write(data) + os.replace(tmp, path) + finally: + if fd >= 0: + os.close(fd) + try: + os.remove(tmp) + except FileNotFoundError: + pass + + +class _BoundedOutput: + _MARKER = b"\n... truncated ..." + + def __init__(self, limit: int) -> None: + self._limit = limit + self._data = bytearray() + self._truncated = False + self._lock = threading.Lock() + + def append(self, chunk: bytes) -> None: + with self._lock: + remaining = max(0, self._limit - len(self._MARKER) - len(self._data)) + if remaining > 0: + self._data.extend(chunk[:remaining]) + if len(chunk) > remaining: + self._truncated = True + + def text(self) -> str: + with self._lock: + text = bytes(self._data).decode("utf-8", errors="replace") + if self._truncated: + text += self._MARKER.decode("ascii") + return text + + +def _drain_output(stream: Any, output: _BoundedOutput) -> None: + if stream is None: + return + try: + while True: + chunk = stream.read(65536) + if not chunk: + return + output.append(chunk) + finally: + stream.close() + + +def _kill_process(proc: subprocess.Popen) -> None: + try: + os.killpg(proc.pid, signal.SIGKILL) + except (OSError, AttributeError): + proc.kill() diff --git a/src/arkruntime/selfhosted/types.py b/src/arkruntime/selfhosted/types.py new file mode 100644 index 0000000..4a7c198 --- /dev/null +++ b/src/arkruntime/selfhosted/types.py @@ -0,0 +1,418 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import secrets +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, Dict, Iterable, List, Mapping, Optional + +from arkruntime.types.environment import ( + HeartbeatWorkResponse, + VolcTag, + WorkState, +) +from arkruntime.types.environment import ( + WorkData as WorkData, +) +from arkruntime.types.environment import ( + WorkItem as WorkItem, +) + +EVENT_TYPE_AGENT_TOOL_USE = "agent.tool_use" +EVENT_TYPE_AGENT_CUSTOM_TOOL_USE = "agent.custom_tool_use" +EVENT_TYPE_USER_TOOL_CONFIRMATION = "user.tool_confirmation" +EVENT_TYPE_USER_TOOL_RESULT = "user.tool_result" +EVENT_TYPE_USER_CUSTOM_TOOL_RESULT = "user.custom_tool_result" +EVENT_TYPE_SESSION_STATUS_IDLE = "session.status_idle" +EVENT_TYPE_SESSION_STATUS_TERMINATED = "session.status_terminated" +EVENT_TYPE_SESSION_DELETED = "session.deleted" + +PERMISSION_ALLOW = "allow" +PERMISSION_DENY = "deny" + +CONFIRMATION_ALLOW = "allow" +CONFIRMATION_DENY = "deny" + +EVENT_LIST_ORDER_ASC = "asc" +EVENT_LIST_ORDER_DESC = "desc" + +SESSION_STOP_REASON_END_TURN = "end_turn" +SESSION_STOP_REASON_REQUIRES_ACTION = "requires_action" +SESSION_STOP_REASON_RETRIES_EXHAUSTED = "retries_exhausted" + +WORK_STATE_QUEUED = WorkState.queued +WORK_STATE_STARTING = WorkState.starting +WORK_STATE_ACTIVE = WorkState.active +WORK_STATE_STOPPING = WorkState.stopping +WORK_STATE_STOPPED = WorkState.stopped + +HeartbeatResponse = HeartbeatWorkResponse +WorkTag = VolcTag + +EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT = "NO_HEARTBEAT" + +WORKER_ERROR_KIND_AUTH = "auth" +WORKER_ERROR_KIND_PERMISSION = "permission" +WORKER_ERROR_KIND_LEASE_CONFLICT = "lease_conflict" +WORKER_ERROR_KIND_RATE_LIMIT = "rate_limit" +WORKER_ERROR_KIND_TIMEOUT = "timeout" +WORKER_ERROR_KIND_NETWORK = "network" +WORKER_ERROR_KIND_TOOL_ERROR = "tool_error" +WORKER_ERROR_KIND_INVALID_RESPONSE = "invalid_response" + +DEFAULT_MAX_IDLE_SECONDS = 60.0 +DEFAULT_TOOL_TIMEOUT_SECONDS = 120.0 +DEFAULT_HEARTBEAT_SECONDS = 30.0 + + +class EventStreamUnsupported(RuntimeError): + """Raised when an API adapter does not support SSE event streaming.""" + + +class SessionTerminated(RuntimeError): + """Raised when the session was terminated by the control plane.""" + + +class IdleTimeout(RuntimeError): + """Raised after an end_turn idle event remains idle for max_idle.""" + + +class APIError(RuntimeError): + """Control-plane API error with HTTP status and request id.""" + + def __init__(self, status_code: int, message: str, request_id: str = "") -> None: + super().__init__(f"worker api status {status_code}: {message}") + self.status_code = status_code + self.message = message + self.request_id = request_id + + +class WorkerError(RuntimeError): + """Stable worker error classification for callers.""" + + def __init__( + self, + kind: str, + message: str, + *, + request_id: str = "", + retryable: bool = False, + cause: Optional[BaseException] = None, + ) -> None: + super().__init__(message) + self.kind = kind + self.message = message + self.request_id = request_id + self.retryable = retryable + self.__cause__ = cause + + +def is_status(exc: BaseException, status_code: int) -> bool: + return isinstance(exc, APIError) and exc.status_code == status_code + + +def is_fatal_4xx(exc: BaseException) -> bool: + return isinstance(exc, APIError) and 400 <= exc.status_code < 500 and exc.status_code not in (408, 409, 412, 429) + + +def classify_worker_error(exc: BaseException) -> WorkerError: + if isinstance(exc, WorkerError): + return exc + kind = WORKER_ERROR_KIND_NETWORK + retryable = True + request_id = "" + if isinstance(exc, TimeoutError): + kind = WORKER_ERROR_KIND_TIMEOUT + if isinstance(exc, APIError): + request_id = exc.request_id + retryable = False + if exc.status_code == 401: + kind = WORKER_ERROR_KIND_AUTH + elif exc.status_code == 403: + kind = WORKER_ERROR_KIND_PERMISSION + elif exc.status_code in (409, 412): + kind = WORKER_ERROR_KIND_LEASE_CONFLICT + elif exc.status_code == 408: + kind = WORKER_ERROR_KIND_TIMEOUT + retryable = True + elif exc.status_code == 429: + kind = WORKER_ERROR_KIND_RATE_LIMIT + retryable = True + elif exc.status_code >= 500: + kind = WORKER_ERROR_KIND_NETWORK + retryable = True + else: + kind = WORKER_ERROR_KIND_INVALID_RESPONSE + return WorkerError(kind, str(exc), request_id=request_id, retryable=retryable, cause=exc) + + +def work_session_id(item: WorkItem) -> str: + if item.data.id and (not item.data.type or item.data.type == "session"): + return item.data.id + return "" + + +@dataclass +class SkillRef: + name: str = "" + display_name: str = "" + id: str = "" + skill_id: str = "" + type: str = "" + version: str = "" + download_url: str = "" + + @classmethod + def from_mapping(cls, raw: Mapping[str, Any]) -> "SkillRef": + return cls( + name=str(raw.get("name") or ""), + display_name=str(raw.get("display_name") or ""), + id=str(raw.get("id") or ""), + skill_id=str(raw.get("skill_id") or ""), + type=str(raw.get("type") or ""), + version=str(raw.get("version") or ""), + download_url=str(raw.get("download_url") or ""), + ) + + def id_value(self) -> str: + return self.skill_id or self.id + + def name_value(self) -> str: + return self.name or self.display_name or self.id_value() + + +@dataclass +class AgentConfig: + skills: List[SkillRef] = field(default_factory=list) + + +@dataclass +class Session: + id: str + agent: AgentConfig = field(default_factory=AgentConfig) + skills: List[SkillRef] = field(default_factory=list) + raw: Dict[str, Any] = field(default_factory=dict) + + @classmethod + def from_mapping(cls, raw: Mapping[str, Any]) -> "Session": + agent_raw = raw.get("agent") if isinstance(raw.get("agent"), Mapping) else {} + return cls( + id=str(raw.get("id") or ""), + agent=AgentConfig(skills=_skill_refs(agent_raw.get("skills") or [])), + skills=_skill_refs(raw.get("skills") or []), + raw=dict(raw), + ) + + def skill_refs(self) -> List[SkillRef]: + return self.skills or self.agent.skills + + +@dataclass +class ContentBlock: + type: str + text: str = "" + media_type: str = "" + data: Any = None + + def to_dict(self) -> Dict[str, Any]: + out: Dict[str, Any] = {"type": self.type} + if self.text: + out["text"] = self.text + if self.media_type: + out["media_type"] = self.media_type + if self.data is not None: + out["data"] = self.data + return out + + +@dataclass +class Event: + type: str + id: str = "" + name: str = "" + input: Any = None + processed_at: str = "" + evaluated_permission: str = "" + session_thread_id: str = "" + tool_use_id: str = "" + custom_tool_use_id: str = "" + result: str = "" + deny_message: str = "" + stop_reason: Any = None + content: List[ContentBlock] = field(default_factory=list) + is_error: Optional[bool] = None + extra: Dict[str, Any] = field(default_factory=dict) + + @classmethod + def from_mapping(cls, raw: Mapping[str, Any]) -> "Event": + known = { + "id", + "type", + "name", + "input", + "processed_at", + "evaluated_permission", + "session_thread_id", + "tool_use_id", + "custom_tool_use_id", + "result", + "deny_message", + "stop_reason", + "content", + "is_error", + } + content = [] + for block in raw.get("content") or []: + if isinstance(block, Mapping): + content.append( + ContentBlock( + type=str(block.get("type") or ""), + text=str(block.get("text") or ""), + media_type=str(block.get("media_type") or ""), + data=block.get("data"), + ) + ) + return cls( + id=str(raw.get("id") or ""), + type=str(raw.get("type") or ""), + name=str(raw.get("name") or ""), + input=raw.get("input"), + processed_at=str(raw.get("processed_at") or ""), + evaluated_permission=str(raw.get("evaluated_permission") or ""), + session_thread_id=str(raw.get("session_thread_id") or ""), + tool_use_id=str(raw.get("tool_use_id") or ""), + custom_tool_use_id=str(raw.get("custom_tool_use_id") or ""), + result=str(raw.get("result") or ""), + deny_message=str(raw.get("deny_message") or ""), + stop_reason=raw.get("stop_reason"), + content=content, + is_error=raw.get("is_error") if isinstance(raw.get("is_error"), bool) else None, + extra={k: v for k, v in raw.items() if k not in known}, + ) + + def stop_reason_type(self) -> str: + if isinstance(self.stop_reason, Mapping): + return str(self.stop_reason.get("type") or "") + if isinstance(self.stop_reason, str): + return self.stop_reason + return "" + + def to_dict(self) -> Dict[str, Any]: + out = dict(self.extra) + out.update( + { + "id": self.id, + "type": self.type, + "processed_at": self.processed_at, + } + ) + for key in ( + "name", + "evaluated_permission", + "session_thread_id", + "tool_use_id", + "custom_tool_use_id", + "result", + "deny_message", + ): + value = getattr(self, key) + if value: + out[key] = value + if self.input is not None: + out["input"] = self.input + if self.stop_reason is not None: + out["stop_reason"] = self.stop_reason + if self.content: + out["content"] = [block.to_dict() for block in self.content] + if self.is_error is not None: + out["is_error"] = self.is_error + return {k: v for k, v in out.items() if v is not None and v != ""} + + +@dataclass +class ListEventsResponse: + events: List[Event] = field(default_factory=list) + next_page: str = "" + + +@dataclass +class SkillContent: + body: Any + content_length: int = -1 + file_name: str = "" + content_type: str = "" + + +@dataclass +class ToolCallResult: + tool_use_id: str + name: str + custom: bool + confirmation: str = "" + posted: bool = False + event: Optional[Event] = None + result: Optional[Event] = None + + +def new_event_id(prefix: str = "evt") -> str: + return f"{prefix}-{secrets.token_hex(8)}" + + +def utc_now_iso() -> str: + return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") + + +def new_user_tool_result_event( + tool_use_id: str, + content: Iterable[ContentBlock], + is_error: bool, + thread_id: str = "", +) -> Event: + return Event( + id=new_event_id("evt"), + type=EVENT_TYPE_USER_TOOL_RESULT, + tool_use_id=tool_use_id, + content=list(content), + is_error=is_error, + processed_at=utc_now_iso(), + session_thread_id=thread_id, + ) + + +def new_user_custom_tool_result_event( + custom_tool_use_id: str, + content: Iterable[ContentBlock], + is_error: bool, + thread_id: str = "", +) -> Event: + return Event( + id=new_event_id("evt"), + type=EVENT_TYPE_USER_CUSTOM_TOOL_RESULT, + custom_tool_use_id=custom_tool_use_id, + content=list(content), + is_error=is_error, + processed_at=utc_now_iso(), + session_thread_id=thread_id, + ) + + +def tool_confirmation_call_id(event: Event) -> str: + return event.tool_use_id or event.custom_tool_use_id + + +def tool_use_call_id(event: Event) -> str: + return event.tool_use_id or event.custom_tool_use_id or event.id + + +def tool_result_call_id(event: Event) -> str: + return event.tool_use_id or event.custom_tool_use_id + + +def _skill_refs(raw: Iterable[Any]) -> List[SkillRef]: + out: List[SkillRef] = [] + for item in raw: + if isinstance(item, Mapping): + out.append(SkillRef.from_mapping(item)) + return out diff --git a/src/arkruntime/selfhosted/worker.py b/src/arkruntime/selfhosted/worker.py new file mode 100644 index 0000000..d1ff5fb --- /dev/null +++ b/src/arkruntime/selfhosted/worker.py @@ -0,0 +1,482 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import hashlib +import logging +import os +import random +import re +import socket +import threading +import time +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Dict, Optional + +from .envinit import Initializer, InitializerOptions +from .session_tool_runner import SessionToolRunner, SessionToolRunnerOptions +from .tool_result_store import FileToolResultStore +from .tools import Tool, ToolContext, ToolSet, default_toolset +from .types import ( + DEFAULT_HEARTBEAT_SECONDS, + DEFAULT_MAX_IDLE_SECONDS, + EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT, + WORK_STATE_STOPPED, + WORK_STATE_STOPPING, + IdleTimeout, + SessionTerminated, + WorkItem, + is_fatal_4xx, + is_status, + work_session_id, +) + +DEFAULT_POLL_BLOCK_MS = 999 +POLL_BACKOFF_CAP_SECONDS = 60.0 + + +@dataclass +class WorkPollerOptions: + environment_id: str + worker_id: str = "" + block_ms: int = DEFAULT_POLL_BLOCK_MS + reclaim_older_than_ms: int = 0 + drain: bool = False + auto_stop: bool = True + stop_event: Optional[threading.Event] = None + logger: logging.Logger = logging.getLogger("arkruntime.selfhosted.work_poller") + + +class WorkPoller: + def __init__(self, api: Any, options: WorkPollerOptions) -> None: + if api is None: + raise ValueError("api is required") + if not options.environment_id: + raise ValueError("environment_id is required") + if not options.worker_id: + options.worker_id = default_worker_id() + self.api = api + self.options = options + self.current: Optional[WorkItem] = None + self.error: Optional[BaseException] = None + self.closed = False + self._pending_stop = None + self._failures = 0 + self._discards = 0 + + def close(self) -> None: + self.closed = True + self._run_pending_stop() + + def next(self) -> Optional[WorkItem]: + self._run_pending_stop() + if self._is_closed(): + return None + while not self._is_closed(): + try: + item = self.api.poll_work( + self.options.environment_id, + worker_id=self.options.worker_id, + block_ms=self.options.block_ms, + reclaim_older_than_ms=self.options.reclaim_older_than_ms, + ) + except Exception as exc: # noqa: BLE001 - worker owns retry classification. + if is_fatal_4xx(exc): + self.error = exc + return None + self._failures += 1 + sleep_seconds = _jitter(_backoff(self._failures) / 2, _backoff(self._failures)) + self.options.logger.warning("poll work failed err=%s sleep=%.3fs", exc, sleep_seconds) + self._sleep(sleep_seconds) + continue + self._failures = 0 + if item is None or not item.id: + if self.options.drain: + return None + self._sleep(_jitter(1, 3)) + continue + if not item.environment_id: + item.environment_id = self.options.environment_id + if not work_session_id(item): + self.options.logger.warning( + "discard invalid work work_id=%s reason=missing session id", + item.id, + ) + self._discard_invalid_work(item) + continue + try: + self.api.ack_work(item.environment_id, item.id, worker_id=self.options.worker_id) + except Exception as exc: # noqa: BLE001 - poller owns retry and discard behavior. + self.options.logger.warning("ack work failed work_id=%s err=%s", item.id, exc) + # ACK is the queued -> starting ownership race. A failed ACK + # does not prove that this worker owns the work, so it must + # never stop the item claimed by another worker. + if _is_resolved_status(exc): + continue + if is_fatal_4xx(exc): + self.error = exc + return None + self._backoff_discard() + continue + self.current = item + if self.options.auto_stop: + self._pending_stop = lambda item=item: self._stop_item(item, force=False) + self._discards = 0 + self.options.logger.info( + "claimed work work_id=%s session_id=%s", + item.id, + work_session_id(item), + ) + return item + return None + + def _run_pending_stop(self) -> None: + pending = self._pending_stop + self._pending_stop = None + self.current = None + if pending is not None: + pending() + + def _discard_invalid_work(self, item: WorkItem) -> None: + try: + self.api.ack_work(item.environment_id, item.id, worker_id=self.options.worker_id) + except Exception as exc: # noqa: BLE001 - invalid work still obeys ACK ownership. + self.options.logger.warning("ack invalid work failed work_id=%s err=%s", item.id, exc) + return + self._stop_item(item, force=True) + self._backoff_discard() + + def _stop_item(self, item: WorkItem, *, force: bool) -> None: + try: + self.api.stop_work(item.environment_id, item.id, force=force) + except Exception as exc: # noqa: BLE001 + if not _is_resolved_status(exc): + self.options.logger.warning("stop work failed work_id=%s err=%s", item.id, exc) + + def _backoff_discard(self) -> None: + self._discards += 1 + self._sleep(_jitter(_backoff(self._discards) / 2, _backoff(self._discards))) + + def _is_closed(self) -> bool: + return self.closed or bool(self.options.stop_event and self.options.stop_event.is_set()) + + def _sleep(self, seconds: float) -> None: + if self.options.stop_event is not None: + self.options.stop_event.wait(max(seconds, 0)) + return + time.sleep(max(seconds, 0)) + + +@dataclass +class HandleItemOptions: + work_id: str = "" + environment_id: str = "" + session_id: str = "" + latest_heartbeat_at: str = "" + + +@dataclass +class _ClaimedWork: + id: str + environment_id: str + session_id: str + latest_heartbeat_at: str = "" + + +@dataclass +class EnvironmentWorkerOptions: + environment_id: str = "" + worker_id: str = "" + workdir: str = "." + unrestricted_paths: bool = False + tool_context: Optional[ToolContext] = None + tools: Optional[ToolSet] = None + max_idle_seconds: Optional[float] = DEFAULT_MAX_IDLE_SECONDS + custom_tools: Dict[str, Tool] = field(default_factory=dict) + logger: logging.Logger = logging.getLogger("arkruntime.selfhosted.environment_worker") + + +class EnvironmentWorker: + def __init__(self, api: Any, options: EnvironmentWorkerOptions) -> None: + if api is None: + raise ValueError("api is required") + if not options.worker_id: + options.worker_id = default_worker_id() + self.api = api + self.options = options + self._stop = threading.Event() + + def close(self) -> None: + self._stop.set() + + def run(self) -> None: + if not self.options.environment_id: + raise ValueError("environment_id is required") + poller = WorkPoller( + self.api, + WorkPollerOptions( + environment_id=self.options.environment_id, + worker_id=self.options.worker_id, + auto_stop=False, + stop_event=self._stop, + logger=self.options.logger, + ), + ) + try: + while not self._stop.is_set(): + item = poller.next() + if item is None: + if poller.error is not None: + raise poller.error + return + try: + self._handle_item(_claimed_work_from_item(item), use_workdir_as_session=False) + except (IdleTimeout, SessionTerminated): + pass + except Exception as exc: # noqa: BLE001 - continue polling after a bad work item. + self.options.logger.warning("handle work failed: %s", exc) + finally: + poller.close() + + def handle_item(self, options: HandleItemOptions) -> None: + work = self._claimed_work_from_options(options) + try: + self._handle_item(work, use_workdir_as_session=True) + except (IdleTimeout, SessionTerminated): + return + + def _handle_item(self, work: _ClaimedWork, *, use_workdir_as_session: bool) -> 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 + try: + workdir = self._workdir_for(work.session_id, use_workdir_as_session) + heartbeat_thread = threading.Thread( + target=self._heartbeat_loop, + args=(work, heartbeat_stop, heartbeat_done, heartbeat_cause), + daemon=True, + ) + heartbeat_thread.start() + heartbeat = heartbeat_thread + session = self.api.get_session(work.session_id) + if work_stop.is_set(): + return + if session is None: + raise ValueError("session response is empty") + if not session.id: + session.id = work.session_id + Initializer( + self.api, + InitializerOptions(workdir=workdir, logger=self.options.logger), + ).setup(session) + if work_stop.is_set(): + return + tool_context = self._tool_context(workdir, work_stop) + store = FileToolResultStore(workdir) + runner = SessionToolRunner( + self.api, + work.session_id, + SessionToolRunnerOptions( + work_id=work.id, + tools=self.options.tools or default_toolset(), + tool_context=tool_context, + custom_tools=self.options.custom_tools, + result_store=store, + max_idle_seconds=self.options.max_idle_seconds, + stop_event=work_stop, + logger=self.options.logger, + ), + ) + runner.run() + finally: + heartbeat_stop.set() + if heartbeat is not None: + heartbeat_done.wait(timeout=DEFAULT_HEARTBEAT_SECONDS + 1) + cause = heartbeat_cause["value"] + if _should_stop_item(cause): + try: + self.api.stop_work(work.environment_id, work.id, force=True) + except Exception as exc: # noqa: BLE001 + if not _is_resolved_status(exc): + self.options.logger.warning("stop work failed: %s", exc) + else: + self.options.logger.info( + "skip stop work after heartbeat ownership became uncertain cause=%s", + cause, + ) + + def _heartbeat_loop(self, work: _ClaimedWork, stop: threading.Event, done: threading.Event, cause: dict) -> None: + interval = max(1.0, min(DEFAULT_HEARTBEAT_SECONDS / 2, DEFAULT_HEARTBEAT_SECONDS)) + ttl = DEFAULT_HEARTBEAT_SECONDS + last = work.latest_heartbeat_at or EXPECTED_LAST_HEARTBEAT_NO_HEARTBEAT + last_success = time.monotonic() + try: + while not stop.is_set(): + try: + resp = self.api.heartbeat_work( + work.environment_id, + work.id, + expected_last_heartbeat=last, + desired_ttl_seconds=int(ttl), + ) + except Exception as exc: # noqa: BLE001 + if is_status(exc, 412): + cause["value"] = "lease_lost" + stop.set() + return + if is_fatal_4xx(exc): + cause["value"] = "heartbeat_permanent_failure" + stop.set() + return + if time.monotonic() - last_success > ttl: + cause["value"] = "heartbeat_lost" + stop.set() + return + self.options.logger.warning( + "heartbeat failed work_id=%s session_id=%s since_last_success=%.3fs ttl=%.3fs err=%s", + work.id, + work.session_id, + time.monotonic() - last_success, + ttl, + exc, + ) + stop.wait(interval) + continue + if resp is None: + if time.monotonic() - last_success > ttl: + cause["value"] = "heartbeat_lost" + stop.set() + return + self.options.logger.warning( + "heartbeat empty response work_id=%s session_id=%s", + work.id, + work.session_id, + ) + stop.wait(interval) + continue + last_success = time.monotonic() + if resp.last_heartbeat: + last = resp.last_heartbeat + if resp.ttl_seconds > 0: + ttl = float(resp.ttl_seconds) + interval = max(1.0, min(ttl / 2, DEFAULT_HEARTBEAT_SECONDS)) + if resp.state in (WORK_STATE_STOPPING, WORK_STATE_STOPPED): + cause["value"] = "stop_requested" + stop.set() + return + if resp.lease_extended is False: + cause["value"] = "lease_not_extended" + stop.set() + return + stop.wait(interval) + finally: + done.set() + + def _tool_context(self, workdir: str, cancel_event: Any) -> ToolContext: + base = self.options.tool_context or ToolContext(workdir=workdir) + env = None if base.env is None else dict(base.env) + return ToolContext( + workdir=workdir, + env=env, + unrestricted_paths=self.options.unrestricted_paths or base.unrestricted_paths, + tool_timeout_seconds=base.tool_timeout_seconds, + cancel_event=cancel_event, + ) + + def _workdir_for(self, session_id: str, use_workdir_as_session: bool) -> str: + 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 + + def _claimed_work_from_options(self, options: HandleItemOptions) -> _ClaimedWork: + work_id = options.work_id or os.environ.get("MA_WORK_ID", "") + environment_id = options.environment_id or os.environ.get("MA_ENVIRONMENT_ID", "") + session_id = options.session_id or os.environ.get("MA_SESSION_ID", "") + latest_heartbeat = options.latest_heartbeat_at or os.environ.get("MA_LATEST_HEARTBEAT_AT", "") + if not work_id: + raise ValueError("work id is required") + if not environment_id: + raise ValueError("environment id is required") + if not session_id: + raise ValueError("session id is required") + return _ClaimedWork( + id=work_id, + environment_id=environment_id, + session_id=session_id, + latest_heartbeat_at=latest_heartbeat, + ) + + +def _claimed_work_from_item(item: WorkItem) -> _ClaimedWork: + session_id = work_session_id(item) + if not item.id: + raise ValueError("work item id must not be empty") + if not session_id: + raise ValueError("work item does not contain session id") + return _ClaimedWork( + id=item.id, + environment_id=item.environment_id, + session_id=session_id, + latest_heartbeat_at=item.latest_heartbeat_at or "", + ) + + +def default_worker_id() -> str: + return f"{socket.gethostname()}-{os.getpid()}" + + +def _backoff(failures: int) -> float: + value = min(POLL_BACKOFF_CAP_SECONDS, 2 ** max(failures - 1, 0)) + return float(value) + + +def _is_resolved_status(exc: BaseException) -> bool: + return is_status(exc, 404) or is_status(exc, 409) or is_status(exc, 412) + + +def _should_stop_item(heartbeat_cause: str) -> bool: + return heartbeat_cause not in { + "lease_lost", + "lease_not_extended", + "heartbeat_lost", + "heartbeat_permanent_failure", + } + + +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 + + def is_set(self) -> bool: + return any(event.is_set() for event in self._events) + + def wait(self, timeout: Optional[float] = None) -> bool: + deadline = None if timeout is None else time.monotonic() + max(timeout, 0) + while not self.is_set(): + if deadline is not None and time.monotonic() >= deadline: + return False + time.sleep(0.05) + return True + + +def _jitter(low: float, high: float) -> float: + if high <= low: + return max(high, 0) + return low + random.random() * (high - low) diff --git a/src/arkruntime/types/agent/__init__.py b/src/arkruntime/types/agent/__init__.py index fc5f6e1..92ea5cd 100644 --- a/src/arkruntime/types/agent/__init__.py +++ b/src/arkruntime/types/agent/__init__.py @@ -11,6 +11,7 @@ from .agent_ref_type import AgentRefType from .agent_skill_ref import AgentSkillRef from .create_agent_request import CreateAgentRequest +from .custom_tool_input_schema import CustomToolInputSchema from .delete_agent_response import DeleteAgentResponse from .list_agents_response import ListAgentsResponse from .mcp_server import MCPServer @@ -23,6 +24,7 @@ from .permission_policy_type import PermissionPolicyType from .skill_ref_type import SkillRefType from .tag import Tag +from .token_limits import TokenLimits from .tool_config import ToolConfig from .tool_default_config import ToolDefaultConfig from .tool_item import ToolItem @@ -34,6 +36,7 @@ "AgentRefType", "AgentSkillRef", "CreateAgentRequest", + "CustomToolInputSchema", "DeleteAgentResponse", "ListAgentsResponse", "MCPServer", @@ -46,6 +49,7 @@ "PermissionPolicyType", "SkillRefType", "Tag", + "TokenLimits", "ToolConfig", "ToolDefaultConfig", "ToolItem", @@ -53,4 +57,4 @@ ] # Hand-written extras (preserved across regen via Makefile rsync --exclude=*_shim.py). -from ._init_extras_shim import * # noqa: F401,F403 +from ._init_extras_shim import * # noqa: F401,F403,E402 diff --git a/src/arkruntime/types/agent/agent.py b/src/arkruntime/types/agent/agent.py index 502743e..aec62f7 100644 --- a/src/arkruntime/types/agent/agent.py +++ b/src/arkruntime/types/agent/agent.py @@ -76,6 +76,10 @@ class Agent(BaseModel): """ 资源标签。 """ + display_name: Optional[str] = None + """ + 展示名。 + """ created_at: str """ RFC 3339 时间。 diff --git a/src/arkruntime/types/agent/agent_ref.py b/src/arkruntime/types/agent/agent_ref.py index ea84e42..fba8cb0 100644 --- a/src/arkruntime/types/agent/agent_ref.py +++ b/src/arkruntime/types/agent/agent_ref.py @@ -6,11 +6,15 @@ from __future__ import annotations -from typing import Optional +from typing import List, Optional from arkruntime._models import BaseModel from .agent_ref_type import AgentRefType +from .agent_skill_ref import AgentSkillRef +from .mcp_server import MCPServer +from .model_config import ModelConfig +from .tool_item import ToolItem class AgentRef(BaseModel): @@ -30,3 +34,35 @@ class AgentRef(BaseModel): """ 被引用 Agent 的版本号。 """ + name: Optional[str] = None + """ + Session 响应中冻结的成员 Agent 名称。 + """ + description: Optional[str] = None + """ + Session 响应中冻结的成员 Agent 描述。 + """ + model: Optional[ModelConfig] = None + """ + Session 响应中冻结的成员 Agent 模型配置。 + """ + system: Optional[str] = None + """ + Session 响应中冻结的成员 Agent system prompt。 + """ + tools: Optional[List[ToolItem]] = None + """ + Session 响应中冻结的成员 Agent 工具配置。 + """ + mcp_servers: Optional[List[MCPServer]] = None + """ + Session 响应中冻结的成员 Agent MCP servers。 + """ + skills: Optional[List[AgentSkillRef]] = None + """ + Session 响应中冻结的成员 Agent skills。 + """ + display_name: Optional[str] = None + """ + Session 响应中冻结的成员 Agent 展示名。 + """ diff --git a/src/arkruntime/types/agent/agent_skill_ref.py b/src/arkruntime/types/agent/agent_skill_ref.py index 30fc672..f5747fe 100644 --- a/src/arkruntime/types/agent/agent_skill_ref.py +++ b/src/arkruntime/types/agent/agent_skill_ref.py @@ -30,3 +30,7 @@ class AgentSkillRef(BaseModel): """ Skill 版本号,可选;不传走最新。 """ + use_latest: Optional[bool] = None + """ + Session 快照中标识创建时用户选择的是使用最新版本。 + """ diff --git a/src/arkruntime/types/agent/create_agent_request.py b/src/arkruntime/types/agent/create_agent_request.py index 455ef37..8d71918 100644 --- a/src/arkruntime/types/agent/create_agent_request.py +++ b/src/arkruntime/types/agent/create_agent_request.py @@ -35,6 +35,10 @@ class CreateAgentRequest(BaseModel): """ 描述信息。 """ + display_name: Optional[str] = None + """ + 展示名。 + """ system: Optional[str] = None """ System prompt。 diff --git a/src/arkruntime/types/agent/custom_tool_input_schema.py b/src/arkruntime/types/agent/custom_tool_input_schema.py new file mode 100644 index 0000000..7ccd5ab --- /dev/null +++ b/src/arkruntime/types/agent/custom_tool_input_schema.py @@ -0,0 +1,30 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Dict, List, Optional + +from arkruntime._models import BaseModel + + +class CustomToolInputSchema(BaseModel): + """ + Custom tool 的输入 JSON Schema。 + """ + + type: Optional[str] = None + """ + JSON Schema 顶层类型;缺省按 `"object"` 处理。 + """ + properties: Optional[Dict[str, object]] = None + """ + JSON Schema properties 对象。 + """ + required: Optional[List[str]] = None + """ + 必填属性名列表。 + """ diff --git a/src/arkruntime/types/agent/model_config.py b/src/arkruntime/types/agent/model_config.py index 35f6479..43c768a 100644 --- a/src/arkruntime/types/agent/model_config.py +++ b/src/arkruntime/types/agent/model_config.py @@ -6,11 +6,12 @@ from __future__ import annotations -from typing import Optional +from typing import List, Optional from arkruntime._models import BaseModel from .model_speed import ModelSpeed +from .token_limits import TokenLimits class ModelConfig(BaseModel): @@ -26,3 +27,23 @@ class ModelConfig(BaseModel): """ 速度档位。空字符串走默认(并非所有模型都支持 `fast`)。 """ + token_limits: Optional[TokenLimits] = None + """ + 模型 token 限制快照。 + """ + input_modalities: Optional[List[str]] = None + """ + 底模支持的输入模态列表。 + """ + provider: Optional[str] = None + """ + 模型提供方。 + """ + thinking: Optional[str] = None + """ + thinking 配置。 + """ + reasoning_effort: Optional[str] = None + """ + 推理努力程度。 + """ diff --git a/src/arkruntime/types/agent/token_limits.py b/src/arkruntime/types/agent/token_limits.py new file mode 100644 index 0000000..61b232d --- /dev/null +++ b/src/arkruntime/types/agent/token_limits.py @@ -0,0 +1,30 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class TokenLimits(BaseModel): + """ + 模型上下文与输入输出 token 限制。 + """ + + context_window: Optional[int] = None + """ + 模型上下文窗口。 + """ + max_input_token_length: Optional[int] = None + """ + 最大输入 token 长度。 + """ + max_output_token_length: Optional[int] = None + """ + 最大输出 token 长度。 + """ diff --git a/src/arkruntime/types/agent/tool_item.py b/src/arkruntime/types/agent/tool_item.py index 3d141f7..cf6175a 100644 --- a/src/arkruntime/types/agent/tool_item.py +++ b/src/arkruntime/types/agent/tool_item.py @@ -10,6 +10,7 @@ from arkruntime._models import BaseModel +from .custom_tool_input_schema import CustomToolInputSchema from .tool_config import ToolConfig from .tool_default_config import ToolDefaultConfig @@ -20,6 +21,7 @@ class ToolItem(BaseModel): - `agent_toolset_`:内置工具集(当前默认 `agent_toolset_20260701`; 存量 `agent_toolset_20260401` 仍兼容) - `mcp_toolset`:来自 `mcp_servers[]` 的工具集 + - `evolution`:自演进类工具 - `custom`:客户端执行的自定义工具 所有变体字段合并在一个 model 里,未使用的字段留空即可(proto oneof @@ -50,8 +52,7 @@ class ToolItem(BaseModel): """ `custom` 专用;1–1024 字符。 """ - input_schema: Optional[str] = None + input_schema: Optional[CustomToolInputSchema] = None """ - `custom` 专用;承载 JSON Schema 的字符串形态 - (wire 上是 JSON-encoded string,非 nested object)。 + `custom` 专用;承载 JSON Schema 对象。 """ diff --git a/src/arkruntime/types/agent/update_agent_request.py b/src/arkruntime/types/agent/update_agent_request.py index a8cc128..08b7940 100644 --- a/src/arkruntime/types/agent/update_agent_request.py +++ b/src/arkruntime/types/agent/update_agent_request.py @@ -14,6 +14,7 @@ from .mcp_server import MCPServer from .model_config import ModelConfig from .multiagent_config import MultiagentConfig +from .tag import Tag from .tool_item import ToolItem @@ -36,6 +37,10 @@ class UpdateAgentRequest(BaseModel): """ 人类可读名称。 """ + display_name: Optional[str] = None + """ + 展示名。 + """ model: Optional[ModelConfig] = None """ 模型配置。 @@ -68,3 +73,7 @@ class UpdateAgentRequest(BaseModel): """ 用户自定义键值对元数据(patch)。 """ + tags: Optional[List[Tag]] = None + """ + 资源标签(整体替换)。 + """ diff --git a/src/arkruntime/types/environment/__init__.py b/src/arkruntime/types/environment/__init__.py index 131b57f..3988ba4 100644 --- a/src/arkruntime/types/environment/__init__.py +++ b/src/arkruntime/types/environment/__init__.py @@ -12,11 +12,20 @@ from .env_config_type import EnvConfigType from .environment import Environment from .environment_scope import EnvironmentScope +from .environment_with_overrides import EnvironmentWithOverrides +from .heartbeat_work_response import HeartbeatWorkResponse from .list_environments_response import ListEnvironmentsResponse from .networking_config import NetworkingConfig from .networking_type import NetworkingType from .packages_config import PackagesConfig +from .poll_work_empty_response import PollWorkEmptyResponse +from .stop_work_body import StopWorkBody +from .tos_config import TosConfig from .update_environment_request import UpdateEnvironmentRequest +from .volc_tag import VolcTag +from .work_data import WorkData +from .work_item import WorkItem +from .work_state import WorkState __all__ = [ "CreateEnvironmentRequest", @@ -25,12 +34,21 @@ "EnvConfigType", "Environment", "EnvironmentScope", + "EnvironmentWithOverrides", + "HeartbeatWorkResponse", "ListEnvironmentsResponse", "NetworkingConfig", "NetworkingType", "PackagesConfig", + "PollWorkEmptyResponse", + "StopWorkBody", + "TosConfig", "UpdateEnvironmentRequest", + "VolcTag", + "WorkData", + "WorkItem", + "WorkState", ] # Hand-written extras (preserved across regen via Makefile rsync --exclude=*_shim.py). -from ._init_extras_shim import * # noqa: F401,F403 +from ._init_extras_shim import * # noqa: F401,F403,E402 diff --git a/src/arkruntime/types/environment/env_config.py b/src/arkruntime/types/environment/env_config.py index 5e3e60b..a177f81 100644 --- a/src/arkruntime/types/environment/env_config.py +++ b/src/arkruntime/types/environment/env_config.py @@ -13,6 +13,7 @@ from .env_config_type import EnvConfigType from .networking_config import NetworkingConfig from .packages_config import PackagesConfig +from .tos_config import TosConfig class EnvConfig(BaseModel): @@ -36,3 +37,11 @@ class EnvConfig(BaseModel): """ 容器启动时注入的环境变量。 """ + setup_script: Optional[str] = None + """ + 沙箱启动阶段执行的初始化脚本。 + """ + tos: Optional[TosConfig] = None + """ + Environment outputs 的 TOS 存储配置。 + """ diff --git a/src/arkruntime/types/environment/environment.py b/src/arkruntime/types/environment/environment.py index 98ef3b3..27d7f91 100644 --- a/src/arkruntime/types/environment/environment.py +++ b/src/arkruntime/types/environment/environment.py @@ -6,7 +6,7 @@ from __future__ import annotations -from typing import Dict, Literal, Optional +from typing import Dict, List, Literal, Optional from arkruntime._models import BaseModel @@ -55,3 +55,7 @@ class Environment(BaseModel): """ RFC 3339 时间。 """ + overridden_fields: Optional[List[str]] = None + """ + Session 使用 EnvironmentWithOverrides 时,本次被覆写的 config 子字段。 + """ diff --git a/src/arkruntime/types/environment/environment_with_overrides.py b/src/arkruntime/types/environment/environment_with_overrides.py new file mode 100644 index 0000000..52a48cb --- /dev/null +++ b/src/arkruntime/types/environment/environment_with_overrides.py @@ -0,0 +1,33 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Literal, Optional + +from arkruntime._models import BaseModel + +from .env_config import EnvConfig + + +class EnvironmentWithOverrides(BaseModel): + """ + CreateSession 时的 Environment 覆写引用;以已有 environment 为底,只覆写 + 运行时 config。 + """ + + type: Literal["environment_with_overrides"] + """ + 固定 `"environment_with_overrides"`。 + """ + id: str + """ + Base Environment ID。 + """ + config: Optional[EnvConfig] = None + """ + 运行时配置覆写;省略表示继承 base。 + """ diff --git a/src/arkruntime/types/environment/heartbeat_work_response.py b/src/arkruntime/types/environment/heartbeat_work_response.py new file mode 100644 index 0000000..7a4a440 --- /dev/null +++ b/src/arkruntime/types/environment/heartbeat_work_response.py @@ -0,0 +1,40 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Literal + +from arkruntime._models import BaseModel + +from .work_state import WorkState + + +class HeartbeatWorkResponse(BaseModel): + """ + Heartbeat work 的响应体。 + """ + + last_heartbeat: str + """ + 控制面接受的 heartbeat 时间,RFC 3339。 + """ + lease_extended: bool + """ + Lease 是否被刷新。 + """ + state: WorkState + """ + Work 生命周期状态。 + """ + ttl_seconds: int + """ + Lease TTL 秒数。 + """ + type: Literal["work_heartbeat"] + """ + 对象类型,固定为 `work_heartbeat`。 + """ diff --git a/src/arkruntime/types/environment/poll_work_empty_response.py b/src/arkruntime/types/environment/poll_work_empty_response.py new file mode 100644 index 0000000..3d6cfaa --- /dev/null +++ b/src/arkruntime/types/environment/poll_work_empty_response.py @@ -0,0 +1,16 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Dict + +from typing_extensions import TypeAliasType + +PollWorkEmptyResponse = TypeAliasType("PollWorkEmptyResponse", Dict[str, object]) +""" +Poll 没有可领取 work 时的响应;MA `/api/v3` 返回 HTTP 200 + `{}`。 +""" diff --git a/src/arkruntime/types/environment/stop_work_body.py b/src/arkruntime/types/environment/stop_work_body.py new file mode 100644 index 0000000..8f8693b --- /dev/null +++ b/src/arkruntime/types/environment/stop_work_body.py @@ -0,0 +1,22 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class StopWorkBody(BaseModel): + """ + Stop work 的请求体。 + """ + + force: Optional[bool] = None + """ + 是否强制停止。 + """ diff --git a/src/arkruntime/types/environment/tos_config.py b/src/arkruntime/types/environment/tos_config.py new file mode 100644 index 0000000..9daab64 --- /dev/null +++ b/src/arkruntime/types/environment/tos_config.py @@ -0,0 +1,27 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class TosConfig(BaseModel): + """ + Environment 产物存储位置。设置后 outputs 文件会注册到用户指定的 TOS + bucket/prefix;不设置则走方舟默认存储。 + """ + + bucket: Optional[str] = None + """ + TOS bucket 名称。 + """ + prefix: Optional[str] = None + """ + TOS 前缀。 + """ diff --git a/src/arkruntime/types/environment/volc_tag.py b/src/arkruntime/types/environment/volc_tag.py new file mode 100644 index 0000000..a974650 --- /dev/null +++ b/src/arkruntime/types/environment/volc_tag.py @@ -0,0 +1,26 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class VolcTag(BaseModel): + """ + Work 关联的标签。 + """ + + key: str + """ + 标签 key。 + """ + value: Optional[str] = None + """ + 标签 value。 + """ diff --git a/src/arkruntime/types/environment/work_data.py b/src/arkruntime/types/environment/work_data.py new file mode 100644 index 0000000..379c96e --- /dev/null +++ b/src/arkruntime/types/environment/work_data.py @@ -0,0 +1,24 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from arkruntime._models import BaseModel + + +class WorkData(BaseModel): + """ + Work 的业务载荷。 + """ + + id: str + """ + 业务对象 ID,例如 session ID。 + """ + type: str + """ + 业务载荷类型,例如 `session`。 + """ diff --git a/src/arkruntime/types/environment/work_item.py b/src/arkruntime/types/environment/work_item.py new file mode 100644 index 0000000..23690d8 --- /dev/null +++ b/src/arkruntime/types/environment/work_item.py @@ -0,0 +1,74 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import List, Literal, Optional + +from arkruntime._models import BaseModel + +from .volc_tag import VolcTag +from .work_data import WorkData +from .work_state import WorkState + + +class WorkItem(BaseModel): + """ + Worker queue 中的一条 work。 + """ + + id: str + """ + Work ID。 + """ + acknowledged_at: Optional[str] = None + """ + Work ack 时间,RFC 3339。 + """ + created_at: str + """ + Work 创建时间,RFC 3339。 + """ + data: WorkData + """ + 业务载荷。 + """ + environment_id: str + """ + Environment ID。 + """ + latest_heartbeat_at: Optional[str] = None + """ + 最近 heartbeat 时间,RFC 3339。 + """ + tags: Optional[List[VolcTag]] = None + """ + Work 标签。 + """ + secret: Optional[str] = None + """ + Work secret;仅 poll 时返回,ack / stop 响应会抹掉。 + """ + started_at: Optional[str] = None + """ + Work 开始时间,RFC 3339。 + """ + state: WorkState + """ + Work 生命周期状态。 + """ + stop_requested_at: Optional[str] = None + """ + 控制面请求停止的时间,RFC 3339。 + """ + stopped_at: Optional[str] = None + """ + Work 停止时间,RFC 3339。 + """ + type: Literal["work"] + """ + 对象类型,固定为 `work`。 + """ diff --git a/src/arkruntime/types/environment/work_state.py b/src/arkruntime/types/environment/work_state.py new file mode 100644 index 0000000..b7e74b8 --- /dev/null +++ b/src/arkruntime/types/environment/work_state.py @@ -0,0 +1,21 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from enum import Enum + + +class WorkState(str, Enum): + """ + Work 生命周期状态。 + """ + + queued = "queued" + starting = "starting" + active = "active" + stopping = "stopping" + stopped = "stopped" diff --git a/src/arkruntime/types/session/__init__.py b/src/arkruntime/types/session/__init__.py index d26a553..a504f85 100644 --- a/src/arkruntime/types/session/__init__.py +++ b/src/arkruntime/types/session/__init__.py @@ -15,6 +15,11 @@ from .create_session_resource_request import CreateSessionResourceRequest from .delete_session_response import DeleteSessionResponse from .document_source import DocumentSource +from .environment_config_override import EnvironmentConfigOverride +from .environment_networking_config import EnvironmentNetworkingConfig +from .environment_packages_config import EnvironmentPackagesConfig +from .environment_tos_config import EnvironmentTosConfig +from .environment_with_overrides import EnvironmentWithOverrides from .file_document_source import FileDocumentSource from .file_image_source import FileImageSource from .image_source import ImageSource @@ -66,6 +71,7 @@ from .managed_agents_user_tool_result_event_params import ( ManagedAgentsUserToolResultEventParams, ) +from .model_overrides import ModelOverrides from .plain_text_document_source import PlainTextDocumentSource from .search_result_citations import SearchResultCitations from .search_result_content import SearchResultContent @@ -95,6 +101,11 @@ "CreateSessionResourceRequest", "DeleteSessionResponse", "DocumentSource", + "EnvironmentConfigOverride", + "EnvironmentNetworkingConfig", + "EnvironmentPackagesConfig", + "EnvironmentTosConfig", + "EnvironmentWithOverrides", "FileDocumentSource", "FileImageSource", "ImageSource", @@ -130,6 +141,7 @@ "ManagedAgentsUserMessageEventParams", "ManagedAgentsUserToolConfirmationEventParams", "ManagedAgentsUserToolResultEventParams", + "ModelOverrides", "PlainTextDocumentSource", "SearchResultCitations", "SearchResultContent", @@ -151,4 +163,4 @@ ] # Hand-written extras (preserved across regen via Makefile rsync --exclude=*_shim.py). -from ._init_extras_shim import * # noqa: F401,F403 +from ._init_extras_shim import * # noqa: F401,F403,E402 diff --git a/src/arkruntime/types/session/agent_ref.py b/src/arkruntime/types/session/agent_ref.py index 274a943..7eae635 100644 --- a/src/arkruntime/types/session/agent_ref.py +++ b/src/arkruntime/types/session/agent_ref.py @@ -6,22 +6,24 @@ from __future__ import annotations -from typing import Literal, Optional +from typing import Dict, List, Optional from arkruntime._models import BaseModel +from .model_overrides import ModelOverrides + class AgentRef(BaseModel): """ - Agent 引用(对象形态):`type: "agent"` + id + optional version。 - 与 CreateSessionRequest.agent 联合使用。 + Agent 引用(对象形态):`type: "agent"` 或 `"agent_with_overrides"`。 + MA wire 上两种对象形态都走同一个 JSON object 承载,避免 SDK 生成复杂 union。 """ - type: Literal["agent"] + type: str """ - 固定 `"agent"`。 + `"agent"` 或 `"agent_with_overrides"`。 """ - id: str + id: Optional[str] = None """ Agent ID。 """ @@ -29,3 +31,31 @@ class AgentRef(BaseModel): """ Agent 版本号;不传走最新。 """ + system: Optional[str] = None + """ + System prompt 覆写。 + """ + tools: Optional[List[Dict[str, object]]] = None + """ + 工具配置覆写。 + """ + mcp_servers: Optional[List[Dict[str, object]]] = None + """ + MCP server 配置覆写。 + """ + skills: Optional[List[Dict[str, object]]] = None + """ + Skill 配置覆写。 + """ + multiagent: Optional[Dict[str, object]] = None + """ + 多 Agent 配置覆写。 + """ + display_name: Optional[str] = None + """ + Session 响应中冻结的 Agent 展示名。 + """ + model: Optional[ModelOverrides] = None + """ + 模型运行参数覆写。 + """ diff --git a/src/arkruntime/types/session/create_session_request.py b/src/arkruntime/types/session/create_session_request.py index 5e3c271..0df50ff 100644 --- a/src/arkruntime/types/session/create_session_request.py +++ b/src/arkruntime/types/session/create_session_request.py @@ -11,6 +11,7 @@ from arkruntime._models import BaseModel from .agent_identifier import AgentIdentifier +from .environment_with_overrides import EnvironmentWithOverrides from .session_resource import SessionResource from .tag import Tag @@ -28,9 +29,13 @@ class CreateSessionRequest(BaseModel): """ Agent 标识。 """ - environment_id: str + environment_id: Optional[str] = None """ - 关联的 Environment ID。 + 关联的 Environment ID。与 `environment` 二选一。 + """ + environment: Optional[EnvironmentWithOverrides] = None + """ + 关联 Environment 的覆写引用。与 `environment_id` 二选一。 """ tags: Optional[List[Tag]] = None """ diff --git a/src/arkruntime/types/session/environment_config_override.py b/src/arkruntime/types/session/environment_config_override.py new file mode 100644 index 0000000..ed09c0e --- /dev/null +++ b/src/arkruntime/types/session/environment_config_override.py @@ -0,0 +1,46 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Dict, Optional + +from arkruntime._models import BaseModel + +from .environment_networking_config import EnvironmentNetworkingConfig +from .environment_packages_config import EnvironmentPackagesConfig +from .environment_tos_config import EnvironmentTosConfig + + +class EnvironmentConfigOverride(BaseModel): + """ + Environment 覆写时使用的运行环境配置。 + """ + + type: str + """ + 运行环境类型。 + """ + networking: Optional[EnvironmentNetworkingConfig] = None + """ + 容器出网策略。 + """ + packages: Optional[EnvironmentPackagesConfig] = None + """ + 启动时预装的依赖包。 + """ + env: Optional[Dict[str, str]] = None + """ + 容器环境变量。 + """ + setup_script: Optional[str] = None + """ + 沙箱启动脚本。 + """ + tos: Optional[EnvironmentTosConfig] = None + """ + Environment outputs 的 TOS 存储配置。 + """ diff --git a/src/arkruntime/types/session/environment_networking_config.py b/src/arkruntime/types/session/environment_networking_config.py new file mode 100644 index 0000000..9ab87a9 --- /dev/null +++ b/src/arkruntime/types/session/environment_networking_config.py @@ -0,0 +1,34 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import List, Optional + +from arkruntime._models import BaseModel + + +class EnvironmentNetworkingConfig(BaseModel): + """ + Environment 覆写时使用的容器出网策略。 + """ + + type: str + """ + 出网策略类型。 + """ + allow_mcp_servers: Optional[bool] = None + """ + 是否允许出网到 MCP servers。 + """ + allow_package_managers: Optional[bool] = None + """ + 是否允许访问包管理器。 + """ + allowed_hosts: Optional[List[str]] = None + """ + 显式允许的出网域名。 + """ diff --git a/src/arkruntime/types/session/environment_packages_config.py b/src/arkruntime/types/session/environment_packages_config.py new file mode 100644 index 0000000..d16f50d --- /dev/null +++ b/src/arkruntime/types/session/environment_packages_config.py @@ -0,0 +1,46 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import List, Literal, Optional + +from arkruntime._models import BaseModel + + +class EnvironmentPackagesConfig(BaseModel): + """ + Environment 覆写时使用的预装依赖包配置。 + """ + + type: Optional[Literal["packages"]] = None + """ + 固定 `"packages"`。 + """ + pip: Optional[List[str]] = None + """ + pip 依赖。 + """ + apt: Optional[List[str]] = None + """ + apt 依赖。 + """ + npm: Optional[List[str]] = None + """ + npm 依赖。 + """ + cargo: Optional[List[str]] = None + """ + cargo 依赖。 + """ + gem: Optional[List[str]] = None + """ + gem 依赖。 + """ + go: Optional[List[str]] = None + """ + go module 依赖。 + """ diff --git a/src/arkruntime/types/session/environment_tos_config.py b/src/arkruntime/types/session/environment_tos_config.py new file mode 100644 index 0000000..6a9f43f --- /dev/null +++ b/src/arkruntime/types/session/environment_tos_config.py @@ -0,0 +1,26 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class EnvironmentTosConfig(BaseModel): + """ + Environment 覆写时使用的 TOS 配置。 + """ + + bucket: Optional[str] = None + """ + TOS bucket 名称。 + """ + prefix: Optional[str] = None + """ + TOS 前缀。 + """ diff --git a/src/arkruntime/types/session/environment_with_overrides.py b/src/arkruntime/types/session/environment_with_overrides.py new file mode 100644 index 0000000..f831cf0 --- /dev/null +++ b/src/arkruntime/types/session/environment_with_overrides.py @@ -0,0 +1,32 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Literal, Optional + +from arkruntime._models import BaseModel + +from .environment_config_override import EnvironmentConfigOverride + + +class EnvironmentWithOverrides(BaseModel): + """ + CreateSession 时的 Environment 覆写引用。 + """ + + type: Literal["environment_with_overrides"] + """ + 固定 `"environment_with_overrides"`。 + """ + id: str + """ + Base Environment ID。 + """ + config: Optional[EnvironmentConfigOverride] = None + """ + 运行时配置覆写。 + """ diff --git a/src/arkruntime/types/session/model_overrides.py b/src/arkruntime/types/session/model_overrides.py new file mode 100644 index 0000000..1b632d9 --- /dev/null +++ b/src/arkruntime/types/session/model_overrides.py @@ -0,0 +1,30 @@ +# Generated by datamodel-code-generator from ark-apis typespec/openapi. +# DO NOT EDIT — regenerate with `make gen-py-` in ark-apis. +# +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Optional + +from arkruntime._models import BaseModel + + +class ModelOverrides(BaseModel): + """ + Session 创建时允许临时覆写的模型运行参数。 + """ + + speed: Optional[str] = None + """ + 模型速度档位。 + """ + thinking: Optional[str] = None + """ + thinking 配置。 + """ + reasoning_effort: Optional[str] = None + """ + 推理努力程度。 + """ diff --git a/src/arkruntime/types/session/send_session_events_response.py b/src/arkruntime/types/session/send_session_events_response.py index bd4d02e..46cbfdf 100644 --- a/src/arkruntime/types/session/send_session_events_response.py +++ b/src/arkruntime/types/session/send_session_events_response.py @@ -6,7 +6,7 @@ from __future__ import annotations -from typing import Optional +from typing import Dict, List from arkruntime._models import BaseModel @@ -16,7 +16,7 @@ class SendSessionEventsResponse(BaseModel): SendSessionEvents 响应体(回执)。 """ - success: Optional[bool] = None + data: List[Dict[str, object]] """ - 是否成功接收(server 决定语义)。 + 服务端落库 / 转发完成后的事件回声。 """ diff --git a/src/arkruntime/types/session/session.py b/src/arkruntime/types/session/session.py index ea69740..93a24af 100644 --- a/src/arkruntime/types/session/session.py +++ b/src/arkruntime/types/session/session.py @@ -78,3 +78,7 @@ class Session(BaseModel): """ 资源标签。 """ + environment: Optional[Dict[str, object]] = None + """ + Session 创建时冻结的 Environment 快照。 + """ diff --git a/src/arkruntime/types/session/session_stream_shim.py b/src/arkruntime/types/session/session_stream_shim.py index 376b9b6..ffa2bda 100644 --- a/src/arkruntime/types/session/session_stream_shim.py +++ b/src/arkruntime/types/session/session_stream_shim.py @@ -538,3 +538,9 @@ def _decode_wire_events(cls, data: Any) -> Any: else: typed.append(item) return {"events": typed, "next_page": data.get("next_page")} + + @classmethod + def construct(cls, _fields_set: Any = None, **values: Any) -> "ListSessionEventsResponse": + """Decode the wire envelope when the SDK uses Pydantic's fast path.""" + decoded = cls._decode_wire_events(values) + return cls.model_construct(_fields_set=_fields_set, **decoded) diff --git a/src/arkruntime/types/session/session_thread_status.py b/src/arkruntime/types/session/session_thread_status.py index 2d3f189..9cd5398 100644 --- a/src/arkruntime/types/session/session_thread_status.py +++ b/src/arkruntime/types/session/session_thread_status.py @@ -17,4 +17,4 @@ class SessionThreadStatus(str, Enum): idle = "idle" running = "running" terminated = "terminated" - archived = "archived" + rescheduling = "rescheduling" diff --git a/src/arkruntime/types/skill/__init__.py b/src/arkruntime/types/skill/__init__.py index 4312b39..4088a7d 100644 --- a/src/arkruntime/types/skill/__init__.py +++ b/src/arkruntime/types/skill/__init__.py @@ -15,4 +15,4 @@ ] # Hand-written extras (preserved across regen via Makefile rsync --exclude=*_shim.py). -from ._init_extras_shim import * # noqa: F401,F403 +from ._init_extras_shim import * # noqa: F401,F403,E402 diff --git a/src/arkruntime/types/skill/create_skill_request.py b/src/arkruntime/types/skill/create_skill_request.py index 1dcac90..d19d8e0 100644 --- a/src/arkruntime/types/skill/create_skill_request.py +++ b/src/arkruntime/types/skill/create_skill_request.py @@ -16,3 +16,7 @@ class CreateSkillRequest(BaseModel): """ Skill 展示名。 """ + protection_enabled: Optional[bool] = None + """ + 是否启用 Skill 内容保护。 + """ diff --git a/src/arkruntime/types/skill/skill.py b/src/arkruntime/types/skill/skill.py index 525bd22..c220038 100644 --- a/src/arkruntime/types/skill/skill.py +++ b/src/arkruntime/types/skill/skill.py @@ -40,7 +40,23 @@ class Skill(BaseModel): """ 最新版本号。 """ - name: str + display_title: str """ - 人类可读名称。 + Skill 展示名。 + """ + source: str + """ + Skill 来源,例如 `custom` / `skill_hub` / `ark`。 + """ + updated_at: int + """ + 更新时间(Unix 秒)。 + """ + name: Optional[str] = None + """ + SKILL.md 中解析出的 name。 + """ + protection_enabled: Optional[bool] = None + """ + 是否启用内容保护。 """ diff --git a/tests/selfhosted/test_client.py b/tests/selfhosted/test_client.py new file mode 100644 index 0000000..3619f6d --- /dev/null +++ b/tests/selfhosted/test_client.py @@ -0,0 +1,310 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import httpx + +from arkruntime import Ark +from arkruntime.selfhosted import ClientAPI, SkillRef +from arkruntime.selfhosted.types import work_session_id + + +def _test_credential() -> str: + return "placeholder" + + +def _ark_client(transport: httpx.MockTransport, *, max_retries: int = 0) -> Ark: + return Ark( + api_key=_test_credential(), + base_url="https://ark.example.com/api/v3", + http_client=httpx.Client( + transport=transport, + headers={"Authorization": "Bearer inherited-test-key"}, + ), + max_retries=max_retries, + ) + + +def test_poll_work_preserves_nested_session_data() -> None: + work = { + "id": "sesn-20260814050521-zb4l4", + "created_at": "2026-08-14T05:05:21Z", + "data": { + "id": "sesn-20260814050521-zb4l4", + "type": "session", + }, + "environment_id": "env-20260813132936-wq8d4", + "state": "queued", + "type": "work", + } + + def handle(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/api/v3/environments/env-20260813132936-wq8d4/work/poll" + return httpx.Response(httpx.codes.OK, json=work) + + client = _ark_client(httpx.MockTransport(handle)) + try: + item = ClientAPI(client).poll_work("env-20260813132936-wq8d4") + finally: + client.close() + + assert item is not None + assert item.id == work["id"] + assert item.data.id == work["data"]["id"] + assert item.data.type == "session" + assert work_session_id(item) == work["data"]["id"] + + +def test_poll_work_rejects_payload_outside_generated_contract() -> None: + def handle(request: httpx.Request) -> httpx.Response: + return httpx.Response(httpx.codes.OK, json={"id": "work-without-required-fields"}) + + client = _ark_client(httpx.MockTransport(handle)) + try: + try: + ClientAPI(client).poll_work("env-1") + except Exception as exc: + assert "invalid WorkItem response" in str(exc) + else: + raise AssertionError("expected generated model validation failure") + finally: + client.close() + + +def test_heartbeat_uses_generated_response_model() -> None: + heartbeat = { + "last_heartbeat": "2026-08-24T10:00:00Z", + "lease_extended": True, + "state": "active", + "ttl_seconds": 30, + "type": "work_heartbeat", + } + + def handle(request: httpx.Request) -> httpx.Response: + return httpx.Response(httpx.codes.OK, json=heartbeat) + + client = _ark_client(httpx.MockTransport(handle)) + try: + response = ClientAPI(client).heartbeat_work( + "env-1", + "work-1", + expected_last_heartbeat="NO_HEARTBEAT", + desired_ttl_seconds=30, + ) + finally: + client.close() + + assert response.last_heartbeat == heartbeat["last_heartbeat"] + assert response.lease_extended is True + assert response.state == "active" + assert response.ttl_seconds == 30 + + +def test_atomic_environment_work_api_matches_openapi_contract() -> None: + work = { + "id": "work-1", + "created_at": "2026-08-24T10:00:00Z", + "data": {"id": "test", "type": "session"}, + "environment_id": "env-1", + "state": "active", + "type": "work", + } + heartbeat = { + "last_heartbeat": "2026-08-24T10:00:01Z", + "lease_extended": True, + "state": "active", + "ttl_seconds": 30, + "type": "work_heartbeat", + } + calls = [] + + def handle(request: httpx.Request) -> httpx.Response: + calls.append(request) + if request.url.path.endswith("/poll"): + assert request.method == "GET" + assert request.url.params["block_ms"] == "999" + assert request.url.params["reclaim_older_than_ms"] == "5000" + assert request.headers["Ark-Worker-ID"] == "worker-1" + return httpx.Response(httpx.codes.OK, json=work) + if request.url.path.endswith("/ack"): + assert request.method == "POST" + assert request.headers["Ark-Worker-ID"] == "worker-1" + return httpx.Response(httpx.codes.OK, json=work) + if request.url.path.endswith("/heartbeat"): + assert request.method == "POST" + assert request.url.params["expected_last_heartbeat"] == "NO_HEARTBEAT" + assert request.url.params["desired_ttl_seconds"] == "30" + return httpx.Response(httpx.codes.OK, json=heartbeat) + assert request.url.path.endswith("/stop") + assert request.method == "POST" + assert request.content == b"{}" + return httpx.Response(httpx.codes.OK, json=work) + + client = _ark_client(httpx.MockTransport(handle)) + try: + resource = client.environments.work + assert ( + resource.poll( + "env-1", + worker_id="worker-1", + block_ms=999, + reclaim_older_than_ms=5000, + ).id + == "work-1" + ) + assert resource.ack("env-1", "work-1", worker_id="worker-1").id == "work-1" + assert ( + resource.heartbeat( + "env-1", + "work-1", + expected_last_heartbeat="NO_HEARTBEAT", + desired_ttl_seconds=30, + ).lease_extended + is True + ) + assert resource.stop("env-1", "work-1").id == "work-1" + finally: + client.close() + + assert [request.url.path for request in calls] == [ + "/api/v3/environments/env-1/work/poll", + "/api/v3/environments/env-1/work/work-1/ack", + "/api/v3/environments/env-1/work/work-1/heartbeat", + "/api/v3/environments/env-1/work/work-1/stop", + ] + + +def test_list_events_decodes_wire_data_into_tool_use_events() -> None: + tool_use = { + "id": "call_test", + "type": "agent.tool_use", + "name": "bash", + "input": {"command": "python --version"}, + "session_thread_id": "sthr_test", + } + + def handle(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/api/v3/sessions/sesn_test/events" + assert request.url.params["limit"] == "100" + assert request.url.params["order"] == "asc" + return httpx.Response(httpx.codes.OK, json={"data": [tool_use]}) + + http_client = httpx.Client(transport=httpx.MockTransport(handle)) + client = Ark( + api_key=_test_credential(), + base_url="https://ark.example.com/api/v3", + http_client=http_client, + ) + try: + response = ClientAPI(client).list_events("sesn_test") + finally: + client.close() + + assert len(response.events) == 1 + event = response.events[0] + assert event.id == tool_use["id"] + assert event.type == tool_use["type"] + assert event.name == tool_use["name"] + assert event.input == tool_use["input"] + assert event.session_thread_id == tool_use["session_thread_id"] + + +def test_resolve_skill_uses_control_plane_metadata() -> None: + def handle(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + assert request.url.path == "/api/v3/skills/skill-1" + return httpx.Response( + httpx.codes.OK, + json={ + "id": "skill-1", + "object": "skill", + "created_at": 1786506774, + "name": "canonical-skill-name", + "latest_version": "1.0.0", + }, + ) + + http_client = httpx.Client(transport=httpx.MockTransport(handle)) + client = Ark( + api_key=_test_credential(), + base_url="https://ark.example.com/api/v3", + http_client=http_client, + ) + try: + resolved = ClientAPI(client).resolve_skill(SkillRef(skill_id="skill-1", type="skill_hub")) + finally: + client.close() + + assert resolved.name == "canonical-skill-name" + assert resolved.version == "1.0.0" + assert resolved.type == "skill_hub" + + +def test_open_skill_hub_resolves_metadata_and_downloads_version() -> None: + requests = [] + + def handle(request: httpx.Request) -> httpx.Response: + assert "authorization" not in request.headers + requests.append(request.url) + if request.url.path == "/v1/skills": + assert request.url.params["skillIds"] == "skill-1" + return httpx.Response( + httpx.codes.OK, + json={ + "Skills": [ + {"Id": "other-skill", "Slug": "wrong/slug"}, + {"Id": "skill-1", "Slug": "volcengine/ark/demo"}, + ], + "Total": 2, + }, + ) + assert request.url.path == "/v1/skills/download/volcengine/ark/demo" + assert request.url.params["version"] == "1.0.0" + return httpx.Response( + httpx.codes.OK, + content=b"skill-hub-zip", + headers={"Content-Type": "application/zip"}, + ) + + client = _ark_client(httpx.MockTransport(handle)) + try: + content = ClientAPI(client).open_skill( + "test", + SkillRef(type="skill_hub", skill_id="skill-1", version="1.0.0"), + ) + try: + assert content.body.read() == b"skill-hub-zip" + finally: + content.body.close() + finally: + client.close() + + assert len(requests) == 2 + + +def test_heartbeat_has_lease_bounded_timeout_and_no_hidden_retry() -> None: + requests = [] + + def handle(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.extensions["timeout"]["read"] == 15.0 + return httpx.Response(httpx.codes.INTERNAL_SERVER_ERROR, json={"error": "temporary"}) + + client = _ark_client(httpx.MockTransport(handle), max_retries=3) + try: + try: + ClientAPI(client).heartbeat_work( + "env-1", + "work-1", + expected_last_heartbeat="NO_HEARTBEAT", + desired_ttl_seconds=30, + ) + except Exception: + pass + else: + raise AssertionError("expected heartbeat failure") + finally: + client.close() + + assert len(requests) == 1 diff --git a/tests/selfhosted/test_envinit.py b/tests/selfhosted/test_envinit.py new file mode 100644 index 0000000..18d820d --- /dev/null +++ b/tests/selfhosted/test_envinit.py @@ -0,0 +1,157 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import io +import logging +import os +import zipfile + +from arkruntime.selfhosted import Initializer, InitializerOptions, Session, SkillRef +from arkruntime.selfhosted.envinit import _replace_skill_dir +from arkruntime.selfhosted.types import SkillContent + + +class FailingSkillAPI: + def open_skill(self, session_id: str, skill: SkillRef) -> None: + raise RuntimeError("content endpoint unavailable") + + +class ResolvingSkillAPI: + def __init__(self, archive: bytes) -> None: + self.archive = archive + + def resolve_skill(self, skill: SkillRef) -> SkillRef: + return SkillRef( + name="canonical-skill-name", + skill_id=skill.id_value(), + type=skill.type, + version=skill.version, + ) + + def open_skill(self, session_id: str, skill: SkillRef) -> SkillContent: + return SkillContent(body=io.BytesIO(self.archive), content_length=len(self.archive)) + + +class _CloseTrackingBody(io.BytesIO): + was_closed = False + + def close(self) -> None: + self.was_closed = True + super().close() + + +def test_setup_logs_skill_download_failure_and_continues(tmp_path, caplog) -> None: + logger = logging.getLogger("test.selfhosted.envinit") + session = Session.from_mapping( + { + "id": "session-1", + "skills": [ + { + "type": "skill_hub", + "skill_id": "skill-1", + "display_name": "demo", + "version": "2", + } + ], + } + ) + initializer = Initializer( + FailingSkillAPI(), + InitializerOptions(workdir=str(tmp_path), logger=logger), + ) + + with caplog.at_level(logging.WARNING, logger=logger.name): + initializer.setup(session) + + assert "failed to install skill" in caplog.text + assert "session_id=session-1" in caplog.text + assert "skill=demo" in caplog.text + assert "version=2" in caplog.text + assert "download skill demo: content endpoint unavailable" in caplog.text + + +def test_setup_installs_skill_under_resolved_metadata_name(tmp_path) -> None: + archive = io.BytesIO() + with zipfile.ZipFile(archive, "w") as output: + output.writestr("SKILL.md", "hello") + session = Session.from_mapping( + { + "id": "session-1", + "skills": [{"type": "custom", "skill_id": "skill-1", "version": "1"}], + } + ) + initializer = Initializer( + ResolvingSkillAPI(archive.getvalue()), + InitializerOptions(workdir=str(tmp_path)), + ) + + initializer.setup(session) + + assert (tmp_path / "skills" / "canonical-skill-name" / "SKILL.md").read_text() == "hello" + assert not (tmp_path / "skills" / "skill-1").exists() + + +def test_install_closes_skill_body_when_archive_copy_fails(tmp_path) -> None: + body = _CloseTrackingBody(b"archive-too-large") + api = ResolvingSkillAPI(b"") + api.open_skill = lambda session_id, skill: SkillContent(body=body, content_length=17) + initializer = Initializer( + api, + InitializerOptions(workdir=str(tmp_path), max_archive_bytes=1), + ) + + try: + initializer.install_skill("session-1", SkillRef(skill_id="skill-1", version="1")) + except ValueError as exc: + assert "archive too large" in str(exc) + else: + raise AssertionError("expected archive size failure") + + assert body.was_closed + + +def test_zip_archive_entry_limit_is_enforced(tmp_path) -> None: + archive = io.BytesIO() + with zipfile.ZipFile(archive, "w") as output: + output.writestr("one", "1") + output.writestr("two", "2") + initializer = Initializer( + ResolvingSkillAPI(archive.getvalue()), + InitializerOptions(workdir=str(tmp_path), max_archive_entries=1), + ) + + try: + initializer.install_skill("session-1", SkillRef(skill_id="skill-1", version="1")) + except ValueError as exc: + assert "too many entries" in str(exc) + else: + raise AssertionError("expected archive entry limit failure") + + +def test_replace_skill_rolls_back_old_version_when_commit_fails(tmp_path, monkeypatch) -> None: + source = tmp_path / "new-skill" + target = tmp_path / "installed-skill" + source.mkdir() + target.mkdir() + (source / "marker.txt").write_text("new") + (target / "marker.txt").write_text("old") + real_replace = os.replace + + def fail_new_commit(from_path, to_path): + if os.fspath(from_path) == os.fspath(source) and os.fspath(to_path) == os.fspath(target): + raise OSError("commit failed") + real_replace(from_path, to_path) + + monkeypatch.setattr("arkruntime.selfhosted.envinit.os.replace", fail_new_commit) + + try: + _replace_skill_dir(source, target) + except OSError as exc: + assert "commit failed" in str(exc) + else: + raise AssertionError("expected skill commit failure") + + assert (target / "marker.txt").read_text() == "old" + assert (source / "marker.txt").read_text() == "new" diff --git a/tests/selfhosted/test_session_tool_runner.py b/tests/selfhosted/test_session_tool_runner.py new file mode 100644 index 0000000..94208ac --- /dev/null +++ b/tests/selfhosted/test_session_tool_runner.py @@ -0,0 +1,205 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import threading +import time + +import pytest + +from arkruntime.selfhosted import Event, ListEventsResponse, SessionToolRunner, SessionToolRunnerOptions +from arkruntime.selfhosted.tools import FunctionTool, ToolContext, ToolSet + + +class _ListAPI: + def __init__(self) -> None: + self.calls = 0 + self.created_at_values = [] + self.sent = [] + self.sent_event = threading.Event() + + def list_events(self, session_id, **kwargs): + self.calls += 1 + self.created_at_values.append(kwargs.get("created_at_gt")) + if self.calls == 1: + raise RuntimeError("temporary list failure") + return ListEventsResponse( + events=[ + Event( + id="event-1", + type="agent.custom_tool_use", + name="custom", + custom_tool_use_id="call-1", + session_thread_id="thread-1", + input={}, + ) + ] + ) + + def send_event(self, session_id, event): + self.sent.append(event) + self.sent_event.set() + + +def test_list_fallback_retries_full_history_and_converts_custom_tool_error(tmp_path) -> None: + api = _ListAPI() + + def fail(_input, _context): + raise RuntimeError("custom tool failed") + + runner = SessionToolRunner( + api, + "session-1", + SessionToolRunnerOptions( + tools=ToolSet(), + tool_context=ToolContext(workdir=str(tmp_path)), + custom_tools={"custom": FunctionTool("custom", fail)}, + prefer_stream=False, + event_poll_interval_seconds=0.01, + ), + ) + thread = threading.Thread(target=runner.run) + thread.start() + assert api.sent_event.wait(2) + runner.close() + thread.join(timeout=2) + + assert not thread.is_alive() + assert api.calls >= 2 + assert all(value is None for value in api.created_at_values) + assert len(api.sent) == 1 + assert api.sent[0].is_error is True + assert api.sent[0].content[0].text == "custom tool failed" + + +def test_runner_requires_tool_context() -> None: + with pytest.raises(ValueError, match="tool context"): + SessionToolRunner(object(), "session-1", SessionToolRunnerOptions(tools=ToolSet())) + + +def test_confirmed_tool_is_released_once_when_post_fails(tmp_path) -> None: + executions = [] + + def execute(_input, _context): + executions.append(True) + return ToolSet().execute("missing", {}, _context) + + runner = SessionToolRunner( + object(), + "session-1", + SessionToolRunnerOptions( + tools=ToolSet([FunctionTool("count", execute)]), + tool_context=ToolContext(workdir=str(tmp_path)), + ), + ) + tool_use = Event( + id="event-tool-use", + type="agent.tool_use", + name="count", + tool_use_id="call-1", + evaluated_permission="ask", + input={}, + ) + confirmation = Event(type="user.tool_confirmation", tool_use_id="call-1", result="allow") + runner._state.pending_ask["call-1"] = tool_use + runner._state.confirmations["call-1"] = confirmation + runner._state.retry_send_event = lambda _event, _call_id: False + + runner._state.release_confirmed_tool_uses() + runner._state.release_confirmed_tool_uses() + + assert len(executions) == 1 + assert "call-1" not in runner._state.pending_ask + + +def test_duplicate_stream_idle_event_does_not_reset_idle_deadline(tmp_path) -> None: + runner = SessionToolRunner( + object(), + "session-1", + SessionToolRunnerOptions(tools=ToolSet(), tool_context=ToolContext(workdir=str(tmp_path))), + ) + event = Event( + id="idle-1", + type="session.status_idle", + stop_reason={"type": "end_turn"}, + ) + + runner._state.handle_stream_event(event) + armed_at = runner._state.idle_armed_at + time.sleep(0.001) + runner._state.handle_stream_event(event) + + assert runner._state.idle_armed_at == armed_at + + +def test_reconcile_does_not_reset_idle_deadline_for_seen_history(tmp_path) -> None: + runner = SessionToolRunner( + object(), + "session-1", + SessionToolRunnerOptions(tools=ToolSet(), tool_context=ToolContext(workdir=str(tmp_path))), + ) + event = Event( + id="idle-1", + type="session.status_idle", + stop_reason={"type": "end_turn"}, + ) + + runner._state.process_listed_events([event], reconcile=True) + armed_at = runner._state.idle_armed_at + time.sleep(0.001) + runner._state.process_listed_events([event], reconcile=True) + + assert runner._state.idle_armed_at == armed_at + + +def test_tool_execution_copies_context_and_preserves_configured_timeout(tmp_path) -> None: + contexts = [] + + def capture(_input, context): + contexts.append(context) + return ToolSet().execute("missing", {}, context) + + original = ToolContext(workdir=str(tmp_path), tool_timeout_seconds=7) + runner = SessionToolRunner( + object(), + "session-1", + SessionToolRunnerOptions( + tools=ToolSet([FunctionTool("capture", capture)]), + tool_context=original, + ), + ) + event = Event(id="tool-1", type="agent.tool_use", name="capture", tool_use_id="call-1", input={}) + + runner._state.execute_tool(event, custom=False) + + assert contexts[0] is not original + assert contexts[0].tool_timeout_seconds == 7 + assert original.tool_timeout_seconds == 7 + + +def test_successful_send_stays_answered_when_mark_sent_fails(tmp_path, caplog) -> None: + class FailingStore: + def mark_sent(self, _call_id): + raise OSError("ledger unavailable") + + runner = SessionToolRunner( + object(), + "session-1", + SessionToolRunnerOptions( + tools=ToolSet(), + tool_context=ToolContext(workdir=str(tmp_path)), + result_store=FailingStore(), + ), + ) + source = Event(id="tool-1", type="agent.tool_use", name="bash", tool_use_id="call-1") + result = Event(id="result-1", type="user.tool_result", tool_use_id="call-1") + runner._state.pending_results["call-1"] = result + runner._state.retry_send_event = lambda _event, _call_id: True + + with caplog.at_level("WARNING"): + runner._state.send_result("call-1", source, False, "", result) + + assert runner._state.answered["call-1"] is True + assert "call-1" not in runner._state.pending_results + assert "mark tool result sent failed" in caplog.text diff --git a/tests/selfhosted/test_tool_result_store.py b/tests/selfhosted/test_tool_result_store.py new file mode 100644 index 0000000..5bf592d --- /dev/null +++ b/tests/selfhosted/test_tool_result_store.py @@ -0,0 +1,23 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from arkruntime.selfhosted import Event, FileToolResultStore + + +def test_recovery_uses_persisted_call_id(tmp_path) -> None: + store = FileToolResultStore(str(tmp_path)) + store.begin("call-1", Event(id="event-1", type="agent.tool_use", name="bash")) + + pending, _ = store.recover() + + assert pending["call-1"].tool_use_id == "call-1" + + +def test_recovery_removes_stale_temporary_records(tmp_path) -> None: + store = FileToolResultStore(str(tmp_path)) + stale = store.dir / ".tool-result-stale.tmp" + stale.write_text("partial") + + store.recover() + + assert not stale.exists() diff --git a/tests/selfhosted/test_tools.py b/tests/selfhosted/test_tools.py new file mode 100644 index 0000000..de8e5fa --- /dev/null +++ b/tests/selfhosted/test_tools.py @@ -0,0 +1,107 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import threading +import time + +import pytest + +from arkruntime.selfhosted.tools import ( + BashTool, + EditFileTool, + GrepTool, + ToolContext, + WriteFileTool, + _is_sensitive_env_key, +) + + +def _text(result) -> str: + return "".join(block.text for block in result.content) + + +def _env_name(*parts: str) -> str: + return "_".join(parts) + + +def test_bash_scrubs_inherited_and_explicit_credentials(tmp_path, monkeypatch) -> None: + monkeypatch.setenv(_env_name("ARK", "API", "KEY"), "redacted") + monkeypatch.setenv("SAFE_INHERITED", "safe") + context = ToolContext( + workdir=str(tmp_path), + env={ + _env_name("ARK", "API", "KEY"): "redacted", + "SAFE_EXPLICIT": "ok", + }, + ) + + result = BashTool().execute( + {"command": ('printf \'%s/%s/%s\' "${ARK_API_KEY-unset}" "$SAFE_INHERITED" "$SAFE_EXPLICIT"')}, + context, + ) + + assert not result.is_error + assert _text(result) == "unset//ok" + + +def test_bash_scrubs_extended_credential_names() -> None: + assert _is_sensitive_env_key(_env_name("AIME", "SESSION")) + assert _is_sensitive_env_key(_env_name("X", "CODE", "AUTH")) + assert _is_sensitive_env_key(_env_name("GITHUB", "JWT")) + assert _is_sensitive_env_key(_env_name("GITHUB", "PAT")) + assert not _is_sensitive_env_key("SAFE_VALUE") + + +def test_bash_honors_worker_cancellation(tmp_path) -> None: + canceled = threading.Event() + context = ToolContext(workdir=str(tmp_path), cancel_event=canceled, tool_timeout_seconds=10) + timer = threading.Timer(0.1, canceled.set) + timer.start() + started = time.monotonic() + try: + result = BashTool().execute({"command": "sleep 10"}, context) + finally: + timer.cancel() + + assert time.monotonic() - started < 2 + assert result.is_error + assert "canceled" in _text(result) + + +def test_grep_skips_symlink_that_escapes_workdir(tmp_path) -> None: + outside = tmp_path.parent / f"{tmp_path.name}-outside" + outside.mkdir() + (outside / "secret.txt").write_text("SELFHOST_SECRET_MARKER\n") + (tmp_path / "escape.txt").symlink_to(outside / "secret.txt") + + result = GrepTool().execute( + {"path": ".", "pattern": "SELFHOST_SECRET_MARKER"}, + ToolContext(workdir=str(tmp_path)), + ) + + assert not result.is_error + assert "SELFHOST_SECRET_MARKER" not in _text(result) + + +@pytest.mark.parametrize( + "tool, tool_input", + [ + (WriteFileTool(), {"path": "example.txt", "content": "new"}), + (EditFileTool(), {"path": "example.txt", "old_string": "old", "new_string": "new"}), + ], +) +def test_file_mutation_keeps_old_content_when_atomic_replace_fails(tmp_path, monkeypatch, tool, tool_input) -> None: + target = tmp_path / "example.txt" + target.write_text("old") + + def fail_replace(_source, _target): + raise OSError("replace failed") + + monkeypatch.setattr("arkruntime.selfhosted.tools.os.replace", fail_replace) + + with pytest.raises(OSError, match="replace failed"): + tool.execute(tool_input, ToolContext(workdir=str(tmp_path))) + + assert target.read_text() == "old" diff --git a/tests/selfhosted/test_worker.py b/tests/selfhosted/test_worker.py new file mode 100644 index 0000000..0456a3c --- /dev/null +++ b/tests/selfhosted/test_worker.py @@ -0,0 +1,163 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time + +import pytest + +from arkruntime.selfhosted import ( + APIError, + EnvironmentWorker, + EnvironmentWorkerOptions, + HandleItemOptions, + HeartbeatResponse, + Session, + WorkItem, + WorkPoller, + WorkPollerOptions, +) +from arkruntime.selfhosted.types import WorkData, is_fatal_4xx + + +def _work_item(environment_id: str) -> WorkItem: + return WorkItem( + id="work-1", + created_at="2026-08-24T10:00:00Z", + environment_id=environment_id, + data=WorkData(type="session", id="session-1"), + state="queued", + type="work", + ) + + +class _PollAPI: + def __init__(self, ack_error: BaseException) -> None: + self.ack_error = ack_error + self.polls = 0 + self.stops = [] + + def poll_work(self, environment_id, **kwargs): + self.polls += 1 + if self.polls > 1: + return None + return _work_item(environment_id) + + def ack_work(self, environment_id, work_id, **kwargs): + raise self.ack_error + + def stop_work(self, environment_id, work_id, **kwargs): + self.stops.append((environment_id, work_id, kwargs)) + + +def test_ack_conflict_does_not_stop_unowned_work() -> None: + api = _PollAPI(APIError(409, "already claimed")) + poller = WorkPoller(api, WorkPollerOptions(environment_id="env-1", drain=True)) + + assert poller.next() is None + assert poller.error is None + assert api.stops == [] + + +def test_fatal_ack_error_stops_poller_without_stopping_work() -> None: + error = APIError(403, "forbidden") + api = _PollAPI(error) + poller = WorkPoller(api, WorkPollerOptions(environment_id="env-1", drain=True)) + + assert poller.next() is None + assert poller.error is error + assert api.stops == [] + + +class _StoppingHeartbeatAPI: + def __init__(self) -> None: + self.stops = [] + + def heartbeat_work(self, environment_id, work_id, **kwargs): + return HeartbeatResponse( + last_heartbeat="2026-08-24T10:00:00Z", + state="stopping", + lease_extended=True, + ttl_seconds=30, + type="work_heartbeat", + ) + + def get_session(self, session_id): + return Session(id=session_id) + + def list_events(self, session_id, **kwargs): + raise AssertionError("runner must observe the heartbeat stop before listing events") + + def stop_work(self, environment_id, work_id, **kwargs): + self.stops.append((environment_id, work_id, kwargs)) + + +def test_heartbeat_stop_cancels_session_runner(tmp_path) -> None: + api = _StoppingHeartbeatAPI() + worker = EnvironmentWorker( + api, + EnvironmentWorkerOptions(environment_id="env-1", workdir=str(tmp_path)), + ) + + started = time.monotonic() + worker.handle_item(HandleItemOptions(work_id="work-1", environment_id="env-1", session_id="session-1")) + + assert time.monotonic() - started < 1 + assert api.stops == [("env-1", "work-1", {"force": True})] + + +class _LeaseLostHeartbeatAPI(_StoppingHeartbeatAPI): + def heartbeat_work(self, environment_id, work_id, **kwargs): + raise APIError(412, "lease lost") + + +def test_lease_lost_does_not_stop_work(tmp_path) -> None: + api = _LeaseLostHeartbeatAPI() + worker = EnvironmentWorker( + api, + EnvironmentWorkerOptions(environment_id="env-1", workdir=str(tmp_path)), + ) + + worker.handle_item(HandleItemOptions(work_id="work-1", environment_id="env-1", session_id="session-1")) + + assert api.stops == [] + + +class _ClaimedPollAPI: + def __init__(self) -> None: + self.stops = [] + + def poll_work(self, environment_id, **kwargs): + return _work_item(environment_id) + + def ack_work(self, environment_id, work_id, **kwargs): + return None + + def stop_work(self, environment_id, work_id, **kwargs): + self.stops.append((environment_id, work_id, kwargs)) + + +@pytest.mark.parametrize("auto_stop, expected_stops", [(True, 1), (False, 0)]) +def test_poller_auto_stop_is_configurable(auto_stop, expected_stops) -> None: + api = _ClaimedPollAPI() + poller = WorkPoller(api, WorkPollerOptions(environment_id="env-1", auto_stop=auto_stop)) + + assert poller.next() is not None + poller.close() + + assert len(api.stops) == expected_stops + + +def test_session_id_cannot_escape_worker_root(tmp_path) -> None: + worker = EnvironmentWorker(object(), EnvironmentWorkerOptions(workdir=str(tmp_path))) + + workdir = worker._workdir_for("../../outside", use_workdir_as_session=False) + + assert str(tmp_path.resolve()) in workdir + assert ".." not in workdir + + +@pytest.mark.parametrize("status_code", [408, 409, 412, 429]) +def test_recoverable_client_status_is_not_fatal(status_code) -> None: + assert not is_fatal_4xx(APIError(status_code, "recoverable")) diff --git a/tests/test_client_base_url.py b/tests/test_client_base_url.py new file mode 100644 index 0000000..321addb --- /dev/null +++ b/tests/test_client_base_url.py @@ -0,0 +1,29 @@ +# Copyright (c) 2026 ByteDance Ltd. and/or its affiliates. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from arkruntime import Ark + + +def _test_credential() -> str: + return "placeholder" + + +def test_client_uses_production_base_url_by_default() -> None: + client = Ark(api_key=_test_credential()) + try: + assert str(client._base_url).rstrip("/") == "https://ark.cn-beijing.volces.com/api/v3" + finally: + client.close() + + +def test_client_accepts_base_url_override() -> None: + client = Ark( + api_key=_test_credential(), + base_url="https://example.com/api/v3", + ) + try: + assert str(client._base_url).rstrip("/") == "https://example.com/api/v3" + finally: + client.close()