diff --git a/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_llm_proposition_extractor_sync.py b/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_llm_proposition_extractor_sync.py index bbef8ae5..aa8abf94 100644 --- a/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_llm_proposition_extractor_sync.py +++ b/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_llm_proposition_extractor_sync.py @@ -67,6 +67,7 @@ def _run_non_batch_extractor(self, nodes): all_nodes = [node for node in nodes] extractor = LLMPropositionExtractor( + llm=self.llm, prompt_template=self.prompt_template, source_metadata_field=self.source_metadata_field ) diff --git a/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_topic_extractor_sync.py b/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_topic_extractor_sync.py index 55894eaa..ea1e3cac 100644 --- a/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_topic_extractor_sync.py +++ b/lexical-graph/src/graphrag_toolkit/lexical_graph/indexing/extract/batch_topic_extractor_sync.py @@ -80,6 +80,7 @@ def _run_non_batch_extractor(self, nodes): all_nodes = [node for node in nodes] extractor = TopicExtractor( + llm=self.llm, prompt_template=self.prompt_template, source_metadata_field=self.source_metadata_field, entity_classification_provider=self.entity_classification_provider, diff --git a/lexical-graph/tests/unit/indexing/extract/test_batch_llm_proposition_extractor_sync.py b/lexical-graph/tests/unit/indexing/extract/test_batch_llm_proposition_extractor_sync.py index 4798359a..7426d70e 100644 --- a/lexical-graph/tests/unit/indexing/extract/test_batch_llm_proposition_extractor_sync.py +++ b/lexical-graph/tests/unit/indexing/extract/test_batch_llm_proposition_extractor_sync.py @@ -70,3 +70,44 @@ def test_update_node_with_missing_node_id(self): result = extractor._update_node(node, node_metadata_map) assert result.metadata[PROPOSITIONS_KEY] == [] + + +class TestBatchLLMPropositionExtractorSyncRunNonBatchExtractor: + """Tests for _run_non_batch_extractor method. + + Regression tests: when a node set falls below Bedrock's minimum batch + size, BatchLLMPropositionExtractorSync falls back to the non-batch + LLMPropositionExtractor. That fallback must reuse the configured LLM + (self.llm) rather than silently defaulting to + GraphRAGConfig.extraction_llm, which is a us.* inference profile and + is invalid in any other Bedrock region. + """ + + def _make_extractor(self, llm, prompt_template="prompt", source_metadata_field=None): + """Create a BatchLLMPropositionExtractorSync instance with fields + populated via model_construct, bypassing validation/__init__ so no + real BatchConfig or AWS setup is needed for this unit test.""" + return BatchLLMPropositionExtractorSync.model_construct( + llm=llm, + prompt_template=prompt_template, + source_metadata_field=source_metadata_field, + ) + + @patch("graphrag_toolkit.lexical_graph.indexing.extract.batch_llm_proposition_extractor_sync.LLMPropositionExtractor") + def test_run_non_batch_extractor_passes_configured_llm(self, mock_extractor_cls): + """Verify the non-batch fallback is constructed with the extractor's + own configured llm, not left to default to GraphRAGConfig.extraction_llm.""" + configured_llm = Mock(name="configured-llm") + extractor = self._make_extractor(llm=configured_llm) + + mock_instance = mock_extractor_cls.return_value + mock_instance.extract.return_value = [{PROPOSITIONS_KEY: []}] + + nodes = [TextNode(text="test", id_="node-1")] + extractor._run_non_batch_extractor(nodes) + + mock_extractor_cls.assert_called_once_with( + llm=configured_llm, + prompt_template="prompt", + source_metadata_field=None, + ) diff --git a/lexical-graph/tests/unit/indexing/extract/test_batch_topic_extractor_sync.py b/lexical-graph/tests/unit/indexing/extract/test_batch_topic_extractor_sync.py index 7f5b21ea..c2feb24a 100644 --- a/lexical-graph/tests/unit/indexing/extract/test_batch_topic_extractor_sync.py +++ b/lexical-graph/tests/unit/indexing/extract/test_batch_topic_extractor_sync.py @@ -68,3 +68,54 @@ def test_update_node_with_missing_node_id(self): result = extractor._update_node(node, node_metadata_map) assert result.metadata[TOPICS_KEY] == {'topics': []} + + +class TestBatchTopicExtractorSyncRunNonBatchExtractor: + """Tests for _run_non_batch_extractor method. + + Regression tests: when a node set falls below Bedrock's minimum batch + size, BatchTopicExtractorSync falls back to the non-batch + TopicExtractor. That fallback must reuse the configured LLM (self.llm) + rather than silently defaulting to GraphRAGConfig.extraction_llm, which + is a us.* inference profile and is invalid in any other Bedrock region. + """ + + def _make_extractor(self, llm, prompt_template="prompt", source_metadata_field=None, + entity_classification_provider=None, topic_provider=None): + """Create a BatchTopicExtractorSync instance with fields populated + via model_construct, bypassing validation/__init__ so no real + BatchConfig or AWS setup is needed for this unit test.""" + return BatchTopicExtractorSync.model_construct( + llm=llm, + prompt_template=prompt_template, + source_metadata_field=source_metadata_field, + entity_classification_provider=entity_classification_provider or Mock(), + topic_provider=topic_provider or Mock(), + ) + + @patch("graphrag_toolkit.lexical_graph.indexing.extract.batch_topic_extractor_sync.TopicExtractor") + def test_run_non_batch_extractor_passes_configured_llm(self, mock_extractor_cls): + """Verify the non-batch fallback is constructed with the extractor's + own configured llm, not left to default to GraphRAGConfig.extraction_llm.""" + configured_llm = Mock(name="configured-llm") + entity_classification_provider = Mock(name="entity-classification-provider") + topic_provider = Mock(name="topic-provider") + extractor = self._make_extractor( + llm=configured_llm, + entity_classification_provider=entity_classification_provider, + topic_provider=topic_provider, + ) + + mock_instance = mock_extractor_cls.return_value + mock_instance.extract.return_value = [{TOPICS_KEY: {'topics': []}}] + + nodes = [TextNode(text="test", id_="node-1")] + extractor._run_non_batch_extractor(nodes) + + mock_extractor_cls.assert_called_once_with( + llm=configured_llm, + prompt_template="prompt", + source_metadata_field=None, + entity_classification_provider=entity_classification_provider, + topic_provider=topic_provider, + )