From 7aff681ca7e561d30eada624ead1d874f95b58a0 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Sun, 12 Apr 2026 16:51:25 +0000 Subject: [PATCH] fix: address all loader PR review comments Agent-Logs-Url: https://github.com/pickwicksoft/pystreamapi/sessions/df7a64c6-e671-45e1-adf6-ae09cef8ec2a Co-authored-by: garlontas <70283087+garlontas@users.noreply.github.com> --- pystreamapi/loaders/__json/__json_loader.py | 14 ++++-- pystreamapi/loaders/__xml/__xml_loader.py | 55 +++++++++------------ tests/_loaders/test_xml_loader.py | 17 ++++--- tests/_loaders/test_yaml_loader.py | 7 +++ 4 files changed, 50 insertions(+), 43 deletions(-) diff --git a/pystreamapi/loaders/__json/__json_loader.py b/pystreamapi/loaders/__json/__json_loader.py index 3a743b3..f506f80 100644 --- a/pystreamapi/loaders/__json/__json_loader.py +++ b/pystreamapi/loaders/__json/__json_loader.py @@ -30,9 +30,13 @@ def generator(): # skipcq: PTC-W6004 with open(file_path, mode='r', encoding='utf-8') as jsonfile: src = jsonfile.read() - if src == '': + if not src.strip(): return - yield from jsonlib.loads(src, object_hook=__dict_to_namedtuple) + result = jsonlib.loads(src, object_hook=__dict_to_namedtuple) + if isinstance(result, list): + yield from result + else: + yield result return generator() @@ -43,7 +47,11 @@ def __lazy_load_json_string(json_string: str) -> Iterator[Any]: def generator(): if not json_string.strip(): return - yield from jsonlib.loads(json_string, object_hook=__dict_to_namedtuple) + result = jsonlib.loads(json_string, object_hook=__dict_to_namedtuple) + if isinstance(result, list): + yield from result + else: + yield result return generator() diff --git a/pystreamapi/loaders/__xml/__xml_loader.py b/pystreamapi/loaders/__xml/__xml_loader.py index f617677..50c2acc 100644 --- a/pystreamapi/loaders/__xml/__xml_loader.py +++ b/pystreamapi/loaders/__xml/__xml_loader.py @@ -10,17 +10,6 @@ from pystreamapi.loaders.__loader_utils import LoaderUtils -class __XmlLoaderUtil: - """Utility class for the XML loader.""" - - def __init__(self): - self.cast_types = True - self.retrieve_children = True - - -config = __XmlLoaderUtil() - - def xml(src: str, read_from_src=False, retrieve_children=True, cast_types=True, encoding="utf-8") -> Iterator[Any]: """ @@ -38,70 +27,70 @@ def xml(src: str, read_from_src=False, retrieve_children=True, cast_types=True, a path to an XML file. :param cast_types: Set as False to disable casting of values to int, bool or float. """ - config.cast_types = cast_types - config.retrieve_children = retrieve_children - if read_from_src: - return _lazy_parse_xml_string(src) + return _lazy_parse_xml_string(src, retrieve_children, cast_types) path = LoaderUtils.validate_path(src) - return _lazy_parse_xml_file(path, encoding) + return _lazy_parse_xml_file(path, encoding, retrieve_children, cast_types) -def _lazy_parse_xml_file(file_path: str, encoding: str) -> Iterator[Any]: +def _lazy_parse_xml_file(file_path: str, encoding: str, + retrieve_children: bool, cast_types: bool) -> Iterator[Any]: def generator(): with open(file_path, mode='r', encoding=encoding) as xmlfile: xml_string = xmlfile.read() - yield from _parse_xml_string_lazy(xml_string) + yield from _parse_xml_string_lazy(xml_string, retrieve_children, cast_types) return generator() -def _lazy_parse_xml_string(xml_string: str) -> Iterator[Any]: +def _lazy_parse_xml_string(xml_string: str, retrieve_children: bool, + cast_types: bool) -> Iterator[Any]: def generator(): - yield from _parse_xml_string_lazy(xml_string) + yield from _parse_xml_string_lazy(xml_string, retrieve_children, cast_types) return generator() -def _parse_xml_string_lazy(xml_string: str) -> Iterator[Any]: +def _parse_xml_string_lazy(xml_string: str, retrieve_children: bool, + cast_types: bool) -> Iterator[Any]: root = ElementTree.fromstring(xml_string) - parsed = __parse_xml(root) - if config.retrieve_children: + parsed = __parse_xml(root, cast_types) + if retrieve_children: yield from __flatten(parsed) else: yield parsed -def __parse_xml(element): +def __parse_xml(element, cast_types: bool): """Parse XML element and convert it into a namedtuple.""" if len(element) == 0: - return __parse_empty_element(element) + return __parse_empty_element(element, cast_types) if len(element) == 1: - return __parse_single_element(element) - return __parse_multiple_elements(element) + return __parse_single_element(element, cast_types) + return __parse_multiple_elements(element, cast_types) -def __parse_empty_element(element): +def __parse_empty_element(element, cast_types: bool): """Parse XML element without children and convert it into a namedtuple.""" - return LoaderUtils.try_cast(element.text) if config.cast_types else element.text + return LoaderUtils.try_cast(element.text) if cast_types else element.text -def __parse_single_element(element): +def __parse_single_element(element, cast_types: bool): """Parse XML element with a single child and convert it into a namedtuple.""" sub_element = element[0] - sub_item = __parse_xml(sub_element) + sub_item = __parse_xml(sub_element, cast_types) Item = namedtuple(element.tag, [sub_element.tag]) return Item(sub_item) -def __parse_multiple_elements(element): +def __parse_multiple_elements(element, cast_types: bool): """Parse XML element with multiple children and convert it into a namedtuple.""" tag_dict = {} for e in element: if e.tag not in tag_dict: tag_dict[e.tag] = [] - tag_dict[e.tag].append(__parse_xml(e)) + tag_dict[e.tag].append(__parse_xml(e, cast_types)) filtered_dict = __filter_single_items(tag_dict) Item = namedtuple(element.tag, filtered_dict.keys()) return Item(*filtered_dict.values()) diff --git a/tests/_loaders/test_xml_loader.py b/tests/_loaders/test_xml_loader.py index 04fb10c..ff6fb1d 100644 --- a/tests/_loaders/test_xml_loader.py +++ b/tests/_loaders/test_xml_loader.py @@ -32,9 +32,12 @@ class TestXmlLoader(TestCase): + def setUp(self): + self.file_content = file_content + @contextmanager - def mock_csv_file(self, content=None, exists=True, is_file=True): - """Context manager for mocking CSV file operations. + def mock_xml_file(self, content=None, exists=True, is_file=True): + """Context manager for mocking XML file operations. Args: content: The content of the mocked file @@ -48,7 +51,7 @@ def mock_csv_file(self, content=None, exists=True, is_file=True): yield def test_xml_loader_from_file_children(self): - with self.mock_csv_file(file_content): + with self.mock_xml_file(file_content): data = xml(file_path) first = next(data) @@ -66,7 +69,7 @@ def test_xml_loader_from_file_children(self): self.assertRaises(StopIteration, next, data) def test_xml_loader_from_file_no_children_false(self): - with self.mock_csv_file(file_content): + with self.mock_xml_file(file_content): data = xml(file_path, retrieve_children=False) first = next(data) @@ -80,7 +83,7 @@ def test_xml_loader_from_file_no_children_false(self): self.assertRaises(StopIteration, next, data) def test_xml_loader_no_casting(self): - with self.mock_csv_file(file_content): + with self.mock_xml_file(file_content): data = xml(file_path, cast_types=False) first = next(data) @@ -98,12 +101,12 @@ def test_xml_loader_no_casting(self): self.assertRaises(StopIteration, next, data) def test_xml_loader_is_iterable(self): - with self.mock_csv_file(file_content): + with self.mock_xml_file(file_content): data = xml(file_path) self.assertEqual(len(list(iter(data))), 3) def test_xml_loader_with_empty_file(self): - with self.mock_csv_file(''): + with self.mock_xml_file(''): data = xml(file_path) self.assertRaises(ParseError, next, data) diff --git a/tests/_loaders/test_yaml_loader.py b/tests/_loaders/test_yaml_loader.py index 6326ea7..ca9f66c 100644 --- a/tests/_loaders/test_yaml_loader.py +++ b/tests/_loaders/test_yaml_loader.py @@ -3,6 +3,8 @@ from unittest import TestCase from unittest.mock import patch, mock_open +import yaml as yaml_lib + from _loaders.file_test import OPEN, PATH_EXISTS, PATH_ISFILE from pystreamapi.loaders import yaml @@ -62,6 +64,11 @@ def test_yaml_loader_is_lazy(self): data = yaml(file_path) self.assertIsInstance(data, GeneratorType) + def test_yaml_loader_with_malformed_yaml(self): + malformed_yaml = "key: : invalid" + with self.assertRaises(yaml_lib.YAMLError): + list(yaml(malformed_yaml, read_from_src=True)) + def _check_extracted_data(self, data): first = next(data) self.assertEqual(first.attr1, 1)