3D-Speaker-MT.Axera / tests /test_lightweight.py
Nnow2024's picture
support long audio
3592794 verified
Raw
History Blame Contribute Delete
9.52 kB
# -*- 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("<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()