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()