Files
doczyai-pipelines/fieldExtraction/tests/test_fieldset.py
T
Katon Minhas c4e519894b Merged in refactor/one-time-io (pull request #648)
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
2025-08-05 20:46:19 +00:00

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