DavidL72Code/UMB_Sustainable_Chatbot
0
1import re2import unittest3 4from conversation_state import ConversationStateMachine, empty_state5 6 7def subject(name, subject_type, unit_id=None):8 return {9 "unit_id": unit_id or name.lower().replace(" ", "-"),10 "name": name,11 "subject_type": subject_type,12 "title": "Projects" if subject_type == "project" else "People",13 "source_path": "projects.txt" if subject_type == "project" else "people.txt",14 }15 16 17class ConversationStateMachineTests(unittest.TestCase):18 def setUp(self):19 self.machine = ConversationStateMachine(rewrite_callable=self.fake_rewrite)20 self.rail = subject("Cape Cod Rail Resilience Project", "project")21 self.c3i = subject("Climate Careers Curricula Initiative", "project")22 self.tim = subject("Tim Cronin", "person")23 self.hannah = subject("Nyingilanyeofori Hannah Brown", "person")24 25 @staticmethod26 def fake_rewrite(message, subject):27 name = subject["name"]28 rewritten = message29 replacements = (30 (r"\b(the\s+)?(former|latter|first one|second one)\b", name),31 (r"\b(the\s+)?(other|previous)\s+(one|project|initiative|person)\b", name),32 (r"\b(that|this)\s+(project|initiative|program|person|one)\b", name),33 (r"\b(its|his|hers|their)\b", f"{name}'s"),34 (r"\b(it|she|he|him|her|they|them)\b", name),35 )36 for pattern, replacement in replacements:37 rewritten = re.sub(pattern, replacement, rewritten, flags=re.IGNORECASE)38 return rewritten if rewritten != message else f"Regarding {name}, {message.strip()}"39 40 def focus(self, active):41 state = empty_state()42 state.update({"mode": "focused", "active_subject": active, "candidate_subjects": [active]})43 return state44 45 def test_project_pronoun_rewrites_to_project(self):46 result = self.machine.resolve(47 "What specifically caused it to be launched?", self.focus(self.rail), []48 )49 self.assertTrue(result["resolved"])50 self.assertTrue(result["used_context"])51 self.assertIn(self.rail["name"], result["rewritten_query"])52 self.assertEqual(result["active_subject"]["subject_type"], "project")53 self.assertEqual(result["intent"], "cause")54 55 def test_person_pronoun_rewrites_to_person(self):56 result = self.machine.resolve(57 "What degree is she pursuing and where?", self.focus(self.hannah), []58 )59 self.assertTrue(result["resolved"])60 self.assertIn(self.hannah["name"], result["rewritten_query"])61 self.assertEqual(result["intent"], "education")62 63 def test_person_details_do_not_replace_active_project(self):64 result = self.machine.resolve(65 "Who leads it?", self.focus(self.rail), []66 )67 self.assertEqual(result["active_subject"], self.rail)68 self.assertIn(self.rail["name"], result["rewritten_query"])69 70 def test_explicit_new_person_switches_from_project(self):71 result = self.machine.resolve(72 "What is Tim Cronin's background?", self.focus(self.rail), [self.tim]73 )74 self.assertEqual(result["active_subject"], self.tim)75 self.assertFalse(result["used_context"])76 77 def test_comparison_retains_candidates_and_clarifies_pronoun(self):78 comparison = self.machine.resolve(79 "Compare C3I and the rail project.", empty_state(), [self.c3i, self.rail]80 )81 self.assertEqual(comparison["state"]["mode"], "comparing")82 follow_up = self.machine.resolve("What does it do?", comparison["state"], [])83 self.assertTrue(follow_up["needs_clarification"])84 self.assertEqual(follow_up["clarifying_question"], "Which project are you asking about?")85 self.assertEqual(len(follow_up["clarification_options"]), 2)86 87 def test_clarification_selection_resumes_pending_question(self):88 comparison = self.machine.resolve("Compare them.", empty_state(), [self.c3i, self.rail])89 clarification = self.machine.resolve("What year was it launched?", comparison["state"], [])90 selected = self.machine.resolve("the rail project", clarification["state"], [self.rail])91 self.assertEqual(selected["active_subject"], self.rail)92 93 def test_ordinal_selection_resumes_pending_question(self):94 comparison = self.machine.resolve("Compare them.", empty_state(), [self.c3i, self.rail])95 clarification = self.machine.resolve("What does it do?", comparison["state"], [])96 selected = self.machine.resolve("the second one", clarification["state"], [])97 self.assertTrue(selected["resolved"])98 self.assertEqual(selected["active_subject"], self.rail)99 self.assertIn(self.rail["name"], selected["rewritten_query"])100 101 def test_incompatible_person_pronoun_does_not_select_project(self):102 result = self.machine.resolve("What degree is she pursuing?", self.focus(self.rail), [])103 self.assertTrue(result["needs_clarification"])104 self.assertEqual(result["clarifying_question"], "Which person are you asking about?")105 self.assertEqual(result["clarification_options"], [])106 107 def test_standalone_topic_does_not_inherit_stale_subject(self):108 result = self.machine.resolve("Which students are currently at SSL?", self.focus(self.rail), [])109 self.assertFalse(result["resolved"])110 self.assertFalse(result["needs_clarification"])111 self.assertEqual(result["rewritten_query"], "Which students are currently at SSL?")112 113 def test_scope_reference_rewrites_without_inventing_person(self):114 state = empty_state()115 state.update({116 "mode": "scoped",117 "active_scope": {"name": "SSL Board of Directors", "title": "BoardOfDirectors"},118 })119 result = self.machine.resolve("Who on it works in policy?", state, [])120 self.assertTrue(result["scope_context"])121 self.assertIn("SSL Board of Directors", result["rewritten_query"])122 self.assertIsNone(result["state"]["active_subject"])123 124 def comparison_state(self):125 return self.machine.resolve(126 "Compare C3I and rail.", empty_state(), [self.c3i, self.rail]127 )["state"]128 129 def test_named_clarification_choice_resumes_pending_intent(self):130 clarification = self.machine.resolve(131 "What year was it launched?", self.comparison_state(), []132 )133 selected = self.machine.resolve(134 "Cape Cod Rail Resilience Project.", clarification["state"], [self.rail]135 )136 self.assertEqual(selected["active_subject"], self.rail)137 self.assertIn("What year", selected["rewritten_query"])138 self.assertNotEqual(selected["rewritten_query"], "Cape Cod Rail Resilience Project.")139 140 def test_correction_resumes_pending_clarification(self):141 clarification = self.machine.resolve(142 "What year was it launched?", self.comparison_state(), []143 )144 selected = self.machine.resolve("I meant C3I.", clarification["state"], [self.c3i])145 self.assertEqual(selected["active_subject"], self.c3i)146 self.assertIn("What year", selected["rewritten_query"])147 148 def test_former_and_latter_resolve_in_comparison_order(self):149 former = self.machine.resolve("What about the former's goals?", self.comparison_state(), [])150 latter = self.machine.resolve("What about the latter's funding?", self.comparison_state(), [])151 self.assertEqual(former["active_subject"], self.c3i)152 self.assertEqual(latter["active_subject"], self.rail)153 self.assertNotIn("former", former["rewritten_query"].lower())154 self.assertNotIn("latter", latter["rewritten_query"].lower())155 156 def test_plural_comparison_followups_keep_both_subjects(self):157 for question in (158 "What do they have in common?",159 "How do their goals differ?",160 "Which one was launched first?",161 "Which started earlier?",162 "Who leads each?",163 ):164 with self.subTest(question=question):165 result = self.machine.resolve(question, self.comparison_state(), [])166 self.assertFalse(result["needs_clarification"])167 self.assertTrue(result["comparison_context"])168 self.assertIn(self.c3i["name"], result["rewritten_query"])169 self.assertIn(self.rail["name"], result["rewritten_query"])170 171 def test_can_return_to_older_type_compatible_subject(self):172 c3i_state = self.machine.resolve("Tell me about C3I.", empty_state(), [self.c3i])["state"]173 tim_state = self.machine.resolve("Who is Tim Cronin?", c3i_state, [self.tim])["state"]174 project_return = self.machine.resolve(175 "Going back to the initiative, who funds it?", tim_state, []176 )177 self.assertEqual(project_return["active_subject"], self.c3i)178 179 tim_first = self.machine.resolve("Who is Tim Cronin?", empty_state(), [self.tim])["state"]180 c3i_second = self.machine.resolve("Tell me about C3I.", tim_first, [self.c3i])["state"]181 person_return = self.machine.resolve("What is his policy background?", c3i_second, [])182 self.assertEqual(person_return["active_subject"], self.tim)183 184 def test_failed_incompatible_reference_preserves_valid_subject(self):185 tim_state = self.machine.resolve("Who is Tim Cronin?", empty_state(), [self.tim])["state"]186 failed = self.machine.resolve("What does that project do?", tim_state, [])187 self.assertTrue(failed["needs_clarification"])188 recovered = self.machine.resolve("Actually, what is his role?", failed["state"], [])189 self.assertTrue(recovered["resolved"])190 self.assertEqual(recovered["active_subject"], self.tim)191 192 def test_long_continuation_uses_active_subject(self):193 result = self.machine.resolve(194 "And what foundation provided the original funding for the program's launch?",195 self.focus(self.c3i),196 [],197 )198 self.assertTrue(result["resolved"])199 self.assertTrue(result["used_context"])200 self.assertIn(self.c3i["name"], result["rewritten_query"])201 202 def test_correction_and_parallel_ellipsis_reuse_previous_facet(self):203 c3i_state = self.machine.resolve("Tell me about C3I.", empty_state(), [self.c3i])["state"]204 funding = self.machine.resolve("What foundation funds it?", c3i_state, [])["state"]205 for correction in ("Actually, I meant the rail project.", "And the rail project?"):206 with self.subTest(correction=correction):207 switched = self.machine.resolve(correction, funding, [self.rail])208 self.assertEqual(switched["active_subject"], self.rail)209 self.assertEqual(switched["intent"], "funding")210 self.assertIn("fund", switched["rewritten_query"].lower())211 212 def test_same_question_for_new_subject_reuses_previous_facet(self):213 c3i_state = self.machine.resolve("Tell me about C3I.", empty_state(), [self.c3i])["state"]214 funding_state = self.machine.resolve("Who funds it?", c3i_state, [])["state"]215 switched = self.machine.resolve("Same question for the rail project.", funding_state, [self.rail])216 self.assertEqual(switched["active_subject"], self.rail)217 self.assertEqual(switched["intent"], "funding")218 self.assertIn("fund", switched["rewritten_query"].lower())219 220 def test_other_subject_returns_to_recent_compatible_alternative(self):221 compared = self.comparison_state()222 c3i_focused = self.machine.resolve("Tell me about the first one.", compared, [])["state"]223 other = self.machine.resolve("What about the other project?", c3i_focused, [])224 self.assertTrue(other["resolved"])225 self.assertEqual(other["active_subject"], self.rail)226 self.assertIn(self.rail["name"], other["rewritten_query"])227 self.assertNotIn("other project", other["rewritten_query"].lower())228 229 def test_correction_to_new_clarification_option_resumes_pending_query(self):230 forum = subject("Climate Adaptation Forum", "project")231 clarification = self.machine.resolve(232 "What year was it launched?", self.comparison_state(), []233 )234 corrected = self.machine.resolve(235 "No, I meant the Climate Adaptation Forum.", clarification["state"], [forum]236 )237 self.assertTrue(corrected["resolved"])238 self.assertEqual(corrected["active_subject"], forum)239 self.assertIn("What year", corrected["rewritten_query"])240 241 def test_funding_attribute_can_follow_a_non_project_subject(self):242 foundation = subject("Barr Foundation", "person", "topic:barr-foundation")243 result = self.machine.resolve(244 "How did SSL use those funds through the Summer Anti-Racism Research Funding?",245 self.focus(foundation),246 [],247 )248 self.assertTrue(result["resolved"])249 self.assertFalse(result["needs_clarification"])250 self.assertEqual(result["active_subject"], foundation)251 252 def test_named_program_starts_new_topic_instead_of_false_clarification(self):253 result = self.machine.resolve(254 "What was the research on the Massachusetts MVP program about?",255 self.focus(self.tim),256 [],257 )258 self.assertFalse(result["resolved"])259 self.assertFalse(result["needs_clarification"])260 self.assertEqual(261 result["rewritten_query"],262 "What was the research on the Massachusetts MVP program about?",263 )264 self.assertEqual(result["state"]["last_query"], result["rewritten_query"])265 266 def test_relative_that_is_not_replaced_with_active_subject(self):267 result = self.machine.resolve(268 "What is the East Boston study that VanDeVeer's team worked on?",269 self.focus(self.tim),270 [self.tim],271 )272 self.assertTrue(result["resolved"])273 self.assertEqual(274 result["rewritten_query"],275 "Regarding Tim Cronin, What is the East Boston study that VanDeVeer's team worked on?",276 )277 self.assertNotIn("Tim Cronin VanDeVeer", result["rewritten_query"])278 279 def test_former_and_latter_select_clarification_options(self):280 clarification = self.machine.resolve(281 "What year was it launched?", self.comparison_state(), []282 )283 former = self.machine.resolve("the former", clarification["state"], [])284 latter = self.machine.resolve("the latter", clarification["state"], [])285 self.assertEqual(former["active_subject"], self.c3i)286 self.assertEqual(latter["active_subject"], self.rail)287 288 def test_multi_entity_correction_selects_new_subject_not_comparison(self):289 c3i_state = self.machine.resolve("Tell me about C3I.", empty_state(), [self.c3i])["state"]290 funding_state = self.machine.resolve("Who funds it?", c3i_state, [])["state"]291 corrected = self.machine.resolve(292 "No, not C3I, I meant the rail project.",293 funding_state,294 [self.c3i, self.rail],295 )296 self.assertTrue(corrected["resolved"])297 self.assertEqual(corrected["active_subject"], self.rail)298 self.assertEqual(corrected["intent"], "funding")299 300 301if __name__ == "__main__":302 unittest.main()303 