Add unit tests for the Record class
This commit is contained in:
parent
04f635d669
commit
8a19c6b960
2 changed files with 184 additions and 23 deletions
|
|
@ -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
139
tests/test_record.py
Normal 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"
|
||||
Loading…
Add table
Add a link
Reference in a new issue