Skip to content

Commit d12082b

Browse files
OllieinCanadaOliver Slapinski
andauthored
fix(utils): close streamed dataset responses (#799)
Signed-off-by: Oliver Slapinski <olliefromcanada@gmail.com> Co-authored-by: Oliver Slapinski <olliefromcanada@gmail.com>
1 parent 050f9c0 commit d12082b

2 files changed

Lines changed: 47 additions & 3 deletions

File tree

src/cohere/utils.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -280,9 +280,9 @@ def dataset_generator(dataset: Dataset):
280280
for part in dataset.dataset_parts:
281281
if not part.url:
282282
raise ValueError("Dataset part does not have a url")
283-
resp = requests.get(part.url, stream=True)
284-
for record in reader(resp.raw): # type: ignore
285-
yield record
283+
with requests.get(part.url, stream=True) as resp:
284+
for record in reader(resp.raw): # type: ignore
285+
yield record
286286

287287

288288
class SdkUtils:

tests/test_dataset_utils.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
from types import SimpleNamespace
2+
from typing import Any, cast
3+
4+
import cohere.utils
5+
from cohere import Dataset
6+
from cohere.utils import dataset_generator
7+
8+
9+
class TrackingResponse:
10+
def __init__(self) -> None:
11+
self.raw = object()
12+
self.closed = False
13+
14+
def close(self) -> None:
15+
self.closed = True
16+
17+
def __enter__(self) -> "TrackingResponse":
18+
return self
19+
20+
def __exit__(self, *_: Any) -> None:
21+
self.close()
22+
23+
24+
def test_dataset_generator_closes_response_after_exhaustion(monkeypatch: Any) -> None:
25+
response = TrackingResponse()
26+
monkeypatch.setattr(cohere.utils.requests, "get", lambda *_args, **_kwargs: response)
27+
monkeypatch.setattr(cohere.utils, "reader", lambda raw: iter([{"id": 1}]))
28+
dataset = cast(Dataset, SimpleNamespace(dataset_parts=[SimpleNamespace(url="https://example.test/part.avro")]))
29+
30+
assert list(dataset_generator(dataset)) == [{"id": 1}]
31+
assert response.closed
32+
33+
34+
def test_dataset_generator_closes_response_when_consumer_stops(monkeypatch: Any) -> None:
35+
response = TrackingResponse()
36+
monkeypatch.setattr(cohere.utils.requests, "get", lambda *_args, **_kwargs: response)
37+
monkeypatch.setattr(cohere.utils, "reader", lambda raw: iter([{"id": 1}, {"id": 2}]))
38+
dataset = cast(Dataset, SimpleNamespace(dataset_parts=[SimpleNamespace(url="https://example.test/part.avro")]))
39+
records = dataset_generator(dataset)
40+
41+
assert next(records) == {"id": 1}
42+
records.close()
43+
44+
assert response.closed

0 commit comments

Comments
 (0)