import csv import json import re from datetime import date from io import StringIO from typing import Any, Callable, Collection, Iterator from airflow.utils.log.logging_mixin import LoggingMixin __all__ = ["CSV"] class CSV(LoggingMixin): date_pattern: re.Pattern = re.compile(r"\d{4}-\d{2}-\d{2}") def __init__(self, delimiter: str, quote_char: str, null_value: str, normalizers: dict[str, Callable[[Any], Any]] | None = None, **kwargs): super().__init__(**kwargs) self.delimiter = delimiter self.quote_char = quote_char self.null_value = null_value if normalizers is None: normalizers = {} self.normalizers = normalizers @classmethod def _is_date(cls, value: Any) -> bool: return isinstance(value, str) and bool(re.match(cls.date_pattern, value)) @classmethod def _is_json(cls, value: Any) -> bool: return ( isinstance(value, str) and len(value) >= 2 and ((value[0] == "[" and value[-1] == "]") or (value[0] == "{" and value[-1] == "}")) ) @staticmethod def _parse_date(value: str) -> Any: try: return date.fromisoformat(value) except ValueError: pass return value @classmethod def _parse_json(cls, value: str) -> Any: try: value = json.loads(value) if isinstance(value, list) and value and cls._is_date(value[0]): value = [cls._parse_date(v) for v in value] except ValueError: pass return value def _default_normalizer(self, value: Any) -> Any: if not isinstance(value, str): return value # if value looks like null/na/nan elif value == self.null_value: value = None # if value looks like numeric elif value.lstrip("-").isdigit(): try: value = int(value) except ValueError: pass # if value looks like json elif self._is_json(value): value = self._parse_json(value) # if value looks like date elif self._is_date(value): value = self._parse_date(value) return value def _normalize_record(self, record: dict[str, Any]) -> dict[str, Any]: normalized_record = {} for key, value in record.items(): key = key.lower() normalizer = self.normalizers.get(key, self._default_normalizer) normalized_record[key] = normalizer(value) return normalized_record def load(self, data: Collection[str]) -> Iterator[dict[str, Any]]: for record in csv.DictReader(data, delimiter=self.delimiter, quotechar=self.quote_char): yield self._normalize_record(record) def dump(self, data: Collection[dict[str, Any]], header: Collection[str], write_header: bool = True) -> bytes: with StringIO() as buf: writer = csv.DictWriter(buf, fieldnames=header, delimiter=self.delimiter, quotechar=self.quote_char) if write_header: writer.writeheader() writer.writerows(data) result = buf.getvalue().encode() return result