diff --git a/src/roe/api/agents.py b/src/roe/api/agents.py index 3c1cfcd..ed1f634 100644 --- a/src/roe/api/agents.py +++ b/src/roe/api/agents.py @@ -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 @@ -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( diff --git a/src/roe/auth.py b/src/roe/auth.py index f80b839..9279236 100644 --- a/src/roe/auth.py +++ b/src/roe/auth.py @@ -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__}", } diff --git a/src/roe/config.py b/src/roe/config.py index e857963..e13be64 100644 --- a/src/roe/config.py +++ b/src/roe/config.py @@ -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 diff --git a/src/roe/models/job.py b/src/roe/models/job.py index 3d3fd80..8dc8087 100644 --- a/src/roe/models/job.py +++ b/src/roe/models/job.py @@ -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 @@ -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)) def retrieve_status(self) -> AgentJobSingleStatus: """Generated ``AgentJobSingleStatus`` for the job.""" @@ -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 = [ @@ -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) @@ -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 @@ -249,7 +261,8 @@ 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( @@ -257,7 +270,7 @@ def wait( ) break - time.sleep(interval) + time.sleep(min(interval, remaining)) return [self._completed_jobs.get(job_id) for job_id in self._job_ids] diff --git a/src/roe/utils/inputs.py b/src/roe/utils/inputs.py index b1026ab..137d0bf 100644 --- a/src/roe/utils/inputs.py +++ b/src/roe/utils/inputs.py @@ -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(), + 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): @@ -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) diff --git a/src/roe/utils/transport.py b/src/roe/utils/transport.py index d18fce6..f045a98 100644 --- a/src/roe/utils/transport.py +++ b/src/roe/utils/transport.py @@ -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. diff --git a/tests/unit/test_agents_wrapper_transport.py b/tests/unit/test_agents_wrapper_transport.py index 8e6cdcd..4bc39e4 100644 --- a/tests/unit/test_agents_wrapper_transport.py +++ b/tests/unit/test_agents_wrapper_transport.py @@ -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" @@ -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(): diff --git a/tests/unit/test_auth.py b/tests/unit/test_auth.py new file mode 100644 index 0000000..df03c2f --- /dev/null +++ b/tests/unit/test_auth.py @@ -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__}" diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py new file mode 100644 index 0000000..6fc8191 --- /dev/null +++ b/tests/unit/test_config.py @@ -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 diff --git a/tests/unit/test_inputs.py b/tests/unit/test_inputs.py new file mode 100644 index 0000000..2de8358 --- /dev/null +++ b/tests/unit/test_inputs.py @@ -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} diff --git a/tests/unit/test_job_wait.py b/tests/unit/test_job_wait.py new file mode 100644 index 0000000..2f59d26 --- /dev/null +++ b/tests/unit/test_job_wait.py @@ -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