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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 23 additions & 4 deletions src/art/tinker/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,24 @@
from .backend import TinkerBackend
from .renderers import get_renderer_name
from .server import OpenAICompatibleTinkerServer
"""Tinker integrations; importing the inference client needs no training extras."""

__all__ = ["TinkerBackend", "get_renderer_name", "OpenAICompatibleTinkerServer"]
from importlib import import_module
from typing import TYPE_CHECKING, Any

if TYPE_CHECKING:
from .backend import TinkerBackend
from .renderers import get_renderer_name
from .server import OpenAICompatibleTinkerServer

_EXPORTS = {
"TinkerBackend": ".backend",
"get_renderer_name": ".renderers",
"OpenAICompatibleTinkerServer": ".server",
}
__all__ = list(_EXPORTS)


def __getattr__(name: str) -> Any:
if name not in _EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
value = getattr(import_module(_EXPORTS[name], __name__), name)
globals()[name] = value
return value
15 changes: 15 additions & 0 deletions src/art/trajectories/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,21 @@ class ResponsesExchange(_Exchange):
request: _Preserved[ResponsesRequest]
response: Response

@pydantic.field_serializer("response")
def serialize_response(
self, response: Response, info: pydantic.SerializationInfo
) -> dict[str, Any]:
# Provider omissions and explicit nulls must survive compact/full replay.
return response.model_dump(
mode=info.mode,
include=info.include,
exclude=info.exclude,
context=info.context,
by_alias=info.by_alias,
exclude_unset=True,
exclude_none=info.exclude_none,
)

