Files
doczyai-pipelines/fieldExtraction/tests/fieldset_test.py
T
Katon Minhas 839f13d934 Merged in hotfix/dynamic-error (pull request #470)
Hotfix/dynamic error

* Bugfix

* update tests


Approved-by: Alex Galarce
2025-04-08 20:13:22 +00:00

135 lines
5.2 KiB
Python

import unittest
from unittest.mock import patch, MagicMock, mock_open
from src.prompts.investment_prompts import Field, FieldSet
class TestField(unittest.TestCase):
def setUp(self):
self.field_dict = {
'field_name': 'test_field',
'relationship': 'one_to_one',
'field_type': 'exhibit_level',
'prompt': 'Please provide {valid_values}.',
'valid_values': 'some_value'
}
self.field = Field(self.field_dict)
def test_init(self):
"""Test initialization of Field object."""
self.assertEqual(self.field.field_name, 'test_field')
self.assertEqual(self.field.relationship, 'one_to_one')
self.assertEqual(self.field.field_type, 'exhibit_level')
self.assertEqual(self.field.prompt, 'Please provide {valid_values}.')
self.assertEqual(self.field.valid_values, 'some_value')
def test_from_values(self):
"""Test the class method 'from_values'."""
new_field = Field.from_values(
field_name='new_field',
relationship='one_to_n',
field_type='dynamic',
prompt='Choose only from the following: {valid_values}.',
valid_values='new_value'
)
self.assertEqual(new_field.field_name, 'new_field')
self.assertEqual(new_field.relationship, 'one_to_n')
self.assertEqual(new_field.field_type, 'dynamic')
self.assertEqual(new_field.prompt, 'Choose only from the following: {valid_values}.')
self.assertEqual(new_field.valid_values, 'new_value')
def test_to_dict(self):
"""Test conversion to dictionary."""
field_dict = self.field.to_dict()
self.assertEqual(field_dict['field_name'], 'test_field')
self.assertEqual(field_dict['relationship'], 'one_to_one')
self.assertEqual(field_dict['field_type'], 'exhibit_level')
self.assertEqual(field_dict['prompt'], 'Please provide {valid_values}.')
self.assertEqual(field_dict['valid_values'], 'some_value')
def test_print(self):
"""Test printing of Field object."""
# Normally we'd check print output, but here we'll assert we reach print calls.
with patch('builtins.print') as mocked_print:
self.field.print()
mocked_print.assert_any_call('Field Name: test_field')
def test_get_prompt_dict(self):
"""Test get_prompt_dict method."""
prompt_dict = self.field.get_prompt_dict()
self.assertEqual(prompt_dict['test_field'], 'Please provide some_value.')
def test_get_prompt(self):
"""Test prompt resolution."""
resolved_prompt = self.field.get_prompt()
self.assertEqual(resolved_prompt, 'Please provide some_value.')
def test_update_valid_values(self):
"""Test updating valid values."""
self.field.update_valid_values('new_value')
self.assertEqual(self.field.valid_values, 'new_value')
class TestFieldSet(unittest.TestCase):
def setUp(self):
self.field_dict = {
'field_name': 'test_field',
'relationship': 'one_to_one',
'field_type': 'exhibit_level',
'prompt': 'Please provide {valid_values}.',
'valid_values': 'some_value'
}
self.field = Field(self.field_dict)
self.field_set = FieldSet()
def test_add_field(self):
"""Test adding a field."""
self.field_set.add_field(self.field)
self.assertEqual(len(self.field_set.fields), 1)
def test_add_duplicate_field(self):
"""Test adding a duplicate field (should do nothing)"""
self.field_set.add_field(self.field)
self.assertEqual(len(self.field_set.fields), 1)
def test_get_field(self):
"""Test getting a specific field."""
self.field_set.add_field(self.field)
fetched_field = self.field_set.get_field('test_field')
self.assertEqual(fetched_field.field_name, 'test_field')
def test_get_field_not_found(self):
"""Test getting a non-existent field (should raise ValueError)."""
with self.assertRaises(ValueError):
self.field_set.get_field('non_existent_field')
def test_list_fields(self):
"""Test listing all fields."""
self.field_set.add_field(self.field)
fields = self.field_set.list_fields()
self.assertIn('test_field', fields)
def test_remove_field(self):
"""Test removing a field."""
self.field_set.add_field(self.field)
self.field_set.remove_field('test_field')
self.assertEqual(len(self.field_set.fields), 0)
def test_get_prompt_dict(self):
"""Test getting prompt dictionary."""
self.field_set.add_field(self.field)
prompt_dict = self.field_set.get_prompt_dict()
self.assertIn('test_field', prompt_dict)
def test_load_from_file(self):
"""Test loading fields from a file (mocked)."""
with patch('builtins.open', mock_open(read_data='[{"field_name": "test_field", "relationship": "one_to_one", "field_type": "exhibit_level", "prompt": "Please provide {valid_values}.", "valid_values": "some_value"}]')):
self.field_set.load_from_file('mock_file.json')
self.assertEqual(len(self.field_set.fields), 1)
if __name__ == '__main__':
unittest.main()