54 lines
2.2 KiB
Python
54 lines
2.2 KiB
Python
import unittest
|
|
|
|
from meeting_summary_lab.llm import FakeLLM
|
|
from meeting_summary_lab.pipeline import SummarizationPipeline, build_combine_prompt, chunk_text, rough_token_count
|
|
from prompt_loader import load_prompt
|
|
|
|
|
|
class PipelineTests(unittest.TestCase):
|
|
def test_rough_token_count_matches_source_formula(self):
|
|
self.assertEqual(rough_token_count("a" * 10), 4)
|
|
|
|
def test_chunking_has_multiple_overlapping_windows(self):
|
|
chunks = chunk_text("word " * 500, chunk_size_tokens=50, overlap_tokens=10)
|
|
self.assertGreater(len(chunks), 1)
|
|
self.assertTrue(chunks[0].strip())
|
|
|
|
def test_chunking_keeps_unpunctuated_windows_intact(self):
|
|
chunks = chunk_text("中" * 200, chunk_size_tokens=50, overlap_tokens=10)
|
|
self.assertEqual(len(chunks[0]), 50)
|
|
self.assertTrue(all(len(chunk) > 1 for chunk in chunks))
|
|
|
|
def test_short_text_runs_one_final_combine_call(self):
|
|
llm = FakeLLM()
|
|
result = SummarizationPipeline(llm=llm, context_tokens=1000).summarize("A short meeting transcript.")
|
|
self.assertFalse(result.used_multilevel_strategy)
|
|
self.assertEqual(result.chunk_count, 1)
|
|
self.assertEqual(len(llm.calls), 1)
|
|
|
|
def test_long_text_maps_then_generates_final_report(self):
|
|
llm = FakeLLM()
|
|
result = SummarizationPipeline(llm=llm, context_tokens=1000).summarize("Alice decided to ship next week. " * 200)
|
|
self.assertTrue(result.used_multilevel_strategy)
|
|
self.assertEqual(len(llm.calls), result.chunk_count + 1)
|
|
self.assertTrue(result.markdown.startswith("[combined summary]"))
|
|
|
|
def test_combine_prompt_contains_template_and_summaries(self):
|
|
prompt = build_combine_prompt(["会议摘要"], "# 自定义模板")
|
|
self.assertIn("# 自定义模板", prompt)
|
|
self.assertIn("会议摘要", prompt)
|
|
self.assertIn("模板", prompt)
|
|
|
|
def test_yaml_prompt_config_has_three_prompt_blocks(self):
|
|
prompt = load_prompt("base")
|
|
self.assertEqual(set(prompt["system"]), {"chunk_prompt", "combine_prompt"})
|
|
self.assertEqual(set(prompt["prompts"]), {"chunk_user", "combine_user"})
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|
|
|
|
|
|
|
|
|