| |
| import importlib |
| import asyncio |
| import os |
| from pathlib import Path |
| import sys |
| import tempfile |
| import types |
| import unittest |
| from types import SimpleNamespace |
| from unittest.mock import patch |
|
|
| from ax_meeting.text_cleaner import clean_asr_text |
| from ax_meeting.diar_utils import pick_speaker |
| from ax_meeting import summarizer as summarizer_module |
|
|
|
|
| fake_model_bundle_module = types.ModuleType("ax_meeting.model_bundle") |
|
|
|
|
| class FakeModelBundle: |
| def ensure_loaded(self): |
| pass |
|
|
|
|
| fake_model_bundle_module.ModelBundle = FakeModelBundle |
| sys.modules.setdefault("ax_meeting.model_bundle", fake_model_bundle_module) |
|
|
| fake_engines_module = types.ModuleType("ax_meeting.engines") |
| fake_engines_module.DiarAsrEngine = object |
| sys.modules.setdefault("ax_meeting.engines", fake_engines_module) |
|
|
| fake_pipeline_module = types.ModuleType("ax_meeting.pipeline") |
| fake_pipeline_module.StreamingMeetingSession = object |
| sys.modules.setdefault("ax_meeting.pipeline", fake_pipeline_module) |
|
|
| from ax_meeting import config as config_module |
| from ax_meeting import server |
|
|
|
|
| class TestTextCleaner(unittest.TestCase): |
| def test_angle_tokens(self): |
| s = "<|zh|<|NEUTRAL|<|Speech|<|withitn|嗯,你好你好你好。" |
| self.assertEqual(clean_asr_text(s), "嗯,你好你好你好。") |
|
|
| def test_pipe_tokens(self): |
| s = "zh|NEUTRAL|Speech|withitn|他相当于把一个平台就把一个拼行能力拆分掉了。" |
| self.assertEqual(clean_asr_text(s), "他相当于把一个平台就把一个拼行能力拆分掉了。") |
|
|
| def test_concat_tokens(self): |
| s = "zhEMO_UNKNOWNSpeechwithitn房止追后了对这种方式掉。" |
| self.assertEqual(clean_asr_text(s), "房止追后了对这种方式掉。") |
|
|
| def test_empty_after_clean(self): |
| s = "zh|NEUTRAL|Speech|withitn|" |
| self.assertEqual(clean_asr_text(s), "") |
|
|
|
|
| class TestPickSpeaker(unittest.TestCase): |
| def test_pick_speaker_overlap(self): |
| diar = [ |
| [0.0, 2.0, 0], |
| [2.0, 5.0, 1], |
| ] |
| spk = pick_speaker(1.5, 3.5, diar) |
| self.assertEqual(spk, 1) |
|
|
| def test_pick_speaker_no_overlap(self): |
| diar = [[0.0, 1.0, 0]] |
| spk = pick_speaker(2.0, 3.0, diar) |
| self.assertEqual(spk, 0) |
|
|
|
|
| class TestServerStartup(unittest.TestCase): |
| def test_preload_models_calls_ensure_loaded(self): |
| with patch.object(server.models, "ensure_loaded") as mocked: |
| server.preload_models() |
| mocked.assert_called_once_with() |
|
|
|
|
| class TestServerPersistence(unittest.TestCase): |
| def test_persist_result_texts_writes_timestamped_files(self): |
| with tempfile.TemporaryDirectory() as tmpdir, patch.object(server, "RESULT_DIR", Path(tmpdir)): |
| transcript_path, summary_path = server.persist_result_texts("transcript body", "summary body") |
| transcript_body = transcript_path.read_text(encoding="utf-8") |
| summary_body = summary_path.read_text(encoding="utf-8") |
|
|
| self.assertTrue(transcript_path.name.endswith("_transcript.txt")) |
| self.assertTrue(summary_path.name.endswith("_summary.txt")) |
| self.assertEqual(transcript_body, "transcript body") |
| self.assertEqual(summary_body, "summary body") |
|
|
| def test_clear_session_cache_removes_store_and_releases_session(self): |
| class FakeSession: |
| def __init__(self): |
| self.released = False |
|
|
| def release(self): |
| self.released = True |
|
|
| with tempfile.TemporaryDirectory() as tmpdir, patch.object(server, "release_process_memory") as mocked_release: |
| record_path = Path(tmpdir) / "meeting.mp3" |
| record_path.write_bytes(b"dummy") |
| server.TRANSCRIPT_STORE["sid"] = "hello" |
| server.RECORDING_STORE["sid"] = record_path |
| fake_session = FakeSession() |
|
|
| server.clear_session_cache("sid", session=fake_session) |
|
|
| self.assertNotIn("sid", server.TRANSCRIPT_STORE) |
| self.assertNotIn("sid", server.RECORDING_STORE) |
| self.assertFalse(record_path.exists()) |
| self.assertTrue(fake_session.released) |
| mocked_release.assert_called_once_with() |
|
|
| def test_summary_api_persists_files_and_clears_cache(self): |
| server.TRANSCRIPT_STORE["sid"] = "meeting transcript" |
| fake_transcript_path = Path("/tmp/transcript.txt") |
| fake_summary_path = Path("/tmp/summary.txt") |
|
|
| async def run_test(): |
| with patch.object(server, "summarize_transcript_text", return_value="meeting summary"), patch.object( |
| server, "persist_result_texts", return_value=(fake_transcript_path, fake_summary_path) |
| ) as mocked_persist, patch.object(server, "clear_session_cache") as mocked_clear: |
| result = await server.summary_api( |
| session_id="sid", |
| openai_base_url="", |
| openai_api_key="", |
| openai_model="", |
| ) |
|
|
| self.assertEqual(result["text"], "meeting summary") |
| self.assertEqual(result["transcript_path"], str(fake_transcript_path)) |
| self.assertEqual(result["summary_path"], str(fake_summary_path)) |
| mocked_persist.assert_called_once_with("meeting transcript", "meeting summary") |
| mocked_clear.assert_called_once_with("sid") |
|
|
| asyncio.run(run_test()) |
| server.TRANSCRIPT_STORE.pop("sid", None) |
|
|
|
|
| class TestSummarizerHelpers(unittest.TestCase): |
| def test_split_text_fixed_size_empty(self): |
| self.assertEqual(summarizer_module._split_text_fixed_size(" ", 4), []) |
|
|
| def test_split_text_fixed_size_chunks_by_chars(self): |
| self.assertEqual( |
| summarizer_module._split_text_fixed_size("abcdefghij", 4), |
| ["abcd", "efgh", "ij"], |
| ) |
|
|
| def test_split_text_fixed_size_invalid_chunk_size(self): |
| with self.assertRaises(ValueError): |
| summarizer_module._split_text_fixed_size("abc", 0) |
|
|
| def test_summary_target_range_from_scalar(self): |
| with patch.object(summarizer_module, "SUMMARY_TARGET_CHARS", 120): |
| self.assertEqual(summarizer_module._summary_target_range(), (120, 120)) |
|
|
| def test_summary_target_range_from_tuple(self): |
| with patch.object(summarizer_module, "SUMMARY_TARGET_CHARS", (100, 300)): |
| self.assertEqual(summarizer_module._summary_target_range(), (100, 300)) |
|
|
| def test_summary_target_range_normalizes_reverse_order(self): |
| with patch.object(summarizer_module, "SUMMARY_TARGET_CHARS", [300, 100]): |
| self.assertEqual(summarizer_module._summary_target_range(), (100, 300)) |
|
|
| def test_summary_chunk_chars_uses_environment_override(self): |
| old_value = os.environ.get("SUMMARY_CHUNK_CHARS") |
| try: |
| os.environ["SUMMARY_CHUNK_CHARS"] = "3456" |
| reloaded = importlib.reload(config_module) |
| self.assertEqual(reloaded.SUMMARY_CHUNK_CHARS, 3456) |
| finally: |
| if old_value is None: |
| os.environ.pop("SUMMARY_CHUNK_CHARS", None) |
| else: |
| os.environ["SUMMARY_CHUNK_CHARS"] = old_value |
| importlib.reload(config_module) |
|
|
|
|
| class TestIncrementalSummarizer(unittest.TestCase): |
| def _make_response(self, content: str): |
| return SimpleNamespace( |
| choices=[ |
| SimpleNamespace( |
| message=SimpleNamespace(content=content), |
| ) |
| ] |
| ) |
|
|
| def test_summarize_incrementally_uses_previous_summary_and_returns_last_round(self): |
| fake_client = SimpleNamespace( |
| chat=SimpleNamespace( |
| completions=SimpleNamespace( |
| create=unittest.mock.Mock( |
| side_effect=[ |
| self._make_response("第一轮摘要"), |
| self._make_response("<think>ignored</think>最终摘要"), |
| ] |
| ) |
| ) |
| ) |
| ) |
|
|
| with patch.object(summarizer_module, "OpenAI", return_value=fake_client), patch.object( |
| summarizer_module, "SUMMARY_CHUNK_CHARS", 5 |
| ), patch.object(summarizer_module, "SUMMARY_TARGET_CHARS", (100, 300)): |
| summarizer = summarizer_module.IncrementalSummarizer(api_key="test-key") |
| summary = summarizer.summarize_incrementally("abcdefghij") |
|
|
| self.assertEqual(summary, "最终摘要") |
| create = fake_client.chat.completions.create |
| self.assertEqual(create.call_count, 2) |
|
|
| first_prompt = create.call_args_list[0].kwargs["messages"][1]["content"] |
| second_prompt = create.call_args_list[1].kwargs["messages"][1]["content"] |
|
|
| self.assertIn("<previous_summary>\n无。这是第一轮请求。\n</previous_summary>", first_prompt) |
| self.assertIn("<current_transcript>\n以下内容是本轮新发送的 transcript 原文,请与上一轮摘要衔接后理解:\nabcde", first_prompt) |
| self.assertIn("请务必保留此前各轮与本轮中出现的关键决策、结论、待办事项、负责人、时间点、风险与分歧", first_prompt) |
| self.assertIn("当前是第2/2轮总结请求", second_prompt) |
| self.assertIn("以下内容是上一轮请求返回的摘要", second_prompt) |
| self.assertIn("第一轮摘要", second_prompt) |
| self.assertIn("\nfghij\n</current_transcript>", second_prompt) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|