Skip to content
26 changes: 18 additions & 8 deletions src/roe/api/agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -663,6 +663,9 @@ def run_many(

Set ``skip_cache=True`` to bypass the job-result cache and force
fresh runs (the fresh results still refresh the cache).

If a chunk fails, the raised exception's ``submitted_job_ids`` lists
the jobs already started by earlier chunks.
"""
all_job_ids: list[str] = []
is_first_chunk = True
Expand All @@ -675,14 +678,21 @@ def run_many(
body = AgentRunAsyncManyRequest(inputs=[_build_aer(item) for item in chunk])
if metadata is not None:
body.additional_properties["metadata"] = metadata
response = request_raw(
self._raw,
agents_run_async_many,
UUID(str(agent_id)),
body=body,
organization_id=self._org_id,
extra_headers=_build_run_headers(skip_cache=skip_cache),
)
try:
response = request_raw(
self._raw,
agents_run_async_many,
UUID(str(agent_id)),
body=body,
organization_id=self._org_id,
extra_headers={
**(_build_run_headers(skip_cache=skip_cache) or {}),
"x-roe-skip-retry": "1",
},
)
except Exception as exc:
exc.submitted_job_ids = all_job_ids
raise
chunk_ids = response.json()
if not isinstance(chunk_ids, list):
raise RoeAPIException(
Expand Down
4 changes: 3 additions & 1 deletion src/roe/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,9 @@ def get_headers(self) -> dict[str, str]:
Returns:
Dictionary of headers including Authorization.
"""
from roe import __version__

return {
"Authorization": f"Bearer {self.config.api_key}",
"User-Agent": "roe-python/0.1.0",
"User-Agent": f"roe-python/{__version__}",
}
6 changes: 5 additions & 1 deletion src/roe/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,11 @@ def from_env(
organization_id = organization_id or os.getenv("ROE_ORGANIZATION_ID")
base_url = base_url or os.getenv("ROE_BASE_URL", "https://api.roe-ai.com")
timeout = timeout or float(os.getenv("ROE_TIMEOUT", "60.0"))
max_retries = max_retries or int(os.getenv("ROE_MAX_RETRIES", "3"))
max_retries = (
max_retries
if max_retries is not None
else int(os.getenv("ROE_MAX_RETRIES", "3"))
)
batch_chunk_delay = (
batch_chunk_delay
if batch_chunk_delay is not None
Expand Down
25 changes: 19 additions & 6 deletions src/roe/models/job.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ def wait(
raise ValueError(f"timeout must be positive, got {timeout}")

effective_timeout = timeout if timeout is not None else self._timeout_seconds
start_time = time.time()
deadline = time.monotonic() + effective_timeout

from roe._generated.types import Unset

Expand All @@ -132,12 +132,13 @@ def wait(
return _empty_result(status.status, error_message)
return _attach_status(result, status.status, error_message)

if (time.time() - start_time) > effective_timeout:
remaining = deadline - time.monotonic()
if remaining <= 0:
raise TimeoutError(
f"Job {self._job_id} did not complete within {effective_timeout} seconds"
)

time.sleep(interval)
time.sleep(min(interval, remaining))
Comment thread
jadenfix marked this conversation as resolved.

def retrieve_status(self) -> AgentJobSingleStatus:
"""Generated ``AgentJobSingleStatus`` for the job."""
Expand Down Expand Up @@ -201,7 +202,7 @@ def wait(
raise ValueError(f"timeout must be positive, got {timeout}")

effective_timeout = timeout if timeout is not None else self._timeout_seconds
start_time = time.time()
deadline = time.monotonic() + effective_timeout

while len(self._completed_jobs) < len(self._job_ids):
pending_job_ids = [
Expand All @@ -214,10 +215,12 @@ def wait(
status_batch = self.agents_api.jobs.retrieve_status_many(pending_job_ids)

completed_in_this_batch: list[str] = []
returned_ids: set[str] = set()
for status_item in status_batch:
job_id = self._extract_id(status_item)
if job_id is None:
continue
returned_ids.add(job_id)
stat_code = self._extract_status(status_item)
if stat_code in _TERMINAL_STATUSES:
completed_in_this_batch.append(job_id)
Expand All @@ -228,6 +231,15 @@ def wait(
"timestamp": self._extract_timestamp(status_item),
}

# Otherwise a job the server never returns is polled until the timeout.
missing = [
job_id
for job_id in pending_job_ids
if str(UUID(str(job_id))) not in returned_ids
]
if missing:
raise NotFoundError(f"Jobs {missing} not found in status response")

if completed_in_this_batch:
result_batch = self.agents_api.jobs.retrieve_result_many(
completed_in_this_batch
Expand All @@ -249,15 +261,16 @@ def wait(
self._completed_jobs[job_id] = result_item

if len(self._completed_jobs) < len(self._job_ids):
if (time.time() - start_time) > effective_timeout:
remaining = deadline - time.monotonic()
if remaining <= 0:
if raise_on_timeout:
remaining = set(self._job_ids) - set(self._completed_jobs)
raise TimeoutError(
f"Jobs {remaining} did not complete within {effective_timeout} seconds"
)
break

time.sleep(interval)
time.sleep(min(interval, remaining))

return [self._completed_jobs.get(job_id) for job_id in self._job_ids]

Expand Down
15 changes: 13 additions & 2 deletions src/roe/utils/inputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,17 @@ def build_execution_multipart(

for key, value in inputs.items():
if isinstance(value, FileUpload):
filename, file_obj, mime_type = value.to_multipart_tuple()
files[key] = (filename, file_obj, mime_type)
if value.path:
# Read now so the handle is closed; nothing closes it after the request.
with open(value.path, "rb") as fh:
files[key] = (
value.effective_filename,
fh.read(),
Comment thread
jadenfix marked this conversation as resolved.
value.effective_mime_type,
)
else:
filename, file_obj, mime_type = value.to_multipart_tuple()
files[key] = (filename, file_obj, mime_type)
elif isinstance(value, (io.IOBase, io.BytesIO)) or hasattr(value, "read"):
files[key] = value
elif isinstance(value, str):
Expand All @@ -56,6 +65,8 @@ def build_execution_multipart(
files[key] = (p.name, fh.read(), mime or "application/octet-stream")
else:
form_data[key] = value
elif isinstance(value, (dict, list)):
form_data[key] = _json.dumps(value)
else:
if value is not None:
form_data[key] = str(value)
Expand Down
2 changes: 1 addition & 1 deletion src/roe/utils/transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
rewindable buffers for JSON-encoded bodies).

Multipart agent-run helpers opt out via the ``x-roe-skip-retry`` header so
those POSTs are not retried (non-idempotent streamed bodies).
those POSTs are not retried (non-idempotent); ``run_many`` opts out too.

See TS ``retryMiddleware`` / ``dynamicInputs.postDynamicInputs`` and Go
``doRetried`` for the analogous contract across SDKs.
Expand Down
13 changes: 12 additions & 1 deletion tests/unit/test_agents_wrapper_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
import pytest

from roe.api.agents import AgentsAPI
from roe.exceptions import NotFoundError
from roe.exceptions import NotFoundError, RoeAPIException

ORG_ID = "00000000-0000-0000-0000-000000000123"
AGENT_ID = "00000000-0000-0000-0000-000000000111"
Expand Down Expand Up @@ -157,6 +157,17 @@ def test_run_many_sends_skip_cache_header_on_every_chunk():
assert request.call_count == 2
for call in request.call_args_list:
assert call.kwargs["headers"]["X-Skip-Cache"] == "true"
assert call.kwargs["headers"]["x-roe-skip-retry"] == "1"


def test_run_many_failure_keeps_job_ids_from_earlier_chunks():
api, request = _api(httpx.Response(200, json=[JOB_ID] * 1000))
request.side_effect = [request.return_value, httpx.Response(500, json={})]

with pytest.raises(RoeAPIException) as exc_info:
api.run_many(AGENT_ID, [{"prompt": "hello"}] * 1001)

assert exc_info.value.submitted_job_ids == [JOB_ID] * 1000


def test_sync_and_version_runs_omit_skip_cache_header_by_default():
Expand Down
9 changes: 9 additions & 0 deletions tests/unit/test_auth.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
import roe
from roe.auth import RoeAuth
from roe.config import RoeConfig


def test_user_agent_uses_package_version():
auth = RoeAuth(RoeConfig(api_key="key", organization_id="org"))

assert auth.get_headers()["User-Agent"] == f"roe-python/{roe.__version__}"
7 changes: 7 additions & 0 deletions tests/unit/test_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
from roe.config import RoeConfig


def test_max_retries_zero_is_respected(monkeypatch):
monkeypatch.setenv("ROE_MAX_RETRIES", "5")
config = RoeConfig.from_env(api_key="key", organization_id="org", max_retries=0)
assert config.max_retries == 0
23 changes: 23 additions & 0 deletions tests/unit/test_inputs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
import pytest

from roe.models.file import FileUpload
from roe.utils.inputs import build_execution_multipart


def test_file_upload_path_is_read_not_left_open(tmp_path):
path = tmp_path / "invoice.pdf"
path.write_bytes(b"%PDF-1.4")

_, files = build_execution_multipart({"document": FileUpload(path=str(path))})

assert files["document"] == ("invoice.pdf", b"%PDF-1.4", "application/pdf")


@pytest.mark.parametrize(
("value", "expected"),
[({"a": "b"}, '{"a": "b"}'), (["x", 1], '["x", 1]'), (3, "3"), (True, "True")],
)
def test_non_string_inputs_are_sent_as_form_values(value, expected):
form_data, _ = build_execution_multipart({"field": value})

assert form_data == {"field": expected}
73 changes: 73 additions & 0 deletions tests/unit/test_job_wait.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
from types import SimpleNamespace
from uuid import UUID

import pytest

from roe.exceptions import NotFoundError
from roe.models import job as job_module
from roe.models.job import Job, JobBatch, JobStatus


class FakeClock:
def __init__(self):
self.now = 0.0

def monotonic(self):
return self.now

def sleep(self, seconds):
self.now += seconds


@pytest.fixture
def clock(monkeypatch):
fake = FakeClock()
monkeypatch.setattr(job_module, "time", fake)
return fake


def test_job_wait_does_not_sleep_past_timeout(clock):
jobs = SimpleNamespace(
retrieve_status=lambda job_id: SimpleNamespace(
status=JobStatus.STARTED, error_message=None
)
)
job = Job(SimpleNamespace(jobs=jobs), "job-1")

with pytest.raises(TimeoutError):
job.wait(interval=60, timeout=1)

assert clock.now == 1


def test_job_batch_wait_does_not_sleep_past_timeout(clock):
jobs = SimpleNamespace(
retrieve_status_many=lambda job_ids: [
SimpleNamespace(id=UUID(job_id), status=JobStatus.STARTED)
for job_id in job_ids
]
)
batch = JobBatch(
SimpleNamespace(jobs=jobs), ["00000000-0000-0000-0000-000000000001"]
)

with pytest.raises(TimeoutError):
batch.wait(interval=60, timeout=1)

assert clock.now == 1


def test_job_batch_wait_raises_when_status_response_omits_a_job(clock):
present = "00000000-0000-0000-0000-000000000001"
missing = "00000000-0000-0000-0000-000000000002"
jobs = SimpleNamespace(
retrieve_status_many=lambda job_ids: [
SimpleNamespace(id=UUID(present), status=JobStatus.STARTED)
]
)
batch = JobBatch(SimpleNamespace(jobs=jobs), [present, missing])

with pytest.raises(NotFoundError, match=missing):
batch.wait(interval=5, timeout=60)

assert clock.now == 0
Loading