diff --git a/agent/integrations/langsmith.py b/agent/integrations/langsmith.py index e8adc694..4c41240f 100644 --- a/agent/integrations/langsmith.py +++ b/agent/integrations/langsmith.py @@ -4,6 +4,7 @@ from __future__ import annotations import asyncio import base64 +import json import logging import os import time @@ -87,6 +88,43 @@ def _get_sandbox_snapshot_config() -> tuple[str | None, int, int, int, int, int] ) +def _get_sandbox_create_extra_fields() -> dict[str, Any]: + """Parse SANDBOX_CREATE_EXTRA_JSON into extra fields merged into the + sandbox-create request body, e.g. ``{"_internal_runtime": "v2"}``.""" + raw = os.environ.get("SANDBOX_CREATE_EXTRA_JSON") + if not raw or not raw.strip(): + return {} + try: + parsed = json.loads(raw) + except json.JSONDecodeError as e: + msg = f"SANDBOX_CREATE_EXTRA_JSON must be valid JSON, got {raw!r}" + raise ValueError(msg) from e + if not isinstance(parsed, dict): + msg = f"SANDBOX_CREATE_EXTRA_JSON must be a JSON object, got {type(parsed).__name__}" + raise ValueError(msg) + return parsed + + +def _install_create_extra_fields(client: SandboxClient, extra: dict[str, Any]) -> None: + """Merge ``extra`` into the JSON body of the sandbox-create request. + + The SDK's ``create_sandbox`` builds a fixed payload with no passthrough, so + wrap the HTTP client's ``post`` to inject the fields on the ``POST /boxes`` + request only (other endpoints post to ``/boxes/{name}/...``). + """ + if not extra: + return + original_post = client._http.post + + def post_with_extra(url: Any, *args: Any, **kwargs: Any) -> Any: + payload = kwargs.get("json") + if str(url).endswith("/boxes") and isinstance(payload, dict): + kwargs["json"] = {**payload, **extra} + return original_post(url, *args, **kwargs) + + client._http.post = post_with_extra + + def _github_proxy_rules(github_token: str) -> list[dict[str, Any]]: basic_auth = base64.b64encode(f"x-access-token:{github_token}".encode()).decode() return [ @@ -453,6 +491,7 @@ class LangSmithProvider(SandboxProvider): ): msg = f"{name} must be >= 0, got {value}" raise ValueError(msg) + _get_sandbox_create_extra_fields() def get_or_create( self, @@ -483,6 +522,8 @@ class LangSmithProvider(SandboxProvider): msg = "DEFAULT_SANDBOX_SNAPSHOT_ID must be set when SANDBOX_TYPE=langsmith" raise ValueError(msg) + _install_create_extra_fields(self._client, _get_sandbox_create_extra_fields()) + try: sandbox = self._client.create_sandbox( snapshot_id=snapshot_id, diff --git a/tests/sandbox/test_langsmith_sandbox_config.py b/tests/sandbox/test_langsmith_sandbox_config.py index 37113cb9..a7a97ae4 100644 --- a/tests/sandbox/test_langsmith_sandbox_config.py +++ b/tests/sandbox/test_langsmith_sandbox_config.py @@ -11,7 +11,9 @@ from agent.integrations.langsmith import ( DEFAULT_SANDBOX_VCPUS, DEFAULT_SNAPSHOT_FS_CAPACITY_BYTES, LangSmithProvider, + _get_sandbox_create_extra_fields, _get_sandbox_snapshot_config, + _install_create_extra_fields, ) @@ -97,3 +99,67 @@ def test_validate_startup_accepts_valid_config() -> None: clear=True, ): LangSmithProvider.validate_startup_config() + + +def test_extra_fields_unset_is_empty() -> None: + with patch.dict("os.environ", {}, clear=True): + assert _get_sandbox_create_extra_fields() == {} + with patch.dict("os.environ", {"SANDBOX_CREATE_EXTRA_JSON": " "}, clear=True): + assert _get_sandbox_create_extra_fields() == {} + + +def test_extra_fields_parsed() -> None: + with patch.dict( + "os.environ", + {"SANDBOX_CREATE_EXTRA_JSON": '{"_internal_runtime": "v2"}'}, + clear=True, + ): + assert _get_sandbox_create_extra_fields() == {"_internal_runtime": "v2"} + + +def test_extra_fields_rejects_invalid_json() -> None: + with patch.dict("os.environ", {"SANDBOX_CREATE_EXTRA_JSON": "{not json"}, clear=True): + with pytest.raises(ValueError, match="valid JSON"): + _get_sandbox_create_extra_fields() + + +def test_extra_fields_rejects_non_object() -> None: + with patch.dict("os.environ", {"SANDBOX_CREATE_EXTRA_JSON": "[1, 2]"}, clear=True): + with pytest.raises(ValueError, match="JSON object"): + _get_sandbox_create_extra_fields() + + +def test_install_create_extra_fields_merges_only_boxes_post() -> None: + calls: list[tuple[str, dict]] = [] + + class _FakeHttp: + def post(self, url, **kwargs): # noqa: ANN001, ANN003 + calls.append((url, kwargs.get("json"))) + return "ok" + + class _FakeClient: + def __init__(self) -> None: + self._http = _FakeHttp() + + client = _FakeClient() + _install_create_extra_fields(client, {"_internal_runtime": "v2"}) + + client._http.post("https://api/v2/sandboxes/boxes", json={"snapshot_id": "s"}) + client._http.post("https://api/v2/sandboxes/boxes/abc/start", json={"foo": "bar"}) + + assert calls[0][1] == {"snapshot_id": "s", "_internal_runtime": "v2"} + assert calls[1][1] == {"foo": "bar"} + + +def test_install_create_extra_fields_noop_when_empty() -> None: + class _FakeHttp: + def __init__(self) -> None: + self.post = "sentinel" + + class _FakeClient: + def __init__(self) -> None: + self._http = _FakeHttp() + + client = _FakeClient() + _install_create_extra_fields(client, {}) + assert client._http.post == "sentinel"