CoolFace
Apppublic

DavidL72Code/UMB_Sustainable_Chatbot

sourceHugging Faceupdated 5d agoView on Hugging Face
0likes
test_conversation_state_machine.py303 linesDownload Raw Back to tests
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