From 8a19c6b960da1742c21e440088455793173b1f97 Mon Sep 17 00:00:00 2001 From: Gabriel Luiz Freitas Almeida Date: Fri, 29 Mar 2024 22:15:56 -0300 Subject: [PATCH] Add unit tests for the Record class --- src/backend/base/langflow/schema/schema.py | 68 ++++++---- tests/test_record.py | 139 +++++++++++++++++++++ 2 files changed, 184 insertions(+), 23 deletions(-) create mode 100644 tests/test_record.py diff --git a/src/backend/base/langflow/schema/schema.py b/src/backend/base/langflow/schema/schema.py index 93a733038..0c1572d68 100644 --- a/src/backend/base/langflow/schema/schema.py +++ b/src/backend/base/langflow/schema/schema.py @@ -1,4 +1,5 @@ import copy +from typing import Optional from langchain_core.documents import Document from pydantic import BaseModel, model_validator @@ -12,8 +13,9 @@ class Record(BaseModel): data (dict, optional): Additional data associated with the record. """ + text_key: Optional[str] = "text" data: dict = {} - _default_value: str = "" + default_value: Optional[str] = "" @model_validator(mode="before") def validate_data(cls, values): @@ -21,10 +23,22 @@ class Record(BaseModel): values["data"] = {} # Any other keyword should be added to the data dictionary for key in values: - if key not in values["data"] and key != "data": + if key not in values["data"] and key not in {"text_key", "data", "default_value"}: values["data"][key] = values[key] return values + def get_text(self): + """ + Retrieves the text value from the data dictionary. + + If the text key is present in the data dictionary, the corresponding value is returned. + Otherwise, the default value is returned. + + Returns: + The text value from the data dictionary or the default value. + """ + return self.data.get(self.text_key, self.default_value) + @classmethod def from_document(cls, document: Document) -> "Record": """ @@ -38,19 +52,27 @@ class Record(BaseModel): """ data = document.metadata data["text"] = document.page_content - return cls(data=data) + return cls(data=data, text_key="text") def __add__(self, other: "Record") -> "Record": """ - Concatenates the text of two records and combines their data. - - Args: - other (Record): The other record to concatenate with. - - Returns: - Record: The concatenated record. + Combines the data of two records by attempting to add values for overlapping keys + for all types that support the addition operation. Falls back to the value from 'other' + record when addition is not supported. """ - combined_data = {**self.data, **other.data} + combined_data = self.data.copy() + for key, value in other.data.items(): + # If the key exists in both records and both values support the addition operation + if key in combined_data: + try: + combined_data[key] += value + except TypeError: + # Fallback: Use the value from 'other' record if addition is not supported + combined_data[key] = value + else: + # If the key is not in the first record, simply add it + combined_data[key] = value + return Record(data=combined_data) def to_lc_document(self) -> Document: @@ -60,17 +82,20 @@ class Record(BaseModel): Returns: Document: The converted Document. """ - return Document(page_content=self.text, metadata=self.data) + text = self.data.pop(self.text_key, self.default_value) + return Document(page_content=text, metadata=self.data) def __getattr__(self, key): """ Allows attribute-like access to the data dictionary. """ try: - if key == "data" or key.startswith("_"): + if key.startswith("__"): + return self.__getattribute__(key) + if key in {"data", "text_key"} or key.startswith("_"): return super().__getattr__(key) - return self.data.get(key, self._default_value) + return self.data.get(key, self.default_value) except KeyError: # Fallback to default behavior to raise AttributeError for undefined attributes raise AttributeError(f"'{type(self).__name__}' object has no attribute '{key}'") @@ -80,7 +105,7 @@ class Record(BaseModel): Allows attribute-like setting of values in the data dictionary, while still allowing direct assignment to class attributes. """ - if key == "data" or key.startswith("_"): + if key in {"data", "text_key"} or key.startswith("_"): super().__setattr__(key, value) else: self.data[key] = value @@ -89,7 +114,7 @@ class Record(BaseModel): """ Allows attribute-like deletion from the data dictionary. """ - if key == "data" or key.startswith("_"): + if key in {"data", "text_key"} or key.startswith("_"): super().__delattr__(key) else: del self.data[key] @@ -98,12 +123,8 @@ class Record(BaseModel): """ Custom deepcopy implementation to handle copying of the Record object. """ - cls = self.__class__ - result = cls.__new__(cls) - memo[id(self)] = result - for k, v in self.__dict__.items(): - setattr(result, k, copy.deepcopy(v, memo)) - return result + # Create a new Record object with a deep copy of the data dictionary + return Record(data=copy.deepcopy(self.data, memo), text_key=self.text_key, default_value=self.default_value) def __str__(self) -> str: """ @@ -114,7 +135,8 @@ class Record(BaseModel): # build the string considering all keys in the data dictionary prefix = "Record(" suffix = ")" - text = ", ".join([f"{k}={v}" for k, v in self.data.items()]) + text = f"text_key={self.text_key}, " + text += ", ".join([f"{k}={v}" for k, v in self.data.items()]) return prefix + text + suffix # check which attributes the Record has by checking the keys in the data dictionary diff --git a/tests/test_record.py b/tests/test_record.py new file mode 100644 index 000000000..45afaa5af --- /dev/null +++ b/tests/test_record.py @@ -0,0 +1,139 @@ +from langchain_core.documents import Document + +from langflow.schema import Record + + +def test_record_initialization(): + record = Record(text_key="msg", data={"msg": "Hello, World!", "extra": "value"}) + assert record.msg == "Hello, World!" + assert record.extra == "value" + + +def test_validate_data_with_extra_keys(): + record = Record(dummy_key="dummy", data={"key": "value"}) + assert record.data["dummy_key"] == "dummy" + assert "dummy_key" in record.data + assert record.key == "value" + + +def test_conversion_to_document(): + record = Record(data={"text": "Sample text", "meta": "data"}) + document = record.to_lc_document() + assert document.page_content == "Sample text" + assert document.metadata == {"meta": "data"} + + +def test_conversion_from_document(): + document = Document(page_content="Doc content", metadata={"meta": "info"}) + record = Record.from_document(document) + assert record.text == "Doc content" + assert record.meta == "info" + + +def test_add_method_for_strings(): + record1 = Record(data={"text": "Hello"}) + record2 = Record(data={"text": " World"}) + combined = record1 + record2 + assert combined.text == "Hello World" + + +def test_add_method_for_integers(): + record1 = Record(data={"number": 5}) + record2 = Record(data={"number": 10}) + combined = record1 + record2 + assert combined.number == 15 + + +def test_add_method_with_non_overlapping_keys(): + record1 = Record(data={"text": "Hello"}) + record2 = Record(data={"number": 10}) + combined = record1 + record2 + assert combined.text == "Hello" + assert combined.number == 10 + + +def test_custom_attribute_get_set_del(): + record = Record() + record.custom_attr = "custom_value" + assert record.custom_attr == "custom_value" + del record.custom_attr + assert record.custom_attr == record.default_value + + +def test_deep_copy(): + import copy + + record1 = Record(data={"text": "Hello", "number": 10}) + record2 = copy.deepcopy(record1) + assert record2.text == "Hello" + assert record2.number == 10 + record2.text = "World" + assert record1.text == "Hello" # Ensure original is unchanged + + +def test_custom_attribute_setting_and_getting(): + record = Record() + record.dynamic_attribute = "Dynamic Value" + assert record.dynamic_attribute == "Dynamic Value" + + +def test_str_and_dir_methods(): + record = Record(text_key="text", data={"text": "Test Text", "key": "value"}) + assert "Test Text" in str(record) + assert "key" in dir(record) + assert "data" in dir(record) + + +def test_dir_includes_data_keys(): + record = Record(data={"text": "Hello", "new_attr": "value"}) + dir_output = dir(record) + + # Check for standard attributes + assert "data" in dir_output + assert "text_key" in dir_output + assert "__add__" in dir_output # Checking for a method + + # Check for dynamic attributes from data + assert "text" in dir_output + assert "new_attr" in dir_output + + # Optionally, verify that dynamically added attributes are listed + record.dynamic_attr = "dynamic" + assert "dynamic_attr" in dir_output or "dynamic_attr" in dir(record) # To account for the change + + +def test_dir_reflects_attribute_deletion(): + record = Record(data={"removable": "I can be removed"}) + assert "removable" in dir(record) + + # Delete the attribute and check again + del record.removable + assert "removable" not in dir(record) + + +def test_get_text_with_text_key(): + data = {"text": "Hello, World!"} + schema = Record(data=data, text_key="text", default_value="default") + result = schema.get_text() + assert result == "Hello, World!" + + +def test_get_text_without_text_key(): + data = {"other_key": "Hello, World!"} + schema = Record(data=data, text_key="text", default_value="default") + result = schema.get_text() + assert result == "default" + + +def test_get_text_with_empty_data(): + data = {} + schema = Record(data=data, text_key="text", default_value="default") + result = schema.get_text() + assert result == "default" + + +def test_get_text_with_none_data(): + data = None + schema = Record(data=data, text_key="text", default_value="default") + result = schema.get_text() + assert result == "default"