@pydantic.computed_field
@property
def model(self) -> str | None:
Expand Down
66 changes: 66 additions & 0 deletions tests/unit/test_tinker_import_boundary.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""Exercise the package initializer without importing optional training extras."""

import builtins
from importlib.util import module_from_spec, spec_from_file_location
from pathlib import Path
import sys
from types import SimpleNamespace
import unittest
from unittest.mock import Mock, patch


class TinkerImportBoundaryTests(unittest.TestCase):
def load_package(self):
path = Path(__file__).parents[2] / "src/art/tinker/__init__.py"
spec = spec_from_file_location("_tested_tinker", path)
assert spec is not None and spec.loader is not None
module = module_from_spec(spec)
original_import = builtins.__import__

def import_without_training(name, *args, **kwargs):
level = args[3] if len(args) > 3 else kwargs.get("level", 0)
if level or name.split(".")[0] not in sys.stdlib_module_names:
raise AssertionError(f"Eager optional dependency: {name}")
return original_import(name, *args, **kwargs)

with patch("builtins.__import__", side_effect=import_without_training):
spec.loader.exec_module(module)
return module

def test_import_does_not_load_training_exports(self):
package = self.load_package()
self.assertEqual(
package.__all__,
["TinkerBackend", "get_renderer_name", "OpenAICompatibleTinkerServer"],
)
self.assertFalse(set(package.__all__) & package.__dict__.keys())

def test_exports_resolve_original_objects_once_on_demand(self):
package = self.load_package()
objects = {name: object() for name in package.__all__}
package.import_module = Mock(return_value=SimpleNamespace(**objects))
for name, target in zip(package.__all__, (".backend", ".renderers", ".server")):
self.assertIs(getattr(package, name), objects[name])
self.assertIs(getattr(package, name), objects[name])
package.import_module.assert_called_once_with(target, "_tested_tinker")
package.import_module.reset_mock()

def test_unknown_attribute_does_not_load_dependencies(self):
package = self.load_package()
package.import_module = Mock()
with self.assertRaises(AttributeError):
getattr(package, "missing")
package.import_module.assert_not_called()

def test_requested_export_preserves_missing_dependency_error(self):
package = self.load_package()
error = ModuleNotFoundError("missing training dependency", name="mp_actors")
package.import_module = Mock(side_effect=error)
with self.assertRaises(ModuleNotFoundError) as raised:
getattr(package, "TinkerBackend")
self.assertIs(raised.exception, error)
self.assertNotIn("TinkerBackend", package.__dict__)


if __name__ == "__main__":
unittest.main()
175 changes: 175 additions & 0 deletions tests/unit/trajectories/test_responses_serialization.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
from datetime import UTC, datetime
import json
from typing import Any, Literal

from openai.types.chat import ChatCompletion
from openai.types.responses import Response
import pytest

from art.trajectories import (
ChatCompletionsExchange,
ResponsesExchange,
Trajectory,
)


def _response(status: str) -> dict[str, Any]:
reasoning: dict[str, Any] = {
"type": "reasoning",
"id": "rs_1",
"summary": [],
"content": [],
"encrypted_content": "opaque-reasoning",
"provider_metadata": {"trace": None},
}
if status != "absent":
reasoning["status"] = None if status == "null" else status
return {
"id": "resp_1",
"object": "response",
"created_at": 1,
"model": "behavior",
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
"metadata": None,
"output": [
reasoning,
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "lookup",
"arguments": '{"id":"item_1"}',
"status": "completed",
},
],
}


def _trajectory(raw: dict[str, Any]) -> Trajectory:
now = datetime.now(UTC)
trajectory = Trajectory()
trajectory.exchanges.responses.append(
ResponsesExchange(
request={"model": "behavior", "input": "Find the item."},
response=Response.model_validate(raw),
start_time=now,
end_time=now,
)
)
return trajectory


def _dump(
trajectory: Trajectory,
mode: Literal["python", "json", "json_string"],
**kwargs: Any,
) -> dict[str, Any]:
if mode == "json_string":
return json.loads(trajectory.model_dump_json(**kwargs))
return trajectory.model_dump(mode=mode, **kwargs)


@pytest.mark.parametrize("status", ["absent", "null", "completed"])
@pytest.mark.parametrize("compact", [True, False])
@pytest.mark.parametrize("mode", ["python", "json", "json_string"])
def test_responses_round_trip_preserves_provider_fields(
status: str, compact: bool, mode: Literal["python", "json", "json_string"]
) -> None:
raw = _response(status)
dumped = _dump(_trajectory(raw), mode, exclude_defaults=compact)
assert dumped["exchanges"]["responses"][0]["response"] == raw

restored = Trajectory.model_validate(dumped)
replayed = restored.exchanges.responses[0].response.model_copy(deep=True)
assert replayed.model_dump(mode="json", exclude_unset=True) == raw


@pytest.mark.parametrize("mode", ["python", "json", "json_string"])
def test_responses_nested_include_and_exclude(
mode: Literal["python", "json", "json_string"],
) -> None:
trajectory = _trajectory(_response("null"))
included = _dump(
trajectory,
mode,
include={
"exchanges": {
"responses": {
0: {
"response": {
"output": {0: {"id", "status", "encrypted_content"}}
}
}
}
}
},
exclude={
"exchanges": {"responses": {0: {"response": {"output": {0: {"id"}}}}}}
},
)
assert included == {
"exchanges": {
"responses": [
{
"response": {
"output": [
{"status": None, "encrypted_content": "opaque-reasoning"}
]
}
}
]
}
}
excluded = _dump(trajectory, mode, exclude_none=True)
response = excluded["exchanges"]["responses"][0]["response"]
assert "metadata" not in response
assert "status" not in response["output"][0]


def test_legacy_responses_tapes_keep_explicit_nulls() -> None:
raw = Response.model_validate(_response("absent")).model_dump(mode="json")
assert raw["output"][0]["status"] is None
restored = Trajectory.model_validate_json(
_trajectory(raw).model_dump_json(exclude_defaults=False)
)
assert _dump(restored, "json")["exchanges"]["responses"][0]["response"] == raw


@pytest.mark.parametrize("compact", [True, False])
@pytest.mark.parametrize("mode", ["python", "json", "json_string"])
def test_chat_serialization_is_unchanged(
compact: bool, mode: Literal["python", "json", "json_string"]
) -> None:
now = datetime.now(UTC)
response = ChatCompletion.model_validate(
{
"id": "chat_1",
"object": "chat.completion",
"created": 1,
"model": "behavior",
"choices": [
{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": "Done."},
}
],
}
)
trajectory = Trajectory()
trajectory.exchanges.chat_completions.append(
ChatCompletionsExchange(
request={"model": "behavior", "messages": []},
response=response,
start_time=now,
end_time=now,
)
)
expected = response.model_dump(
mode="python" if mode == "python" else "json",
exclude_defaults=compact if mode == "python" else False,
)
dumped = _dump(trajectory, mode, exclude_defaults=compact)
assert dumped["exchanges"]["chat_completions"][0]["response"] == expected
Loading