# -*- coding: utf-8 -*- 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("ignored最终摘要"), ] ) ) ) ) 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("\n无。这是第一轮请求。\n", first_prompt) self.assertIn("\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", second_prompt) if __name__ == "__main__": unittest.main()