diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 08a70e8..6af3caa 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,6 +1,6 @@ repos: - repo: https://github.com/psf/black - rev: 22.3.0 # Replace by any tag/version: https://github.com/psf/black/tags + rev: 22.6.0 # Replace by any tag/version: https://github.com/psf/black/tags hooks: - id: black language_version: python3 # Should be a command that runs python3.6+ \ No newline at end of file diff --git a/dcm/cli.py b/dcm/cli.py index f939b67..66702aa 100644 --- a/dcm/cli.py +++ b/dcm/cli.py @@ -713,6 +713,7 @@ async def _do_route( break finally: print("Listener shutting down") + report.log_issues() @click.command() @@ -772,6 +773,8 @@ def forward( router = Router(dests) asyncio.run(_do_route(local, router, inactive_timeout)) + log.debug("Forward call finished successfully") + return 0 def make_print_cb(fmt, elem_filter=None): diff --git a/dcm/filt.py b/dcm/filt.py index 546f513..3d28d8e 100644 --- a/dcm/filt.py +++ b/dcm/filt.py @@ -15,7 +15,7 @@ QueryResult, DataNode, get_uid, - uid_elems, + UID_ELEMS, QueryLevelMismatchError, InconsistentDataError, ) @@ -24,7 +24,7 @@ log = logging.getLogger(__name__) -uid_elem_set = FrozenLazySet(uid_elems.values()) +uid_elem_set = FrozenLazySet(UID_ELEMS.values()) @dataclass(frozen=True) @@ -241,7 +241,7 @@ def _add_new(self, old_ds: Dataset, new_ds: Dataset) -> Optional[str]: if new_dupe: return None self.new.add(new_ds) - for lvl, uid_attr in uid_elems.items(): + for lvl, uid_attr in UID_ELEMS.items(): new_uid = get_uid(lvl, new_ds) old_uid = get_uid(lvl, old_ds) self._new_to_old[lvl][new_uid] = old_uid diff --git a/dcm/net.py b/dcm/net.py index a2dfba5..e5669c3 100644 --- a/dcm/net.py +++ b/dcm/net.py @@ -60,9 +60,11 @@ QueryLevel, QueryResult, InconsistentDataError, - uid_elems, - req_elems, - opt_elems, + expand_queries, + get_level_and_query, + UID_ELEMS, + REQ_ELEMS, + OPT_ELEMS, choose_level, minimal_copy, get_all_uids, @@ -73,6 +75,7 @@ MultiListReport, MultiError, ProgressHookBase, + optional_report, ) from .util import ( json_serializer, @@ -837,7 +840,7 @@ def _query_worker( split_level = QueryLevel(level - 1) elif split_level > level: raise ValueError("The split_level can't be higher than the query level") - split_attr = uid_elems[split_level] + split_attr = UID_ELEMS[split_level] last_split_val = None resp_count = 0 @@ -901,7 +904,7 @@ def _query_worker( def _make_move_request(ds: Dataset) -> Dataset: res = Dataset() - for uid_attr in uid_elems.values(): + for uid_attr in UID_ELEMS.values(): uid_val = getattr(ds, uid_attr, None) if uid_val is not None: setattr(res, uid_attr, uid_val) @@ -1178,12 +1181,13 @@ async def query( See documentation for the `queries` method for details """ - level, query = self._prep_query(level, query, query_res) + level, query = get_level_and_query(level, query, query_res) res = QueryResult(level) async for sub_res in self.queries(remote, level, query, query_res, report): res |= sub_res return res + @optional_report async def queries( self, remote: DcmNode, @@ -1215,16 +1219,12 @@ async def queries( report If provided, will store status report from DICOM operations """ - if report is None: - extern_report = False - report = MultiListReport(meta_data={"remote": remote, "level": level}) - else: - extern_report = True + assert report is not None if report._description is None: report.description = "queries" report._meta_data["remote"] = remote - level, query = self._prep_query(level, query, query_res) + level, query = get_level_and_query(level, query, query_res) report._meta_data["level"] = level if query is not None: @@ -1267,13 +1267,13 @@ async def queries( # Add in any required and default optional elements to query for lvl in QueryLevel: - for req_attr in req_elems[lvl]: + for req_attr in REQ_ELEMS[lvl]: if getattr(query, req_attr, None) is None: setattr(query, req_attr, "") if lvl == level: break auto_attrs = set() - for opt_attr in opt_elems[level]: + for opt_attr in OPT_ELEMS[level]: if getattr(query, opt_attr, None) is None: setattr(query, opt_attr, "") auto_attrs.add(opt_attr) @@ -1281,25 +1281,9 @@ async def queries( # Set the QueryRetrieveLevel query.QueryRetrieveLevel = level.name - # Pull out a list of the attributes we are querying on - queried_elems = set(e.keyword for e in query) - - # If QueryResult was given we potentially generate multiple - # queries, one for each dataset referenced by the QueryResult - if query_res is None: - queries = [query] - else: - queries = [] - for path, sub_uids in query_res.walk(): - if path.level == min(level, query_res.level): - q = deepcopy(query) - for lvl in QueryLevel: - if lvl > path.level: - break - setattr(q, uid_elems[lvl], path.uids[lvl]) - queried_elems.add(uid_elems[lvl]) - queries.append(q) - sub_uids.clear() + # Potentially expand one query into multiple based on query_res + queries, queried_elems = expand_queries(level, query, query_res) + if query_res is not None: log.debug("QueryResult expansion results in %d sub-queries" % len(queries)) if len(queries) > 1: report.n_expected = len(queries) @@ -1387,9 +1371,6 @@ async def queries( await rep_builder_task except BaseException as e: log.exception("Exception from report builder task") - if not extern_report: - report.log_issues() - report.check_errors() @asynccontextmanager async def listen( @@ -1446,6 +1427,7 @@ async def listen( await self._cleanup_listen_mgr() log.debug("Listener lock released") + @optional_report async def move( self, source: DcmNode, @@ -1455,11 +1437,7 @@ async def move( report: MultiListReport[DicomOpReport] = None, ) -> None: """Move DICOM files from one network entity to another""" - if report is None: - extern_report = False - report = MultiListReport() - else: - extern_report = True + assert report is not None if report._description is None: report.description = "move" report._meta_data["source"] = source @@ -1516,11 +1494,9 @@ async def move( if rep_builder_task is not None: await rep_builder_task log.debug("Report builder task is done") - if not extern_report: - report.log_issues() - report.check_errors() log.debug("Call to LocalEntity.move has completed") + @optional_report async def retrieve( self, remote: DcmNode, @@ -1546,11 +1522,7 @@ async def retrieve( By default inconsistent, unexpected, and duplicate data are skipped """ - if report is None: - extern_report = False - report = RetrieveReport() - else: - extern_report = True + assert report is not None report.requested = query_res report._meta_data["remote"] = remote self._add_qr_meta(report, query_res) @@ -1602,13 +1574,10 @@ async def retrieve( log.debug("Waiting for move task to finish") await move_task report.done = True - if not extern_report: - report.log_issues() - log.debug("About to check errors") - report.check_errors() log.debug("The LocalEntity.retrieve method has completed") @asynccontextmanager + @optional_report async def send( self, remote: DcmNode, @@ -1637,13 +1606,9 @@ async def send( as it is assumed the caller will handle this themselves (e.g. by calling the `log_issues` and `check_errors` methods on the report). """ + assert report is not None if transfer_syntax is None: transfer_syntax = self._default_ts - if report is None: - extern_report = False - report = DicomOpReport() - else: - extern_report = True report.dicom_op.provider = remote report.dicom_op.user = self._local report.dicom_op.op_type = "c-store" @@ -1674,9 +1639,6 @@ async def send( await rep_builder_task finally: report.done = True - if not extern_report: - report.log_issues() - report.check_errors() async def download( self, @@ -1752,30 +1714,6 @@ async def _associate( log.debug("Releasing association") await loop.run_in_executor(self._thread_pool, assoc.release) - def _prep_query( - self, - level: Optional[QueryLevel], - query: Optional[Dataset], - query_res: Optional[QueryResult], - ) -> Tuple[QueryLevel, Dataset]: - """Resolve/check `level` and `query` args for query methods""" - # Build up our base query dataset - if query is None: - query = Dataset() - else: - query = deepcopy(query) - - # Deterimine level if not specified, otherwise make sure it is valid - if level is None: - if query_res is None: - default_level = QueryLevel.STUDY - else: - default_level = query_res.level - level = choose_level(query, default_level) - elif level not in QueryLevel: - raise ValueError("Unknown 'level' for query: %s" % level) - return level, query - async def _fwd_event(self, event: evt.Event) -> int: for filt, handler in self._event_handlers.items(): if filt.matches(event): diff --git a/dcm/query.py b/dcm/query.py index 2bb3ed3..6241b2c 100644 --- a/dcm/query.py +++ b/dcm/query.py @@ -9,6 +9,7 @@ from enum import IntEnum from dataclasses import dataclass, field from typing import ( + Iterable, Tuple, List, Dict, @@ -19,6 +20,7 @@ Union, Set, ) +from xml.etree.ElementInclude import include from pydicom.dataset import Dataset from tree_format import format_tree @@ -39,7 +41,7 @@ class QueryLevel(IntEnum): IMAGE = 3 -uid_elems = { +UID_ELEMS = { QueryLevel.PATIENT: "PatientID", QueryLevel.STUDY: "StudyInstanceUID", QueryLevel.SERIES: "SeriesInstanceUID", @@ -50,16 +52,19 @@ class QueryLevel(IntEnum): def get_uid(level: QueryLevel, data_set: Dataset) -> str: """Get the UID from the `data_set` for the given `level`""" - return getattr(data_set, uid_elems[level]) + return getattr(data_set, UID_ELEMS[level]) def get_all_uids(data_set: Dataset) -> Tuple[str, ...]: + """Get tuple of UIDs corresponding to levels, with trailing empty UIDs trimmed""" uids = [] + last_found = -1 for lvl in QueryLevel: - lvl_uid = getattr(data_set, uid_elems[lvl], None) - if lvl_uid is not None: - uids.append(lvl_uid) - return tuple(uids) + lvl_uid = getattr(data_set, UID_ELEMS[lvl], "") + if lvl_uid != "": + last_found = lvl + uids.append(lvl_uid) + return tuple(uids[: last_found + 1]) @dataclass(frozen=True) @@ -77,6 +82,10 @@ class DataNode: class DataPath: """Identifies the path to a node in the DICOM data hierarchy""" + @classmethod + def from_uids(cls, uids: Tuple[str, ...]) -> "DataPath": + return cls(QueryLevel(len(uids) - 1), uids) + level: QueryLevel """The level of the node in the hierarchy""" @@ -93,8 +102,25 @@ def __post_init__(self) -> None: ) object.__setattr__(self, "end", DataNode(self.level, self.uids[-1])) + def __add__(self, new_end: DataNode) -> DataPath: + if new_end.level != self.level + 1: + raise ValueError( + "Trying to add node at level %s to path with level %s", + new_end.level, + self.level, + ) + return DataPath(new_end.level, self.uids + (new_end.uid,)) + + @property + def parent_uid(self) -> str: + return self.uids[-2] + + @property + def parent(self) -> DataPath: + return self.from_uids(self.uids[:-1]) -req_elems = { + +REQ_ELEMS = { QueryLevel.PATIENT: [ "PatientID", "PatientName", @@ -117,7 +143,7 @@ def __post_init__(self) -> None: """Required attributes for each query level (accumulates at each level)""" -blankable_req_elems = [ +BLANKABLE_REQ_ELEMS = [ "PatientID", "PatientName", "StudyDate", @@ -129,7 +155,7 @@ def __post_init__(self) -> None: """ -opt_elems: Dict[QueryLevel, List[str]] = { +OPT_ELEMS: Dict[QueryLevel, List[str]] = { QueryLevel.PATIENT: [ "NumberOfPatientRelatedStudies", "NumberOfPatientRelatedSeries", @@ -145,6 +171,7 @@ def __post_init__(self) -> None: QueryLevel.SERIES: [ "SeriesDescription", "ProtocolName", + "SeriesTime", "NumberOfSeriesRelatedInstances", ], QueryLevel.IMAGE: [], @@ -152,13 +179,13 @@ def __post_init__(self) -> None: """Optional attributes we always try to query (exclusive to each level)""" -level_filters = { - lvl: make_elem_filter(req_elems[lvl] + opt_elems[lvl]) for lvl in QueryLevel +LEVEL_FILTERS = { + lvl: make_elem_filter(REQ_ELEMS[lvl] + OPT_ELEMS[lvl]) for lvl in QueryLevel } """Element filters for each level""" -level_identifiers = { +LEVEL_IDENTIFIERS = { QueryLevel.PATIENT: [ "NumberOfPatientRelatedStudies", "NumberOfPatientRelatedSeries", @@ -187,14 +214,17 @@ def __post_init__(self) -> None: """Maps query levels to elements that imply that level is needed""" -def minimal_copy(ds: Dataset) -> Dataset: +MIN_ATTRS = tuple(chain.from_iterable(chain(REQ_ELEMS.values(), OPT_ELEMS.values()))) + + +def minimal_copy(ds: Dataset, include_elems: Iterable[str] = MIN_ATTRS) -> Dataset: """Make reduced copy with only the attributes needed for a QueryResult""" res = Dataset() - for attr in chain.from_iterable(chain(req_elems.values(), opt_elems.values())): + for attr in include_elems: val = getattr(ds, attr, None) if val is not None: setattr(res, attr, val) - elif val in blankable_req_elems: + elif val in BLANKABLE_REQ_ELEMS: setattr(res, attr, "") return res @@ -202,7 +232,7 @@ def minimal_copy(ds: Dataset) -> Dataset: def choose_level(qdat: Dataset, default: QueryLevel = QueryLevel.STUDY) -> QueryLevel: """Try to choose the correct level for a given query""" for lvl in reversed(QueryLevel): - for attr in level_identifiers[lvl]: + for attr in LEVEL_IDENTIFIERS[lvl]: if hasattr(qdat, attr): return lvl return default @@ -447,7 +477,7 @@ def add(self, data_set: Dataset) -> None: last_info = None branch_found = False for lvl in QueryLevel: - lvl_uid = getattr(data_set, uid_elems[lvl]) + lvl_uid = getattr(data_set, UID_ELEMS[lvl]) lvl_info = self._levels[lvl].get(lvl_uid) if lvl_info is None: branch_found = True @@ -455,11 +485,11 @@ def add(self, data_set: Dataset) -> None: if last_info is not None: parent_uid = last_info["level_uid"] lvl_info = OrderedDict() - normed_data = normalize(data_set, level_filters[lvl]) - for attr in req_elems[lvl]: + normed_data = normalize(data_set, LEVEL_FILTERS[lvl]) + for attr in REQ_ELEMS[lvl]: val = normed_data.get(attr, None) if val is None: - if attr not in blankable_req_elems: + if attr not in BLANKABLE_REQ_ELEMS: raise InvalidDicomError( f"Dataset is missing required elem: {attr}" ) @@ -468,7 +498,7 @@ def add(self, data_set: Dataset) -> None: ) val = "" lvl_info[attr] = val - for attr in opt_elems[lvl]: + for attr in OPT_ELEMS[lvl]: val = normed_data.get(attr) if val is not None: lvl_info[attr] = val @@ -498,7 +528,7 @@ def remove(self, data_set: Dataset) -> None: for lvl in reversed(QueryLevel): if lvl > self._level: continue - lvl_uid = getattr(data_set, uid_elems[lvl]) + lvl_uid = getattr(data_set, UID_ELEMS[lvl]) if lvl == self._level: del self._data[lvl_uid] lvl_info = self._levels[lvl][lvl_uid] @@ -522,7 +552,7 @@ def __contains__(self, data_set: Dataset) -> bool: last_info = None for lvl in QueryLevel: try: - lvl_uid = getattr(data_set, uid_elems[lvl]) + lvl_uid = getattr(data_set, UID_ELEMS[lvl]) except AttributeError: # The PatientID can be blank, though many systems won't support it... if lvl == QueryLevel.PATIENT: @@ -684,6 +714,20 @@ def get_path(self, node: DataNode) -> DataPath: parent_uid = parent["parent_uid"] return DataPath(node.level, tuple(reversed(path))) + def check_path(self, path: DataPath) -> bool: + """Return true if this `path` exists in our data hierarchy + + Raises InconsistentDataError if the path is not consistent with the hierarchy + """ + if path.level > self._level: + raise InsufficientQueryLevelError() + n_match = sum(1 if path.uids[l] in self._levels[l] else 0 for l in QueryLevel) + if n_match == 0: + return False + elif n_match < path.level + 1: + raise InconsistentDataError() + return True + def _get_info(self, lvl_info: Dict[str, Any]) -> Dict[str, Any]: res = {} for key, val in lvl_info.items(): @@ -888,9 +932,43 @@ def sub_query( return res def level_sub_queries(self, level: QueryLevel) -> Iterator[QueryResult]: + """Generate sub queries at the given `level`""" for dpath in self.level_paths(level): yield self.sub_query(dpath.end) + def chunk(self, max_instances: int = 1000) -> Iterator[QueryResult]: + """Generate sub queries constrained by size + + If n_instances is unknown, just yield series level (or highest available) sub + queries. + """ + n_inst = self.n_instances() + if n_inst is None: + # We don't have info to constrain by size + for sub_qr in self.level_sub_queries(min(self._level, QueryLevel.SERIES)): + yield sub_qr + elif n_inst < max_instances: + yield deepcopy(self) + else: + chunk_qr = QueryResult(self._level) + for curr_path, sub_uids in self.walk(): + n_inst = self.n_instances(curr_path.end) + assert n_inst is not None + if len(chunk_qr) + n_inst <= max_instances: + chunk_qr |= self.sub_query(curr_path.end) + sub_uids.clear() + elif curr_path.level == self._level: + if len(chunk_qr) != 0: + yield chunk_qr + chunk_qr = QueryResult(self._level) + if n_inst >= max_instances: + # Result is too big but we can't get any smaller + yield self.sub_query(curr_path.end) + else: + chunk_qr |= self.sub_query(curr_path.end) + if chunk_qr: + yield chunk_qr + def reduced(self, level: QueryLevel) -> QueryResult: """Create lower level of detail copy""" if level >= self._level: @@ -949,10 +1027,7 @@ def __and__(self, other: QueryResult) -> QueryResult: else: pref, non_pref = other, self else: - if len(self) > len(other): - pref, non_pref = other, self - else: - pref, non_pref = self, other + pref, non_pref = self, other res = QueryResult(max(self._level, other._level), prov=deepcopy(pref.prov)) for ds in pref._data.values(): if ds in non_pref: @@ -1032,14 +1107,14 @@ def to_json_dict(self) -> Dict[str, Any]: return { "level": self._level.name, "patients": self._levels[QueryLevel.PATIENT], - "prov": self.prov, + "prov": self.prov.to_json_dict(), } @classmethod def from_json_dict(cls, json_dict: Dict[str, Any]) -> QueryResult: """Create a QueryResult from a previous `to_json` call""" level = getattr(QueryLevel, json_dict["level"]) - res = cls(level, prov=json_dict["prov"]) + res = cls(level, prov=QueryProv.from_json_dict(json_dict["prov"])) patients = json_dict["patients"] visit_q = list(patients.values()) visited_stack: List[Dict[str, Any]] = [] @@ -1173,21 +1248,46 @@ def __str__(self) -> str: return "%s Level QR: %s" % (self.level.name, descr) -# TODO: Fix this or remove it -# async def chunk_qrs(qr_gen: AsyncIterator[QueryResult], -# chunk_size: int = 10) -> AsyncIterator[QueryResult]: -# '''Generator wrapper that aggregates QueryResults into larger chunks''' -# try: -# first = await qr_gen.__anext__() -# except StopAsyncIteration: -# return -# level = first.level -# res = QueryResult(level) -# res |= first -# res_size = 1 -# async for qr in qr_gen: -# if res_size == chunk_size: -# yield res -# res = QueryResult(level) -# res |= first -# res_size = 1 +def get_level_and_query( + level: Optional[QueryLevel], + query: Optional[Dataset], + query_res: Optional[QueryResult], +) -> Tuple[QueryLevel, Dataset]: + """Resolve/check `level` and `query` args for query methods""" + # Build up our base query dataset + if query is None: + query = Dataset() + else: + query = deepcopy(query) + + # Deterimine level if not specified, otherwise make sure it is valid + if level is None: + if query_res is None: + default_level = QueryLevel.STUDY + else: + default_level = query_res.level + level = choose_level(query, default_level) + elif level not in QueryLevel: + raise ValueError("Unknown 'level' for query: %s" % level) + return level, query + + +def expand_queries( + level: QueryLevel, query: Dataset, query_res: Optional[QueryResult] = None +) -> Tuple[List[Dataset], Set[str]]: + queried_elems = set(e.keyword for e in query) + if query_res is None: + queries = [query] + else: + queries = [] + for path, sub_uids in query_res.walk(): + if path.level == min(level, query_res.level): + q = deepcopy(query) + for lvl in QueryLevel: + if lvl > path.level: + break + setattr(q, UID_ELEMS[lvl], path.uids[lvl]) + queried_elems.add(UID_ELEMS[lvl]) + queries.append(q) + sub_uids.clear() + return queries, queried_elems diff --git a/dcm/report.py b/dcm/report.py index dc7bbcc..57c080c 100644 --- a/dcm/report.py +++ b/dcm/report.py @@ -6,10 +6,14 @@ variety of "report" classes to capture this kind of information and provide real-time insight into an ongoing async operation. """ -import logging +from collections import deque +from contextlib import contextmanager +import logging, inspect from dataclasses import dataclass, field from datetime import datetime +import typing from typing import ( + Iterable, Optional, Dict, List, @@ -21,10 +25,14 @@ ItemsView, KeysView, ValuesView, + Callable, ) +from typing_extensions import get_args import rich.progress +from .util import Args_Type, decorate_sync_async + log = logging.getLogger(__name__) @@ -278,6 +286,52 @@ def _set_prog_hook(self, val: Optional[ProgressHookBase[Any]]) -> None: self._prog_hook = val +def optional_report(func: Callable[..., Any]) -> Callable[..., Any]: + """Decorator for functions that optionally take a 'report' argument + + If the report is not supplied, one will be created automatically and after the + function is called the report's `log_issues` and `check_errors` methods will be + called. If the user supplies the report themselves, it is up to them to call + these methods if they want to. + """ + sig = inspect.signature(func) + report_type = typing.get_type_hints(func)["report"] + if not isinstance(report_type, type) or not issubclass(report_type, BaseReport): + type_stack = deque([report_type]) + report_type = None + while type_stack: + curr_type = type_stack.popleft() + for sub_type in get_args(curr_type): + if isinstance(sub_type, type): + if issubclass(sub_type, BaseReport): + report_type = sub_type + break + else: + type_stack.append(sub_type) + if report_type is None: + raise ValueError(f"The 'report' arg isn't the correct type: {report_type}") + + @contextmanager + def check_report( + args: Iterable[Any], kwargs: Dict[str, Any] + ) -> Iterator[Args_Type]: + bound_args = sig.bind(*args, **kwargs) + bound_args.apply_defaults() + report = bound_args.arguments["report"] + if report is None: + extern_report = False + report = report_type() + bound_args.arguments["report"] = report + else: + extern_report = True + yield (bound_args.args, bound_args.kwargs) + if not extern_report: + report.log_issues() + report.check_errors() + + return decorate_sync_async(check_report, func) + + class MultiError(Exception): def __init__(self, errors: List[Exception]): self.errors = errors diff --git a/dcm/route.py b/dcm/route.py index a9a5a3d..7278dfc 100644 --- a/dcm/route.py +++ b/dcm/route.py @@ -38,6 +38,7 @@ MultiDictReport, MultiKeyedError, ProgressHookBase, + optional_report, ) from .util import DuplicateDataError, TomlConfigurable from .net import DicomOpReport, IncomingDataError, IncomingErrorType @@ -870,17 +871,14 @@ async def route( await data_q.put(None) await route_task + @optional_report async def _route( self, data_q: "asyncio.Queue[Optional[Dataset]]", keep_errors: Union[bool, Tuple[IncomingErrorType, ...]], report: Optional[DynamicTransferReport], ) -> None: - if report is None: - extern_report = False - report = DynamicTransferReport() - else: - extern_report = True + assert report is not None report.keep_errors = keep_errors # type: ignore assoc_cache = SendAssociationCache(self._assoc_cache_time) try: @@ -897,6 +895,7 @@ async def _route( # What happens if a user pushes None accidentally? Just # use a different sentinel value? if ds is None: + log.debug("The route task got None and is shutting down") break filter_dest_map = self.get_filter_dest_map(ds) n_filt = len([f for f in filter_dest_map if f is not None]) @@ -939,9 +938,6 @@ async def _route( await assoc_cache.empty_cache() report.done = True log.debug("Done with routing") - if not extern_report: - report.log_issues() - report.check_errors() async def _fill_qr( self, diff --git a/dcm/store/base.py b/dcm/store/base.py index ef1e7c6..45302c8 100644 --- a/dcm/store/base.py +++ b/dcm/store/base.py @@ -1,3 +1,4 @@ +"""Base classes and shared functionality for storage abstractions""" from __future__ import annotations import os, enum, logging, asyncio from contextlib import asynccontextmanager @@ -24,7 +25,7 @@ import pydicom from pydicom import Dataset -from ..query import QueryLevel, QueryResult, uid_elems +from ..query import QueryLevel, QueryResult, UID_ELEMS from ..net import ( DcmNode, DicomOpReport, @@ -218,7 +219,7 @@ def clear(self) -> None: def is_valid_dicom(ds: Dataset) -> bool: - for uid_elem in uid_elems.values(): + for uid_elem in UID_ELEMS.values(): if not hasattr(ds, uid_elem): return False return True @@ -270,9 +271,46 @@ async def gen_paths_and_data(self) -> AsyncIterator[Tuple[PathInputType, Dataset self.report.done = True +class LocalRepoChunk(RepoChunk): + """Repo chunk for local files""" + + report: IncomingDataReport + + def __init__( + self, repo: "LocalRepo", qr: QueryResult, description: Optional[str] = None + ): + self.repo = repo + self.qr = qr + self.description = description + if description is not None: + rep_descr = description + else: + rep_descr = str(self.repo.root_path) + self.report = IncomingDataReport( + description=rep_descr, + n_expected=self.n_expected, + ) + + async def gen_data(self) -> AsyncIterator[Dataset]: + async for _, ds in self.gen_paths_and_data(): + yield ds + + async def gen_paths_and_data(self) -> AsyncIterator[Tuple[PathInputType, Dataset]]: + """Generate both the paths and the corresponding data sets""" + loop = asyncio.get_running_loop() + for min_ds in self.qr: + ds_path = min_ds.StorageURL + ds = await loop.run_in_executor(None, _read_f, ds_path) + if not self.report.add(ds): + continue + yield ds_path, ds + self.report.done = True + + T_chunk = TypeVar("T_chunk", bound=DataChunk, covariant=True) T_qreport = TypeVar( - "T_qreport", bound=Union[CountableReport, SummaryReport[Any]], contravariant=True + "T_qreport", + bound=Union[CountableReport, SummaryReport[Any]], ) T_rreport = TypeVar("T_rreport", bound=CountableReport, contravariant=True) T_sreport = TypeVar( @@ -339,6 +377,8 @@ class DataRepo( ): """Protocol for stores with query/retrieve functionality""" + query_report_type: Type[T_qreport] + async def queries( self, level: Optional[QueryLevel] = None, @@ -404,6 +444,8 @@ class DcmRepo( TransferMethod.REMOTE_COPY, ) + query_report_type: Type[MultiListReport[DicomOpReport]] = MultiListReport + @property def remote(self) -> DcmNode: raise NotImplementedError @@ -428,6 +470,18 @@ def __repr__(self) -> str: return f"DcmRepo({self.remote})" +class LocalStore(Protocol): + + _root_path: Path + + @property + def root_path(self) -> Path: + return self._root_path + + def __repr__(self) -> str: + return f"{self.__class__.__name__}({self.root_path})" + + class LocalWriteError(Exception): def __init__(self, write_errors: Dict[Exception, List[PathInputType]]): self.write_errors = write_errors @@ -498,6 +552,7 @@ def clear(self) -> None: class LocalBucket( + LocalStore, DataBucket[LocalChunk, LocalWriteReport], OobCapable[LocalChunk, LocalWriteReport], Protocol, @@ -518,8 +573,105 @@ class LocalBucket( TransferMethod.MOVE, ) + @asynccontextmanager + async def send( + self, report: Optional[LocalWriteReport] = None + ) -> AsyncIterator["janus._AsyncQueueProxy[Dataset]"]: + """Produces a Queue that you can put data sets into for storage""" + raise NotImplementedError + yield + + def get_empty_send_report(self) -> LocalWriteReport: + return LocalWriteReport(meta_data={"root_path": self.root_path}) + + def get_empty_oob_report(self) -> LocalWriteReport: + return LocalWriteReport(meta_data={"root_path": self.root_path}) + + +class LocalQueryReport(CountableReport): + """Capture info about query operation against LocalRepo""" + + def __init__( + self, + description: Optional[str] = None, + meta_data: Optional[Dict[str, Any]] = None, + depth: int = 0, + prog_hook: Optional[ProgressHookBase[Any]] = None, + n_expected: Optional[int] = None, + ): + self._inconsistent: List[Dataset] = [] + super().__init__(description, meta_data, depth, prog_hook, n_expected) + + def add_inconsistent(self, ds: Dataset) -> None: + self._inconsistent.append(ds) + self.count_input() + + def add_success(self, ds: Dataset) -> None: + self.count_input() + @property - def root_path(self) -> Path: + def n_success(self) -> int: + """Number of successfully handled inputs""" + return self._n_input - self.n_warnings + + @property + def n_errors(self) -> int: + """Number of errors, where a single input can cause multiple errors""" + return 0 + + @property + def n_warnings(self) -> int: + """Number of warnings, where a single input can cause multiple warnings""" + return len(self._inconsistent) + + def log_issues(self) -> None: + if self.n_warnings != 0: + log.warning("There were %d inconsitent query data sets", self.n_warnings) + + def check_errors(self) -> None: + pass + + +class IndexInitMode(enum.IntEnum): + """Define how to handle locally managed indices on initialization""" + + ASSUME_CLEAN = 0 # Do nothing + CHECK_INDEXED = 1 # Check if all indexed files exist + SCRUB_INDEXED = 2 # Make sure indexed files exist and meta data matches + + +class LocalRepo( + LocalStore, + DataRepo[LocalRepoChunk, LocalQueryReport, LocalWriteReport, IncomingDataReport], + OobCapable[LocalRepoChunk, LocalWriteReport], + Protocol, +): + """Abstract base class for local files with some sort of meta data index""" + + _supported_methods: Tuple[TransferMethod, ...] = ( + TransferMethod.PROXY, + TransferMethod.LINK, + ) + + _streaming_methods: Tuple[TransferMethod, ...] = ( + TransferMethod.PROXY, + TransferMethod.LINK, + ) + + @classmethod + async def is_repo(cls, path: PathInputType) -> bool: + """Return True if the path looks like a repo for the subclass""" + raise NotImplementedError + + @classmethod + async def build( + cls: Type["LocalRepo"], + path: PathInputType, + index_init: IndexInitMode = IndexInitMode.ASSUME_CLEAN, + scan_fs: bool = False, + **init_kwargs: Any, + ) -> LocalRepo: + """Build LocalRepo with control of index initialization and new file handling""" raise NotImplementedError @asynccontextmanager @@ -535,6 +687,3 @@ def get_empty_send_report(self) -> LocalWriteReport: def get_empty_oob_report(self) -> LocalWriteReport: return LocalWriteReport(meta_data={"root_path": self.root_path}) - - def __repr__(self) -> str: - return f"LocalDir({self.root_path})" diff --git a/dcm/store/local_dir.py b/dcm/store/local_dir.py index 96e7ee5..eb917c3 100644 --- a/dcm/store/local_dir.py +++ b/dcm/store/local_dir.py @@ -5,12 +5,26 @@ from glob import iglob from pathlib import Path from queue import Empty -from typing import Optional, AsyncIterator, Any, Callable, Tuple, cast, Union, Dict +from typing import ( + Iterable, + Optional, + AsyncIterator, + Any, + Callable, + Tuple, + cast, + Union, + Dict, + FrozenSet, +) from pydicom.dataset import Dataset import janus +from dcm.report import optional_report + from .base import LocalBucket, TransferMethod, LocalChunk, LocalWriteReport +from ..query import MIN_ATTRS, InconsistentDataError, QueryResult, minimal_copy from ..util import fstr_eval, PathInputType, InlineConfigurable, create_thread_task @@ -23,6 +37,7 @@ def _dir_crawl_worker( recurse: bool = True, file_ext: str = "dcm", max_chunk: int = 1000, + skip_paths: Optional[FrozenSet[str]] = None, shutdown: Optional[threading.Event] = None, ) -> None: curr_files = [] @@ -31,7 +46,7 @@ def _dir_crawl_worker( if recurse: glob_comps.append("**") if file_ext: - glob_comps.append("*.%s" % file_ext) + glob_comps.append(f"*.{file_ext}") else: glob_comps.append("*") glob_exp = os.path.join(*glob_comps) @@ -41,6 +56,8 @@ def _dir_crawl_worker( return if not os.path.isfile(path): continue + if skip_paths is not None and path in skip_paths: + continue curr_files.append(path) if len(curr_files) == max_chunk: res_q.put(LocalChunk(curr_files)) @@ -72,6 +89,8 @@ def _disk_write_worker( out_fmt: str, force_overwrite: bool, report: LocalWriteReport, + dest_qr: Optional[QueryResult] = None, + include_elems: Iterable[str] = MIN_ATTRS, shutdown: Optional[threading.Event] = None, ) -> None: """Take data sets from a queue and write to disk""" @@ -93,9 +112,24 @@ def _disk_write_worker( if no_input: continue log.debug("disk_writer thread got a data set") + dupe = False + if dest_qr is not None: + try: + dupe = ds in dest_qr + except InconsistentDataError: + pass # TODO + if dupe: + continue out_path = root_path / make_out_path(out_fmt, ds) if os.path.exists(out_path): + if dest_qr is not None and not dupe: + log.warning( + "Skipping existing file that is not in destinaion index: %s", + out_path, + ) + report.add_skipped(out_path) + continue if force_overwrite: log.debug("File exists, overwriting: %s", out_path) else: @@ -112,6 +146,10 @@ def _disk_write_worker( except Exception as e: report.add_error(out_path, e) else: + if dest_qr is not None: + min_ds = minimal_copy(ds, include_elems) + min_ds.StorageURL = str(out_path) + dest_qr.add(min_ds) report.add_success(out_path) @@ -168,6 +206,18 @@ def _oob_transfer_worker( os.remove(existing_backup) +def get_root_dir(in_path: PathInputType, make_missing: bool = True) -> Path: + res = Path(in_path).expanduser() + if not res.exists(): + if make_missing: + res.mkdir(parents=True) + else: + raise ValueError(f"Path doesn't exist: {res}") + elif not res.is_dir(): + raise ValueError(f"Path is a file not a directory: {res}") + return res + + class LocalDir(LocalBucket, InlineConfigurable["LocalDir"]): """Local directory of data without any additional meta data""" @@ -193,14 +243,7 @@ def __init__( force_overwrite: bool = False, make_missing: bool = True, ): - self._root_path = Path(path).expanduser() - if not self._root_path.exists(): - if make_missing: - self._root_path.mkdir(parents=True) - else: - raise ValueError(f"Path doesn't exist: {self._root_path}") - elif not self._root_path.is_dir(): - raise ValueError(f"Path is a file not a directory: {self._root_path}") + self._root_path = get_root_dir(path, make_missing) self._recurse = recurse self._max_chunk = max_chunk self._force_overwrite = force_overwrite @@ -239,13 +282,6 @@ def inline_to_dict(in_str: str) -> Dict[str, Any]: raise ValueError(f"Invalid short form for LocalDir: {in_str}") return res - @property - def root_path(self) -> Path: - return self._root_path - - def __str__(self) -> str: - return f"LocalDir({self._root_path})" - async def gen_chunks(self) -> AsyncIterator[LocalChunk]: res_q: janus.Queue[LocalChunk] = janus.Queue() crawl_fut = create_thread_task( @@ -270,14 +306,11 @@ async def gen_chunks(self) -> AsyncIterator[LocalChunk]: await crawl_fut @asynccontextmanager + @optional_report async def send( self, report: Optional[LocalWriteReport] = None ) -> AsyncIterator["janus._AsyncQueueProxy[Dataset]"]: - if report is None: - extern_report = False - report = LocalWriteReport() - else: - extern_report = True + assert report is not None report._meta_data["root_path"] = self._root_path send_q: janus.Queue[Optional[Dataset]] = janus.Queue(10) send_fut = create_thread_task( @@ -300,23 +333,17 @@ async def send( await send_fut log.debug("The disk writer thread has finished") report.done = True - if not extern_report: - report.log_issues() - report.check_errors() + @optional_report async def oob_transfer( self, method: TransferMethod, chunk: LocalChunk, report: Optional[LocalWriteReport] = None, ) -> None: + assert report is not None if method is TransferMethod.PROXY or method not in self._supported_methods: raise ValueError(f"Invalid transfer method: {method}") - if report is None: - extern_report = False - report = LocalWriteReport() - else: - extern_report = True report._meta_data["root_path"] = self._root_path # At least for now, python seeems to lack the ability to define only # the required args to a callable while ignoring kwargs @@ -346,6 +373,3 @@ async def oob_transfer( await oob_fut log.info("Oob transfer worker shutdown, marking report done") report.done = True - if not extern_report: - report.log_issues() - report.check_errors() diff --git a/dcm/store/qr_repo.py b/dcm/store/qr_repo.py new file mode 100644 index 0000000..6588588 --- /dev/null +++ b/dcm/store/qr_repo.py @@ -0,0 +1,419 @@ +"""Lightweight local DataRepo that just keeps a JSON serialized QueryResult around""" + +import asyncio, json, logging, fnmatch, re, inspect +from datetime import datetime +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from contextlib import asynccontextmanager +from copy import deepcopy +from typing import ( + AsyncIterator, + FrozenSet, + Iterable, + List, + Optional, + Set, + Type, + Dict, + Any, +) +from typing_extensions import Protocol + +import janus +import flufl.lock +from pydicom import Dataset + +from ..util import PathInputType, InlineConfigurable, atomic_open, create_thread_task +from ..query import ( + MIN_ATTRS, + UID_ELEMS, + DataNode, + DataPath, + InconsistentDataError, + QueryLevel, + QueryResult, + expand_queries, + get_level_and_query, + minimal_copy, +) +from ..net import IncomingDataReport +from ..report import optional_report +from .local_dir import _dir_crawl_worker, _disk_write_worker, LocalDir, get_root_dir +from .base import ( + LocalChunk, + LocalRepo, + LocalRepoChunk, + IndexInitMode, + LocalWriteReport, + LocalQueryReport, + _read_f, +) + + +log = logging.getLogger(__name__) + + +class InvalidQrRepoError(Exception): + pass + + +DEFUALT_INDEX_ELEMS = MIN_ATTRS + ( + "EchoTime", + "InversionTime", + "RepetitionTime", + "FlipAngle", + "BodyPartExamined", +) + + +class _SingletonQrRepo(type(Protocol)): # type: ignore + """Make sure we have a single QrRepo for each root path""" + + _instances = {} # type: ignore + _init = {} # type: ignore + + def __init__(cls, name, bases, dct): # type: ignore + super(_SingletonQrRepo, cls).__init__(name, bases, dct) + cls._init[cls] = dct.get("__init__", None) + + def __call__(cls, *args, **kwargs): # type: ignore + init = cls._init[cls] + local_path = inspect.getcallargs(init, None, *args, **kwargs)["path"] + if local_path is None: + raise ValueError("The 'path' arg can't be None") + key = (cls, local_path) + if key not in cls._instances: + cls._instances[key] = super(_SingletonQrRepo, cls).__call__(*args, **kwargs) + return cls._instances[key] + + +class FsLockTimeout(Exception): + """Raised when we timeout trying to acquire a filesystem lock""" + + +class QrRepo(LocalRepo, InlineConfigurable["QrRepo"], metaclass=_SingletonQrRepo): + """Simple local data repository using QueryResult JSON serialization for persistence""" + + default_out_fmt = LocalDir.default_out_fmt + + chunk_type: Type[LocalRepoChunk] = LocalRepoChunk + + query_report_type: Type[LocalQueryReport] = LocalQueryReport + + def __init__( + self, + path: PathInputType, + index_elems: Iterable[str] = DEFUALT_INDEX_ELEMS, + file_ext: str = "dcm", + max_chunk: int = 1000, + out_fmt: Optional[str] = None, + overwrite_existing: bool = False, + make_missing: bool = True, + max_sync_time: int = 30, + ): + self._root_path = get_root_dir(path, make_missing) + self._index_elems = index_elems + self._max_chunk = max_chunk + if out_fmt is None: + self._out_fmt = self.default_out_fmt + else: + self._out_fmt = out_fmt + self._file_ext = file_ext + if self._file_ext: + self._out_fmt += ".%s" % file_ext + self._overwrite = overwrite_existing + self.description = str(self._root_path) + self._qr_path = self._root_path / "dcm_meta.json" + self._lock_path = self._root_path / ".dcm_meta.lock" + self._max_sync_time = max_sync_time + self._fs_lock = flufl.lock.Lock( + str(self._lock_path), + lifetime=max_sync_time, + default_timeout=max_sync_time * 5, + ) + self._sync_pool = ThreadPoolExecutor(1) + if self._qr_path.exists(): + with open(self._qr_path, "rt") as qr_f: + self._qr = QueryResult.from_json_dict(json.load(qr_f)) + if self._qr.level != QueryLevel.IMAGE: + raise InvalidQrRepoError() + else: + if len(list(self._root_path.iterdir())) != 0: + log.warning("No QueryResult JSON found in non-empty dir") + self._qr = QueryResult(level=QueryLevel.IMAGE) + self._path_set: Optional[FrozenSet[str]] = None + + @staticmethod + def inline_to_dict(in_str: str) -> Dict[str, Any]: + """Parse inline string format 'path[:out_fmt][:file_ext]' + + Both the second components are optional + """ + return LocalDir.inline_to_dict(in_str) + + @classmethod + async def is_repo(cls, path: PathInputType) -> bool: + try: + root_path = get_root_dir(path) + except: + return False + qr_path = root_path / "dcm_meta.json" + loop = asyncio.get_running_loop() + return await loop.run_in_executor(None, qr_path.exists) + + @classmethod + async def build( + cls: Type["QrRepo"], + path: PathInputType, + index_init: IndexInitMode = IndexInitMode.ASSUME_CLEAN, + scan_fs: bool = False, + **init_kwargs: Any, + ) -> "QrRepo": + """Build LocalRepo with control of index initialization and new file handling""" + repo = QrRepo(path, **init_kwargs) + if index_init != IndexInitMode.ASSUME_CLEAN: + # TODO: Handle index initialization options + # TODO: Probably also need another option ("update_on_mismatch"?) to control + # whether we update the QR or not when we notice a missing file or a + # file where meta data has changed. + raise NotImplementedError + if scan_fs: + log.info("About to scan FS for new files") + loop = asyncio.get_running_loop() + res_q: janus.Queue[LocalChunk] = janus.Queue() + crawl_fut = create_thread_task( + _dir_crawl_worker, + ( + res_q.sync_q, + repo._root_path, + True, + repo._file_ext, + repo._max_chunk, + repo.path_set, + ), + ) + n_found = 0 + while True: + try: + chunk = await asyncio.wait_for(res_q.async_q.get(), timeout=1.0) + except asyncio.TimeoutError: + # Check if the worker thread exited + if crawl_fut.done(): + break + else: + async for ds_path, ds in chunk.gen_paths_and_data(): + try: + dupe = ds in repo._qr + except InconsistentDataError: + log.warn("Refusing to index inconsistent data: %s", ds_path) + continue + if not dupe: + log.debug("Found a new file to index: %s", ds_path) + min_ds = minimal_copy(ds, repo._index_elems) + min_ds.StorageURL = str(ds_path) + repo._qr.add(min_ds) + n_found += 1 + await crawl_fut + log.info(f"Found {n_found} new files") + if n_found: + repo._path_set = None + await repo.sync() + return repo + + @property + def path_set(self) -> FrozenSet[str]: + """Get a set of all paths currently indexed""" + if self._path_set is None: + self._path_set = frozenset(ds.StorageURL for ds in self._qr) + return self._path_set + + async def gen_chunks(self) -> AsyncIterator[LocalRepoChunk]: + """Generate the data in this bucket, one chunk at a time""" + for chunk_qr in self._qr.chunk(self._max_chunk): + yield LocalRepoChunk(self, chunk_qr) + + async def gen_query_chunks( + self, query_res: QueryResult + ) -> AsyncIterator[LocalRepoChunk]: + """Generate chunks of data corresponding to `query_res`""" + matched = self._qr & query_res + for chunk_qr in matched.chunk(self._max_chunk): + yield LocalRepoChunk(self, chunk_qr) + + @asynccontextmanager + @optional_report + async def send( + self, report: Optional[LocalWriteReport] = None + ) -> AsyncIterator["janus._AsyncQueueProxy[Dataset]"]: + """Produces a Queue that you can put data sets into for storage""" + assert report is not None + loop = asyncio._get_running_loop() + report._meta_data["root_path"] = self._root_path + send_q: janus.Queue[Optional[Dataset]] = janus.Queue(10) + send_fut = create_thread_task( + _disk_write_worker, + ( + send_q.sync_q, + self._root_path, + self._out_fmt, + self._overwrite, + report, + self._qr, + self._index_elems, + ), + ) + try: + yield send_q.async_q # type: ignore + finally: + if not send_fut.done(): + await send_q.async_q.put(None) + log.debug("awaiting send_fut") + await send_fut + log.debug("awaited send_fut") + report.done = True + self._path_set = None + await self.sync() + + @optional_report + async def queries( + self, + level: Optional[QueryLevel] = None, + query: Optional[Dataset] = None, + query_res: Optional[QueryResult] = None, + report: Optional[LocalQueryReport] = None, + ) -> AsyncIterator[QueryResult]: + """Returns async generator that produces partial QueryResult objects""" + # If QueryResult was given we potentially generate multiple + # queries, one for each dataset referenced by the QueryResult + assert report is not None + level, query = get_level_and_query(level, query, query_res) + queries, queried_elems = expand_queries(level, query, query_res) + res = QueryResult(level) + for sub_query in queries: + # Iterate through QueryLevels, narrowing results when able + log.debug("Processing query: %s", sub_query) + last_incl: Optional[Set[DataNode]] = None + incl_nodes: Optional[List[DataNode]] = None + for curr_lvl, uid_elem in UID_ELEMS.items(): + uid_qval = sub_query.get(uid_elem, "*") + if uid_qval == "": + uid_qval = "*" + # TODO: Need to recognize date ranges here too + if "*" not in uid_qval: + log.debug( + "Checking at level %s for exact_uid %s", curr_lvl, uid_qval + ) + # We have precise (one or none) match specification + try: + matched_path = self._qr.get_path(DataNode(curr_lvl, uid_qval)) + except KeyError: + break + if ( + last_incl is not None + and matched_path.parent.end not in last_incl + ): + # We found a match but with the wrong parent + report.add_inconsistent(sub_query) + break + incl_nodes = [matched_path.end] + elif uid_qval == "*": + # We are matching everything at this level + if last_incl is not None: + incl_nodes = [] + for parent in last_incl: + incl_nodes += self._qr.children(parent) + else: + # We are matching zero or more at this level + incl_nodes = [] + uid_regex = re.compile(fnmatch.translate(uid_qval)) + if last_incl is not None: + for parent in last_incl: + for child in self._qr.children(parent): + if uid_regex.match(child.uid): + incl_nodes.append(child) + else: + for sub_path in self._qr.level_paths(level): + if uid_regex.match(sub_path.end.uid): + incl_nodes.append(sub_path.end) + if curr_lvl == level: + # We are done refining results + report.add_success(sub_query) + if incl_nodes is None: + # We matched everything + yield deepcopy(self._qr) + return + res = QueryResult(level) + # TODO: Are underlying datasets getting deepcopied here? + for incl_node in incl_nodes: + res |= self._qr.sub_query(incl_node, level) + yield res + else: + # Prepare to refine results at next QueryLevel + if incl_nodes is not None: + if len(incl_nodes) == 0: + report.add_success(sub_query) + break + last_incl = set(incl_nodes) + incl_nodes = None + + async def query( + self, + level: Optional[QueryLevel] = None, + query: Optional[Dataset] = None, + query_res: Optional[QueryResult] = None, + report: Optional[LocalQueryReport] = None, + ) -> QueryResult: + """Perform a query against the data repo""" + level, query = get_level_and_query(level, query, query_res) + res = QueryResult(level) + async for sub_res in self.queries(level, query, query_res, report): + res |= sub_res + return res + + @optional_report + async def retrieve( + self, query_res: QueryResult, report: Optional[IncomingDataReport] = None + ) -> AsyncIterator[Dataset]: + """Returns an async generator that will produce datasets""" + assert report is not None + loop = asyncio.get_running_loop() + # TODO: Could use QueryProv to avoid duplicate operation here? + match = self._qr & query_res + for min_ds in match: + ds_path = min_ds.StorageURL + ds = await loop.run_in_executor(None, _read_f, ds_path) + if report.add(ds): + yield ds + report.done = True + + async def sync(self) -> None: + """Sync the in-memory and on disk QueryResults + + Shouldn't need to be called manually unless a FsLockTimeoutError was raised on + a previous operation. + """ + loop = asyncio.get_running_loop() + json_str = json.dumps(self._qr.to_json_dict()) + await loop.run_in_executor(self._sync_pool, self._sync_qr, json_str) + + def _sync_qr(self, json_str: str) -> None: + """Sync the in-memory and on disk QueryResults""" + + # TODO: We probably don't want to set the flufl Lock timeout too high here since + # we don't want to block the thread too long on early shutdown, although + # this shouldn't come up much in practice. + try: + self._fs_lock.lock() + except flufl.lock.TimeOutError: + raise FsLockTimeout("Unable to aquire FS lock, did another process die?") + try: + start = datetime.now() + with atomic_open(self._qr_path, mode="wt") as out_f: + out_f.write(json_str) + if (datetime.now() - start).total_seconds() > self._max_sync_time: + log.warning( + "It took longer than max_sync_time (%s) seconds to sync JSON", + self._max_sync_time, + ) + finally: + self._fs_lock.unlock() diff --git a/dcm/sync.py b/dcm/sync.py index 40cb802..7985692 100644 --- a/dcm/sync.py +++ b/dcm/sync.py @@ -33,7 +33,7 @@ DataNode, get_all_uids, minimal_copy, - uid_elems, + UID_ELEMS, ) from .store.base import ( TransferMethod, @@ -57,13 +57,11 @@ from .diff import diff_data_sets, DataDiff from .net import ( IncomingDataReport, - IncomingDataError, IncomingErrorType, - DicomOpReport, RetrieveReport, ) from .report import ( - BaseReport, + CountableReport, MultiAttrReport, MultiListReport, MultiDictReport, @@ -356,15 +354,11 @@ async def _sync_iter_to_async(sync_gen: Iterator[T]) -> AsyncIterator[T]: yield result -TransferReportTypes = Union[ - DynamicTransferReport, - StaticTransferReport, -] - +TransferReportTypes = Union[DynamicTransferReport, StaticTransferReport] DestType = Union[DataBucket, Route] -SourceMissingQueryReportType = MultiListReport[MultiListReport[DicomOpReport]] +SourceMissingQueryReportType = MultiListReport[CountableReport] DestMissingQueryReportType = MultiDictReport[ DataRepo[Any, Any, Any, Any], SourceMissingQueryReportType @@ -375,6 +369,8 @@ class RepoRequiredError(Exception): """Operation requires a DataRepo but a DataBucket was provided""" +# TODO: This class shouldn't be responsible for building various sub-reports, because it +# doesn't know about the src / dests which determines the report types class SyncQueriesReport(MultiAttrReport): """Report for queries being performed during sync""" @@ -385,7 +381,7 @@ def __init__( depth: int = 0, prog_hook: Optional[ProgressHookBase[Any]] = None, ): - self._init_src_qr_report: Optional[MultiListReport[DicomOpReport]] = None + self._init_src_qr_report: Optional[CountableReport] = None self._missing_src_qr_reports: Optional[ MultiListReport[SourceMissingQueryReportType] ] = None @@ -400,13 +396,15 @@ def __init__( super().__init__(description, meta_data, depth, prog_hook) @property - def init_src_qr_report(self) -> MultiListReport[DicomOpReport]: - if self._init_src_qr_report is None: - self._init_src_qr_report = MultiListReport( - "init-src-qr", depth=self._depth + 1, prog_hook=self._prog_hook - ) + def init_src_qr_report(self) -> Optional[CountableReport]: return self._init_src_qr_report + @init_src_qr_report.setter + def init_src_qr_report(self, val: CountableReport) -> None: + if self._init_src_qr_report is not None: + raise ValueError("Report was already set") + self._init_src_qr_report = val + @property def missing_src_qr_reports(self) -> MultiListReport[SourceMissingQueryReportType]: if self._missing_src_qr_reports is None: @@ -493,6 +491,9 @@ class SyncManager: keep_errors Whether or not we try to sync erroneous data + Can be set to `True` to send all data regardless of the error, or a tuple of + `IncomingErrorType` specifying on which types of errors to send the data + report Allows live introspection of the sync process, detailed results """ @@ -618,7 +619,8 @@ async def gen_transfers( expected_sub_qrs = query_res.get_count(gen_level) else: q = dict_to_ds({elem: "*" for elem in self._router.required_elems}) - qr_report = self.report.queries_report.init_src_qr_report + qr_report = self._src.query_report_type() + self.report.queries_report.init_src_qr_report = qr_report qr_gen = self._src.queries(QueryLevel.STUDY, q, report=qr_report) n_sub_qr = 0 @@ -775,7 +777,7 @@ async def _get_missing( if filt is not None: invertible_uids = filt.invertible_uids can_invert_uids = all( - uid in invertible_uids for uid in uid_elems.values() + uid in invertible_uids for uid in UID_ELEMS.values() ) for dest in route.dests: df_tuple = (dest, filt) @@ -795,13 +797,13 @@ async def _get_missing( # Build multi reports for capturing queries expected = QueryLevel.IMAGE - src_qr.level + 1 - src_queries_report: MultiListReport[ - MultiListReport[DicomOpReport] - ] = MultiListReport("missing-src-qr", n_expected=expected) + src_queries_report: MultiListReport[CountableReport] = MultiListReport( + "missing-src-qr", n_expected=expected + ) self.report.queries_report.missing_src_qr_reports.append(src_queries_report) dest_queries_report: MultiDictReport[ DataRepo[Any, Any, Any, Any], - MultiListReport[MultiListReport[DicomOpReport]], + MultiListReport[CountableReport], ] = MultiDictReport("missing-dests-qr") self.report.queries_report.missing_dest_qr_reports.append(dest_queries_report) @@ -839,7 +841,7 @@ async def _get_missing( if curr_level > curr_src_qr.level: # We need more details for the source QueryResult log.debug("Querying src in _get_missing more details") - src_report: MultiListReport[DicomOpReport] = MultiListReport() + src_report: MultiListReport[CountableReport] = MultiListReport() src_queries_report.append(src_report) src_qr_task = asyncio.create_task( self._src.query( @@ -847,13 +849,13 @@ async def _get_missing( ) ) if dest not in dest_queries_report: - dest_reports: MultiListReport[ - MultiListReport[DicomOpReport] - ] = MultiListReport("missing-dest-qr") + dest_reports: MultiListReport[CountableReport] = MultiListReport( + "missing-dest-qr" + ) dest_queries_report[dest] = dest_reports else: dest_reports = dest_queries_report[dest] - dest_report: MultiListReport[DicomOpReport] = MultiListReport() + dest_report: CountableReport = dest.query_report_type() dest_reports.append(dest_report) dest_qr = await dest.query( level=curr_level, query_res=curr_matching[df], report=dest_report @@ -1036,7 +1038,7 @@ async def sync_data( dests: List[DestType], query: Optional[Union[Dataset, List[Dataset]]] = None, query_res: Optional[List[Optional[QueryResult]]] = None, - query_reports: Optional[List[Optional[MultiListReport[DicomOpReport]]]] = None, + query_reports: Optional[List[Optional[MultiListReport[CountableReport]]]] = None, sm_kwargs: Optional[List[Dict[str, Any]]] = None, dry_run: bool = False, ) -> List[SyncReport]: diff --git a/dcm/tests/conftest.py b/dcm/tests/conftest.py index bad2c41..73e0037 100644 --- a/dcm/tests/conftest.py +++ b/dcm/tests/conftest.py @@ -1,10 +1,10 @@ +import os, time, shutil, tarfile, logging, re, asyncio from dataclasses import dataclass -import os, time, shutil, tarfile, logging, re from copy import deepcopy import subprocess as sp from tempfile import TemporaryDirectory, NamedTemporaryFile from pathlib import Path -from typing import BinaryIO +from typing import BinaryIO, Optional import pydicom from pynetdicom import AE @@ -15,8 +15,10 @@ from ..conf import _default_conf, DcmConfig from ..query import QueryLevel, QueryResult from ..net import DcmNode, _make_default_store_scu_pcs +from ..store.base import IndexInitMode from ..store.net_repo import NetRepo from ..store.local_dir import LocalDir +from ..store.qr_repo import QrRepo logging_opts = {} @@ -184,17 +186,17 @@ class TestNode: node_type: str - dcm_node: DcmNode - init_qr: QueryResult store_dir: Path - proc: sp.Popen + dcm_node: Optional[DcmNode] = None + + proc: Optional[sp.Popen] = None - stdout: BinaryIO + stdout: Optional[BinaryIO] = None - stderr: BinaryIO + stderr: Optional[BinaryIO] = None _is_finalized: bool = False @@ -203,7 +205,7 @@ def is_finalized(self) -> bool: return self._is_finalized def finalize(self): - if not self._is_finalized: + if not self._is_finalized and self.proc is not None: self.proc.terminate() self._is_finalized = True self.stdout.flush() @@ -213,6 +215,26 @@ def finalize(self): return self.stdout.read(), self.stderr.read() +@fixture +def make_qr_repo(get_dicom_subset): + """Factory fixutre to build QrRepo stores""" + curr_dirs = [] + with TemporaryDirectory(prefix="dcm-test") as temp_dir: + temp_dir = Path(temp_dir) + + async def _make_qr_repo(local_node=None, subset="all", **kwargs): + store_dir = temp_dir / f"store_dir{len(curr_dirs)}" + os.makedirs(store_dir) + curr_dirs.append(store_dir) + init_qr, init_data = get_dicom_subset(subset) + for dcm_path, _ in init_data: + shutil.copy(dcm_path, store_dir) + repo = await QrRepo.build(store_dir, scan_fs=True, **kwargs) + return (repo, TestNode("qr", init_qr, store_dir)) + + yield _make_qr_repo + + DCMQRSCP_PATH = shutil.which("dcmqrscp") @@ -368,7 +390,7 @@ def _make_dcmtk_node(clients, subset="all"): proc = sp.Popen(dcmqrscp_args, stdout=sout_f, stderr=serr_f) print("Done") res = TestNode( - "dcmtk", dcmtk_node, init_qr, test_store_dir, proc, sout_f, serr_f + "dcmtk", init_qr, test_store_dir, dcmtk_node, proc, sout_f, serr_f ) nodes.append(res) time.sleep(2) @@ -384,7 +406,7 @@ def _make_dcmtk_node(clients, subset="all"): @fixture def make_dcmtk_net_repo(make_local_node, make_dcmtk_nodes): - def _make_net_repo(local_node=None, clients=[], subset="all"): + async def _make_net_repo(local_node=None, clients=[], subset="all"): if local_node is None: local_node = make_local_node() dcmtk_node = make_dcmtk_nodes([local_node] + clients, subset) @@ -513,7 +535,7 @@ def _make_pnd_node(clients, subset="all"): continue init_assoc.send_c_store(ds) init_assoc.release() - res = TestNode("pnd", pnd_node, init_qr, data_dir, proc, sout_f, serr_f) + res = TestNode("pnd", init_qr, data_dir, pnd_node, proc, sout_f, serr_f) nodes.append(res) return res @@ -527,7 +549,7 @@ def _make_pnd_node(clients, subset="all"): @fixture def make_pnd_net_repo(make_local_node, make_pnd_nodes): - def _make_net_repo(local_node=None, clients=[], subset="all"): + async def _make_net_repo(local_node=None, clients=[], subset="all"): if local_node is None: local_node = make_local_node() pnd_node = make_pnd_nodes([local_node] + clients, subset) @@ -554,11 +576,22 @@ def make_net_repo(node_type, make_dcmtk_net_repo, make_pnd_net_repo): return make_pnd_net_repo +@fixture +def make_repo(node_type, make_dcmtk_net_repo, make_pnd_net_repo, make_qr_repo): + if node_type == "dcmtk": + return make_dcmtk_net_repo + elif node_type == "pnd": + return make_pnd_net_repo + else: + assert node_type == "qr" + return make_qr_repo + + def get_stored_files(store_dir): return [ x for x in Path(store_dir).glob("**/*") - if not x.is_dir() and x.name != "index.dat" + if not x.is_dir() and x.name not in ("index.dat", "dcm_meta.json") ] diff --git a/dcm/tests/store/test_net_repo.py b/dcm/tests/store/test_repo.py similarity index 56% rename from dcm/tests/store/test_net_repo.py rename to dcm/tests/store/test_repo.py index 6223dcc..6db24f7 100644 --- a/dcm/tests/store/test_net_repo.py +++ b/dcm/tests/store/test_repo.py @@ -6,12 +6,15 @@ from ..test_net import get_retr_subsets, get_send_subsets, test_query_subsets +NODE_TYPES = ("dcmtk", "pnd", "qr") + + @mark.asyncio -@mark.parametrize("node_type, subset", get_retr_subsets()) -async def test_gen_chunks(make_net_repo, subset): - net_repo, repo_node = make_net_repo(subset=subset) +@mark.parametrize("node_type, subset", get_retr_subsets(NODE_TYPES)) +async def test_gen_chunks(make_repo, subset): + repo, repo_node = await make_repo(subset=subset) n_dcm_gen = 0 - async for chunk in net_repo.gen_chunks(): + async for chunk in repo.gen_chunks(): async for dcm in chunk.gen_data(): n_dcm_gen += 1 assert dcm in repo_node.init_qr @@ -19,11 +22,11 @@ async def test_gen_chunks(make_net_repo, subset): @mark.asyncio -@mark.parametrize("node_type, subset", get_send_subsets()) -async def test_send(make_net_repo, get_dicom_subset, subset): - net_repo, repo_node = make_net_repo(subset=None) +@mark.parametrize("node_type, subset", get_send_subsets(NODE_TYPES)) +async def test_send(make_repo, get_dicom_subset, subset): + repo, repo_node = await make_repo(subset=None) send_qr, send_data = get_dicom_subset(subset) - async with net_repo.send() as send_q: + async with repo.send() as send_q: for send_path, send_ds in send_data: await send_q.put(send_ds) n_files = len(get_stored_files(repo_node.store_dir)) @@ -31,11 +34,11 @@ async def test_send(make_net_repo, get_dicom_subset, subset): @mark.asyncio -@mark.parametrize("node_type", (pytest.param("dcmtk", marks=has_dcmtk), "pnd")) +@mark.parametrize("node_type", (pytest.param("dcmtk", marks=has_dcmtk), "pnd", "qr")) @mark.parametrize("subset", test_query_subsets) -async def test_query(make_net_repo, get_dicom_subset, subset): - net_repo, repo_node = make_net_repo(subset="all") +async def test_query(make_repo, get_dicom_subset, subset): + repo, repo_node = await make_repo(subset="all") req_qr, _ = get_dicom_subset(subset) req_qr = req_qr & repo_node.init_qr - res_qr = await net_repo.query(query_res=req_qr, level=QueryLevel.IMAGE) + res_qr = await repo.query(query_res=req_qr, level=QueryLevel.IMAGE) assert req_qr == res_qr diff --git a/dcm/tests/test_cli.py b/dcm/tests/test_cli.py index 70f001d..d1b01d2 100644 --- a/dcm/tests/test_cli.py +++ b/dcm/tests/test_cli.py @@ -106,6 +106,7 @@ def _run_forward(config_path, local_node, dest_dir): fwd_args = [ "--config", config_path, + "--debug", "forward", "--inactive-timeout", "10", @@ -113,7 +114,7 @@ def _run_forward(config_path, local_node, dest_dir): local_node, dest_dir, ] - return runner.invoke(cli, fwd_args) + return runner.invoke(cli, fwd_args, catch_exceptions=False) @mark.parametrize("node_type", (pytest.param("dcmtk", marks=has_dcmtk), "pnd")) diff --git a/dcm/tests/test_net.py b/dcm/tests/test_net.py index 2861426..df27fb8 100644 --- a/dcm/tests/test_net.py +++ b/dcm/tests/test_net.py @@ -44,17 +44,16 @@ ] -def get_retr_subsets(): +def get_retr_subsets(node_types=("dcmtk", "pnd")): res = [] - for node_type in ("dcmtk", "pnd"): + for node_type in node_types: if node_type == "dcmtk": for sub in test_retr_subsets: marks = [has_dcmtk] if sub in ("all", "PATIENT-0/STUDY-1/SERIES-2"): marks.append(dcmtk_priv_sop_retr_xfail) res.append(pytest.param(node_type, sub, marks=marks)) - else: - assert node_type == "pnd" + elif node_type == "pnd": for sub in test_retr_subsets: if sub in ("all", "PATIENT-0/STUDY-1/SERIES-2"): res.append( @@ -66,6 +65,9 @@ def get_retr_subsets(): ) else: res.append((node_type, sub)) + else: + for sub in test_retr_subsets: + res.append((node_type, sub)) return res @@ -81,17 +83,16 @@ def get_retr_subsets(): ] -def get_send_subsets(): +def get_send_subsets(node_types=("dcmtk", "pnd")): res = [] - for node_type in ("dcmtk", "pnd"): + for node_type in node_types: if node_type == "dcmtk": for sub in test_retr_subsets: marks = [has_dcmtk] if sub in ("all", "PATIENT-0/STUDY-1/SERIES-2"): marks.append(dcmtk_priv_sop_send_xfail) res.append(pytest.param(node_type, sub, marks=marks)) - else: - assert node_type == "pnd" + elif node_type == "pnd": for sub in test_retr_subsets: if sub in ("all", "PATIENT-0/STUDY-1/SERIES-2"): res.append( @@ -103,6 +104,9 @@ def get_send_subsets(): ) else: res.append((node_type, sub)) + else: + for sub in test_retr_subsets: + res.append((node_type, sub)) return res diff --git a/dcm/tests/test_query.py b/dcm/tests/test_query.py index ca5ea20..6d405c7 100644 --- a/dcm/tests/test_query.py +++ b/dcm/tests/test_query.py @@ -2,7 +2,7 @@ from pytest import fixture, mark from pydicom.dataset import Dataset -from ..query import QueryLevel, QueryProv, QueryResult, req_elems +from ..query import QueryLevel, QueryProv, QueryResult, REQ_ELEMS def make_dataset(attrs=None, level=QueryLevel.IMAGE): @@ -10,7 +10,7 @@ def make_dataset(attrs=None, level=QueryLevel.IMAGE): attrs = {} ds = Dataset() for lvl in QueryLevel: - for attr in req_elems[lvl]: + for attr in REQ_ELEMS[lvl]: val = attrs.get(attr) if val is None: if attr == "PatientID": diff --git a/dcm/tests/test_route.py b/dcm/tests/test_route.py index 302d7dc..7d48447 100644 --- a/dcm/tests/test_route.py +++ b/dcm/tests/test_route.py @@ -25,20 +25,22 @@ def lookup_func(ds): ( pytest.param("dcmtk", ["all", None, None, None], marks=has_dcmtk), ("pnd", ["all", None, None, None]), + ("qr", ["all", None, None, None]), ), ) -def test_pre_route(make_local_node, make_net_repo, node_subsets): +@mark.asyncio +async def test_pre_route(make_local_node, make_repo, node_subsets): local_node = make_local_node() - src_repo, _ = make_net_repo(local_node, subset=node_subsets[0]) - dest1_repo, _ = make_net_repo(local_node, subset=node_subsets[1]) - dest2_repo, _ = make_net_repo(local_node, subset=node_subsets[2]) - dest3_repo, _ = make_net_repo(local_node, subset=node_subsets[3]) + src_repo, _ = await make_repo(local_node, subset=node_subsets[0]) + dest1_repo, _ = await make_repo(local_node, subset=node_subsets[1]) + dest2_repo, _ = await make_repo(local_node, subset=node_subsets[2]) + dest3_repo, _ = await make_repo(local_node, subset=node_subsets[3]) static_route = StaticRoute([dest1_repo]) dyn_route = DynamicRoute( make_id_lookup(dest2_repo, dest3_repo), required_elems=["PatientID"] ) router = Router([static_route, dyn_route]) - res = asyncio.run(router.pre_route(src_repo)) + res = await router.pre_route(src_repo) for routes, qr in res.items(): print("%s -> %s" % ([str(r) for r in routes], qr)) assert all(isinstance(r, StaticRoute) for r in routes) @@ -74,14 +76,16 @@ def lookup_func(ds): ( pytest.param("dcmtk", ["all", None, None, None], marks=has_dcmtk), ("pnd", ["all", None, None, None]), + ("qr", ["all", None, None, None]), ), ) -def test_pre_route_with_dl(make_local_node, make_net_repo, node_subsets): +@mark.asyncio +async def test_pre_route_with_dl(make_local_node, make_repo, node_subsets): local_node = make_local_node() - src_repo, src_node = make_net_repo(local_node, subset=node_subsets[0]) - dest1_repo, _ = make_net_repo(local_node, subset=node_subsets[1]) - dest2_repo, _ = make_net_repo(local_node, subset=node_subsets[2]) - dest3_repo, _ = make_net_repo(local_node, subset=node_subsets[3]) + src_repo, src_node = await make_repo(local_node, subset=node_subsets[0]) + dest1_repo, _ = await make_repo(local_node, subset=node_subsets[1]) + dest2_repo, _ = await make_repo(local_node, subset=node_subsets[2]) + dest3_repo, _ = await make_repo(local_node, subset=node_subsets[3]) static_route = StaticRoute([dest1_repo]) # Setup a dynamic route where we route on an element that can't be queried for # thus forcing the router to download example data sets @@ -91,7 +95,7 @@ def test_pre_route_with_dl(make_local_node, make_net_repo, node_subsets): required_elems=["EchoTime"], ) router = Router([static_route, dyn_route]) - res = asyncio.run(router.pre_route(src_repo)) + res = await router.pre_route(src_repo) for routes, qr in res.items(): print("%s -> %s" % ([str(r) for r in routes], qr)) assert all(isinstance(r, StaticRoute) for r in routes) diff --git a/dcm/tests/test_sync.py b/dcm/tests/test_sync.py index 7093bbb..d41899e 100644 --- a/dcm/tests/test_sync.py +++ b/dcm/tests/test_sync.py @@ -43,14 +43,13 @@ def lookup_func(ds): ] -def get_gen_transfer_sets(): +def get_gen_transfer_sets(node_types=("dcmtk", "pnd", "qr")): res = [] - for node_type in ("dcmtk", "pnd"): + for node_type in node_types: if node_type == "dcmtk": for sub in gen_transfer_sets: res.append(pytest.param(node_type, sub, marks=has_dcmtk)) else: - assert node_type == "pnd" for sub in gen_transfer_sets: res.append((node_type, sub)) return res @@ -68,9 +67,9 @@ def get_gen_transfer_sets(): ] -def get_repo_to_repo_subsets(): +def get_repo_to_repo_subsets(node_types=("dcmtk", "pnd", "qr")): res = [] - for node_type in ("dcmtk", "pnd"): + for node_type in node_types: if node_type == "dcmtk": for sub in sync_subsets: marks = [has_dcmtk] @@ -78,34 +77,35 @@ def get_repo_to_repo_subsets(): marks += priv_sop_marks res.append(pytest.param(node_type, sub, marks=marks)) else: - assert node_type == "pnd" for sub in sync_subsets: res.append((node_type, sub)) return res -def get_bucket_to_repo_subsets(): +def get_bucket_to_repo_subsets(node_types=("dcmtk", "pnd", "qr")): res = [] - for node_type in ("dcmtk", "pnd"): + for node_type in node_types: if node_type == "dcmtk": for sub in sync_subsets: marks = [has_dcmtk] + priv_sop_marks res.append(pytest.param(node_type, sub, marks=marks)) - else: - assert node_type == "pnd" + elif node_type == "pnd": for sub in sync_subsets: res.append(pytest.param(node_type, sub, marks=pnd_priv_sop_xfail)) + else: + for sub in sync_subsets: + res.append((node_type, sub)) return res @mark.parametrize("node_type, subset_specs", get_gen_transfer_sets()) @mark.asyncio -async def test_gen_transfers(make_local_node, make_net_repo, subset_specs): +async def test_gen_transfers(make_local_node, make_repo, subset_specs): local_node = make_local_node() - src_repo, src_node = make_net_repo(local_node, subset="all") - dest1_repo, dest1_node = make_net_repo(local_node, subset=subset_specs[0]) - dest2_repo, dest2_node = make_net_repo(local_node, subset=subset_specs[1]) - dest3_repo, dest3_node = make_net_repo(local_node, subset=subset_specs[2]) + src_repo, src_node = await make_repo(local_node, subset="all") + dest1_repo, dest1_node = await make_repo(local_node, subset=subset_specs[0]) + dest2_repo, dest2_node = await make_repo(local_node, subset=subset_specs[1]) + dest3_repo, dest3_node = await make_repo(local_node, subset=subset_specs[2]) static_route = StaticRoute([dest1_repo]) dyn_lookup = make_lookup(dest2_repo, dest3_repo) dyn_route = DynamicRoute(dyn_lookup, required_elems=["PatientID"]) @@ -144,10 +144,10 @@ async def test_gen_transfers(make_local_node, make_net_repo, subset_specs): @mark.parametrize("node_type, subset_specs", get_repo_to_repo_subsets()) @mark.asyncio -async def test_repo_sync_single_static(make_local_node, make_net_repo, subset_specs): +async def test_repo_sync_single_static(make_local_node, make_repo, subset_specs): local_node = make_local_node() - src_repo, src_node = make_net_repo(local_node, subset="all") - dest1_repo, dest1_node = make_net_repo(local_node, subset=subset_specs[0]) + src_repo, src_node = await make_repo(local_node, subset="all") + dest1_repo, dest1_node = await make_repo(local_node, subset=subset_specs[0]) static_route = StaticRoute([dest1_repo]) dests = [static_route] async with SyncManager(src_repo, dests) as sm: @@ -164,12 +164,12 @@ async def test_repo_sync_single_static(make_local_node, make_net_repo, subset_sp @mark.parametrize("node_type, subset_specs", get_repo_to_repo_subsets()) @mark.asyncio -async def test_repo_sync_multi(make_local_node, make_net_repo, subset_specs): +async def test_repo_sync_multi(make_local_node, make_repo, subset_specs): local_node = make_local_node() - src_repo, src_node = make_net_repo(local_node, subset="all") - dest1_repo, dest1_node = make_net_repo(local_node, subset=subset_specs[0]) - dest2_repo, _ = make_net_repo(local_node, subset=subset_specs[1]) - dest3_repo, _ = make_net_repo(local_node, subset=subset_specs[2]) + src_repo, src_node = await make_repo(local_node, subset="all") + dest1_repo, dest1_node = await make_repo(local_node, subset=subset_specs[0]) + dest2_repo, _ = await make_repo(local_node, subset=subset_specs[1]) + dest3_repo, _ = await make_repo(local_node, subset=subset_specs[2]) static_route = StaticRoute([dest1_repo]) dyn_route = DynamicRoute( make_lookup(dest2_repo, dest3_repo), required_elems=["PatientID"] @@ -190,14 +190,12 @@ async def test_repo_sync_multi(make_local_node, make_net_repo, subset_specs): @mark.parametrize("node_type, subset_specs", get_bucket_to_repo_subsets()) @mark.asyncio -async def test_bucket_sync( - make_local_dir, make_local_node, make_net_repo, subset_specs -): +async def test_bucket_sync(make_local_dir, make_local_node, make_repo, subset_specs): src_bucket, init_qr, _ = make_local_dir("all", max_chunk=2) local_node = make_local_node() - dest1_repo, dest1_node = make_net_repo(local_node, subset=subset_specs[0]) - dest2_repo, _ = make_net_repo(local_node, subset=subset_specs[1]) - dest3_repo, _ = make_net_repo(local_node, subset=subset_specs[2]) + dest1_repo, dest1_node = await make_repo(local_node, subset=subset_specs[0]) + dest2_repo, _ = await make_repo(local_node, subset=subset_specs[1]) + dest3_repo, _ = await make_repo(local_node, subset=subset_specs[2]) static_route = StaticRoute([dest1_repo]) dyn_route = DynamicRoute( make_lookup(dest2_repo, dest3_repo), required_elems=["PatientID"] diff --git a/dcm/util.py b/dcm/util.py index 59305c4..006a84f 100644 --- a/dcm/util.py +++ b/dcm/util.py @@ -1,26 +1,25 @@ """Various utility functions""" from __future__ import annotations -import os, sys, json, time, logging, string -import asyncio, threading +from io import IOBase +import os, json, logging, string, asyncio, threading, functools +from pathlib import Path from functools import partial from concurrent.futures import ThreadPoolExecutor from dataclasses import is_dataclass, asdict -from contextlib import asynccontextmanager +from contextlib import AbstractContextManager, asynccontextmanager, contextmanager from typing import ( AsyncGenerator, Any, AsyncIterator, Dict, - List, + Generator, + Iterator, + Tuple, TypeVar, Optional, Union, Generic, - Iterator, Iterable, - KeysView, - ValuesView, - ItemsView, Type, Callable, ) @@ -29,12 +28,32 @@ from pydicom import Dataset from pydicom.tag import BaseTag, Tag, TagType from pydicom.datadict import tag_for_keyword -from rich.progress import Progress, Task log = logging.getLogger(__name__) +@contextmanager +def atomic_open( + path: Path, file_factory: Callable[..., IOBase] = open, **kwargs: Any +) -> Iterator[IOBase]: + """Open hidden file for writing and rename to `path` after closing without error + + Hidden file is deleted if an exception occured + """ + if "w" not in kwargs.get("mode", ""): + raise ValueError("Only valid for writing files") + tmp_path = path.parent / (".tmp_" + path.name) + try: + with file_factory(tmp_path, **kwargs) as tmp_f: + yield tmp_f + except BaseException: + os.remove(tmp_path) + raise + else: + tmp_path.rename(path) + + def dict_to_ds(data_dict: Dict[str, Any]) -> Dataset: """Convert a dict to a pydicom.Dataset""" ds = Dataset() @@ -178,6 +197,31 @@ async def aclosing( await thing.aclose() +Args_Type = Tuple[Iterable[Any], Dict[str, Any]] + +DC_Type = Callable[[Iterable[Any], Dict[str, Any]], Any] + + +def decorate_sync_async( + decorating_context: DC_Type, func: Callable[..., Any] +) -> Callable[..., Any]: + """Helper to decorate functions that are sync or async""" + if asyncio.iscoroutinefunction(func): + + async def adecorated(*args: Iterable[Any], **kwargs: Dict[str, Any]) -> Any: + with decorating_context(args, kwargs) as (args, kwargs): + return await func(*args, **kwargs) + + return functools.wraps(func)(adecorated) + else: + + def decorated(*args: Iterable[Any], **kwargs: Dict[str, Any]) -> Any: + with decorating_context(args, kwargs) as (args, kwargs): + return func(*args, **kwargs) + + return functools.wraps(func)(decorated) + + def fstr_eval( f_str: str, context: Dict[str, Any], diff --git a/poetry.lock b/poetry.lock index 7b9ca70..ce0aad8 100644 --- a/poetry.lock +++ b/poetry.lock @@ -22,6 +22,14 @@ category = "dev" optional = false python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*" +[[package]] +name = "atpublic" +version = "3.0.1" +description = "Keep all y'all's __all__'s in sync" +category = "main" +optional = false +python-versions = ">=3.7" + [[package]] name = "attrs" version = "21.4.0" @@ -167,6 +175,19 @@ python-versions = ">=3.7" docs = ["furo (>=2021.8.17b43)", "sphinx (>=4.1)", "sphinx-autodoc-typehints (>=1.12)"] testing = ["covdefaults (>=1.2.0)", "coverage (>=4)", "pytest (>=4)", "pytest-cov", "pytest-timeout (>=1.4.2)"] +[[package]] +name = "flufl.lock" +version = "7.0" +description = "NFS-safe file locking with timeouts for POSIX and Windows" +category = "main" +optional = false +python-versions = ">=3.7" + +[package.dependencies] +atpublic = ">=2.3" +psutil = ">=5.9.0" +typing_extensions = {version = "*", markers = "python_version < \"3.8\""} + [[package]] name = "greenlet" version = "1.1.2" @@ -432,7 +453,7 @@ wcwidth = "*" name = "psutil" version = "5.9.1" description = "Cross-platform lib for process and system monitoring in Python." -category = "dev" +category = "main" optional = false python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*" @@ -851,7 +872,7 @@ testing = ["pytest (>=6)", "pytest-checkdocs (>=2.4)", "pytest-flake8", "pytest- [metadata] lock-version = "1.1" python-versions = "^3.7" -content-hash = "0d107f96cc28316a8dfee51a597fa8dbfeac222d8be415aa24a028a208422425" +content-hash = "7a46a799d61c8cf2a0b4447900b6d8652e28316d64fcf149ffd27ae608cc33b2" [metadata.files] appdirs = [ @@ -866,6 +887,7 @@ atomicwrites = [ {file = "atomicwrites-1.4.0-py2.py3-none-any.whl", hash = "sha256:6d1784dea7c0c8d4a5172b6c620f40b6e4cbfdf96d783691f2e1302a7b88e197"}, {file = "atomicwrites-1.4.0.tar.gz", hash = "sha256:ae70396ad1a434f9c7046fd2dd196fc04b12f9e91ffb859164193be8b6168a7a"}, ] +atpublic = [] attrs = [ {file = "attrs-21.4.0-py2.py3-none-any.whl", hash = "sha256:2d27e3784d7a565d36ab851fe94887c5eccd6a463168875832a1be79c82828b4"}, {file = "attrs-21.4.0.tar.gz", hash = "sha256:626ba8234211db98e869df76230a137c4c40a12d72445c45d5f5b716f076e2fd"}, @@ -956,6 +978,7 @@ filelock = [ {file = "filelock-3.7.1-py3-none-any.whl", hash = "sha256:37def7b658813cda163b56fc564cdc75e86d338246458c4c28ae84cabefa2404"}, {file = "filelock-3.7.1.tar.gz", hash = "sha256:3a0fd85166ad9dbab54c9aec96737b744106dc5f15c0b09a6744a445299fcf04"}, ] +"flufl.lock" = [] greenlet = [ {file = "greenlet-1.1.2-cp27-cp27m-macosx_10_14_x86_64.whl", hash = "sha256:58df5c2a0e293bf665a51f8a100d3e9956febfbf1d9aaf8c0677cf70218910c6"}, {file = "greenlet-1.1.2-cp27-cp27m-manylinux1_x86_64.whl", hash = "sha256:aec52725173bd3a7b56fe91bc56eccb26fbdff1386ef123abb63c84c5b43b63a"}, diff --git a/pyproject.toml b/pyproject.toml index 468a421..eb302fb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,6 +35,7 @@ dateparser = "^1.0.0" importlib-metadata = "<2.0" # Newer tzlocal versions have issues with Ubuntu (at least 16.04?) see https://github.com/regebro/tzlocal/issues/122 tzlocal = "<3.0" +"flufl.lock" = "^7.0" [tool.poetry.dev-dependencies] pre-commit = "^2.12.0"