Add unit tests for the Record class

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-03-29 22:15:56 -03:00
commit 8a19c6b960
2 changed files with 184 additions and 23 deletions

View file

@ -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

139
tests/test_record.py Normal file
View file

@ -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"