|
| 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