063c95d6ac
feature/dynamic and generalized branch * included txt files for nltk_data * move nltk_data to src * Fix last upload count * Last upload count bugfix * Fixed B processing * Remove client-specific postprocessing * split consolidate_output * Run by file * Fix output * Reconfigure smart_chunk fields * Add Full Context * Fix merge conflicts * Regex, Smart-Chunked, and Full working - not adding smart-chunked-->full when necessary * Modernized run_full_context_fields() * Switched set to list in field_context * Move fields from smart_chunked to full_context as part of 'field_context' function * Working version with placeholders * Update poetry and pyproject * Update s3 output * Remove deprecated unit test * Updated error messages * Updated smart chunk ac name to one to one * Update dependencies - end-to-end test for write s3 functional * Add basic multithreading * Send individual output to s3/local Approved-by: Alex Galarce
135 lines
5.3 KiB
Python
135 lines
5.3 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 raise ValueError)."""
|
|
self.field_set.add_field(self.field)
|
|
with self.assertRaises(ValueError):
|
|
self.field_set.add_field(self.field)
|
|
|
|
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()
|