import csv import logging from typing import Any, Collection, Generator, TextIO from .base import BaseBatchReader, BaseRowReader CsvDictReaderT = dict[str, str] logger = logging.getLogger('users_cleanup') class StringReader(BaseRowReader[str]): def _get_row(self) -> Generator[str, None, None]: assert self._fp is not None, 'File is not opened. Probably you forgot to use context manager.' yield self._fp.readline() class CsvDictReaderMixin: _fp: TextIO | None def __init__(self, *, ensure_keys: Collection[str] | None = None, **kwargs: Any) -> None: self.ensure_keys = set(ensure_keys) if ensure_keys else None super().__init__(**kwargs) def _get_row(self) -> Generator[Any, None, None]: assert self._fp is not None, 'File is not opened. Probably you forgot to use context manager.' reader = csv.DictReader(self._fp) if self.ensure_keys: if reader.fieldnames is None: raise KeyError('Ensure keys are presented in csv') diff = self.ensure_keys.difference(reader.fieldnames) if diff: raise KeyError(f'"{", ".join(diff)}" keys are not present in csv') for row in reader: yield row class CsvDictReader(CsvDictReaderMixin, BaseRowReader[CsvDictReaderT]): pass class BatchCsvDictReader(CsvDictReaderMixin, BaseBatchReader[list[CsvDictReaderT]]): pass