diff --git a/dataconnect/client.py b/dataconnect/client.py index 637121e..9473a30 100644 --- a/dataconnect/client.py +++ b/dataconnect/client.py @@ -43,9 +43,9 @@ def connect( # Public API - def get_studies(self) -> list[Study]: + def get_studies(self, search_study_name: str | None = None) -> list[Study]: """List the studies the client is authorized to access.""" - return self._service.get_studies() + return self._service.get_studies(search_study_name=search_study_name) def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: """List the dataset versions the client is authorized to access.""" diff --git a/dataconnect/service/base.py b/dataconnect/service/base.py index 938c6a9..338d340 100644 --- a/dataconnect/service/base.py +++ b/dataconnect/service/base.py @@ -12,7 +12,7 @@ class DataConnectService(ABC): """Abstract service interface — defines all operations available to the client.""" @abstractmethod - def get_studies(self) -> list[Study]: ... + def get_studies(self, search_study_name: str | None = None) -> list[Study]: ... @abstractmethod def get_dataset_versions(self, dataset_uuid: UUID) -> list[DatasetVersion]: ... diff --git a/dataconnect/service/default.py b/dataconnect/service/default.py index 204586f..5272ea8 100644 --- a/dataconnect/service/default.py +++ b/dataconnect/service/default.py @@ -17,6 +17,7 @@ from dataconnect.models import DatasetVersion, Study from dataconnect.service.base import DataConnectService from dataconnect.service.mappers import resource_to_dataset_version, resource_to_study +from dataconnect.service.validators import validate_search_study_name from dataconnect.transport.base import Transport from dataconnect.transport.errors import ( TransportAuthenticationError, @@ -61,9 +62,13 @@ def __init__(self, transport: Transport) -> None: # DataConnectService - def get_studies(self) -> list[Study]: + def get_studies(self, search_study_name: str | None = None) -> list[Study]: + + validate_search_study_name(search_study_name) request = ResourceQuery(action=_ACTION_LIST_STUDIES) + if search_study_name and search_study_name.strip() != "": + request = request.append_body({"search_study_name": search_study_name}) try: resources = self._transport.list_resources(request) diff --git a/dataconnect/service/validators.py b/dataconnect/service/validators.py new file mode 100644 index 0000000..5909d36 --- /dev/null +++ b/dataconnect/service/validators.py @@ -0,0 +1,15 @@ +"""Input validation helpers for service-layer operations.""" + +from __future__ import annotations + +from dataconnect.exceptions import ValidationError + + +def validate_search_study_name(search_study_name: str | None) -> None: + """Validate optional study-name filter used by ``get_studies``.""" + + if not search_study_name: + return + + if not isinstance(search_study_name, str): + raise ValidationError("search_study_name must be a string") diff --git a/tests/test_client.py b/tests/test_client.py index 38b1563..92ad150 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -28,13 +28,6 @@ def close(self) -> None: self.closed += 1 -def test_get_studies_returns_service_result() -> None: - studies = [Study(uuid=UUID("64a98a9b-1512-44c8-92af-e4cab0183670"), name="Study A")] - client = DataConnectClient(_FakeService(studies=studies)) - - assert client.get_studies() == studies - - def test_get_dataset_versions_forwards_uuid_to_service() -> None: dataset_uuid = UUID("073410b6-79be-3e7d-ae37-92f6e054013e") versions = [ @@ -109,3 +102,45 @@ def close(self) -> None: def test_dummy_benchmark() -> None: # Dummy benchmark test to satisfy CI assert True + + +class StubService: + def __init__(self) -> None: + self.search_study_name: str | None = None + self.was_closed = False + + def get_studies(self, search_study_name: str | None = None) -> list[Study]: + self.search_study_name = search_study_name + return [] + + def close(self) -> None: + self.was_closed = True + + +def test_get_studies_without_filter_delegates_to_service() -> None: + service = StubService() + client = DataConnectClient(service) + + studies = client.get_studies() + + assert studies == [] + assert service.search_study_name is None + + +def test_get_studies_with_filter_delegates_to_service() -> None: + service = StubService() + client = DataConnectClient(service) + + studies = client.get_studies(search_study_name="cardio") + + assert studies == [] + assert service.search_study_name == "cardio" + + +def test_close_delegates_to_service() -> None: + service = StubService() + client = DataConnectClient(service) + + client.close() + + assert service.was_closed diff --git a/tests/test_service_default.py b/tests/test_service_default.py new file mode 100644 index 0000000..d1ac135 --- /dev/null +++ b/tests/test_service_default.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +import pytest + +from dataconnect.exceptions import ValidationError +from dataconnect.service.default import DefaultDataConnectService +from dataconnect.transport.models import DataRef, ResourceInfo, ResourceQuery + + +class StubTransport: + def __init__(self, resources: list[ResourceInfo]) -> None: + self.resources = resources + self.last_request: ResourceQuery | None = None + + def list_resources(self, request: ResourceQuery) -> list[ResourceInfo]: + self.last_request = request + return self.resources + + def close(self) -> None: + return None + + +def _study_resource(name: str = "Study A") -> ResourceInfo: + payload = (f'{{"uuid":"12345678-1234-1234-1234-123456789abc","name":"{name}","environments":[]}}').encode() + + return ResourceInfo( + descriptor=b"", + endpoints=[DataRef(ticket=payload)], + total_records=1, + schema_bytes=b"", + ) + + +def test_get_studies_without_search_name_uses_empty_request_body() -> None: + transport = StubTransport(resources=[_study_resource()]) + service = DefaultDataConnectService(transport) + + studies = service.get_studies() + + assert len(studies) == 1 + assert studies[0].name == "Study A" + assert transport.last_request is not None + assert transport.last_request.action == "studies.list" + assert transport.last_request.body == "" + + +def test_get_studies_with_search_name_sets_request_body() -> None: + transport = StubTransport(resources=[_study_resource("Cardio Study")]) + service = DefaultDataConnectService(transport) + + studies = service.get_studies(search_study_name="Cardio") + + assert len(studies) == 1 + assert studies[0].name == "Cardio Study" + assert transport.last_request is not None + assert transport.last_request.body == '{"search_study_name":"Cardio"}' + + +def test_get_studies_rejects_non_string_search_name() -> None: + transport = StubTransport(resources=[]) + service = DefaultDataConnectService(transport) + + with pytest.raises(ValidationError, match="search_study_name must be a string"): + service.get_studies(search_study_name=123) # type: ignore[arg-type] + + assert transport.last_request is None + + +def test_get_studies_accepts_none_search_name() -> None: + transport = StubTransport(resources=[_study_resource()]) + service = DefaultDataConnectService(transport) + + studies = service.get_studies(search_study_name=None) + + assert len(studies) == 1 + assert transport.last_request is not None + assert transport.last_request.body == ""