c4e519894b
Refactor/one time io * Delete client work * Streamline imports * Fix list class * Update tests * Clear code funcs test * Fix unit tests * Single read code funcs * Single-read model * Single-load embeddings * Successful E2E test * update unit tests * Update importsg * Reload poetry.lock * remove old exhibit header function * Remove aarete_derived generic function * remove align and format tables * Remove strings to dict * references * Clear code * Move preprocessing_funcs * remove keywords * refactor postprocessing_funcs * Pass unit test - remove qa_qc directory * Black and isort * remove print Approved-by: Alex Galarce
142 lines
5.3 KiB
Python
142 lines
5.3 KiB
Python
import unittest
|
|
from unittest.mock import MagicMock, mock_open, patch
|
|
|
|
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()
|