Skip to content
Open
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
24 changes: 8 additions & 16 deletions src/openai/lib/_vector_stores.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,10 +42,8 @@ def poll_vector_store_file(

file = response.parse()
if file.status == "in_progress":
if not is_given(poll_interval_ms):
poll_interval_ms = _get_poll_interval_ms(response.headers)

resource._sleep(poll_interval_ms / 1000)
interval_ms = poll_interval_ms if is_given(poll_interval_ms) else _get_poll_interval_ms(response.headers)
resource._sleep(interval_ms / 1000)
elif file.status == "cancelled" or file.status == "completed" or file.status == "failed":
return file
else:
Expand Down Expand Up @@ -76,10 +74,8 @@ async def async_poll_vector_store_file(

file = response.parse()
if file.status == "in_progress":
if not is_given(poll_interval_ms):
poll_interval_ms = _get_poll_interval_ms(response.headers)

await resource._sleep(poll_interval_ms / 1000)
interval_ms = poll_interval_ms if is_given(poll_interval_ms) else _get_poll_interval_ms(response.headers)
await resource._sleep(interval_ms / 1000)
elif file.status == "cancelled" or file.status == "completed" or file.status == "failed":
return file
else:
Expand Down Expand Up @@ -110,10 +106,8 @@ def poll_vector_store_file_batch(

batch = response.parse()
if batch.file_counts.in_progress > 0:
if not is_given(poll_interval_ms):
poll_interval_ms = _get_poll_interval_ms(response.headers)

resource._sleep(poll_interval_ms / 1000)
interval_ms = poll_interval_ms if is_given(poll_interval_ms) else _get_poll_interval_ms(response.headers)
resource._sleep(interval_ms / 1000)
continue

return batch
Expand All @@ -140,10 +134,8 @@ async def async_poll_vector_store_file_batch(

batch = response.parse()
if batch.file_counts.in_progress > 0:
if not is_given(poll_interval_ms):
poll_interval_ms = _get_poll_interval_ms(response.headers)

await resource._sleep(poll_interval_ms / 1000)
interval_ms = poll_interval_ms if is_given(poll_interval_ms) else _get_poll_interval_ms(response.headers)
await resource._sleep(interval_ms / 1000)
continue

return batch
18 changes: 18 additions & 0 deletions tests/lib/test_vector_store_polling.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,24 @@ async def test_poll_interval_and_headers(
sleep.assert_called_once_with(seconds)


async def test_server_poll_interval_is_read_from_each_response(resource: PollResource) -> None:
terminal = raw_response(resource)
with (
mock.patch.object(
resource.with_raw_response,
"retrieve",
side_effect=[
raw_response(resource, pending=True, headers={"openai-poll-after-ms": "100"}),
raw_response(resource, pending=True, headers={"openai-poll-after-ms": "2000"}),
terminal,
],
),
mock.patch.object(resource, "_sleep") as sleep,
):
assert await poll(resource) is terminal.parse.return_value
assert sleep.call_args_list == [mock.call(0.1), mock.call(2.0)]


async def test_retrieve_error_propagates(resource: PollResource) -> None:
error = RuntimeError("synthetic retrieval failure")
with mock.patch.object(resource.with_raw_response, "retrieve", side_effect=error):
Expand Down