feat: add extensible agent audio capabilities

This commit is contained in:
jxxghp
2026-05-14 19:37:13 +08:00
parent 6834d8b2c7
commit 7b27b7fd16
14 changed files with 845 additions and 351 deletions

View File

@@ -12,10 +12,10 @@ from telebot import apihelper
from app.agent.tools.impl.send_message import SendMessageInput
from app.agent.tools.impl.send_local_file import SendLocalFileInput
from app.agent import MoviePilotAgent, AgentChain
from app.agent.llm import AgentCapabilityManager
from app.chain.message import MessageChain
from app.core.config import settings
from app.agent.llm import LLMHelper
from app.helper.voice import VoiceHelper
from app.modules.discord import DiscordModule
from app.modules.qqbot import QQBotModule
from app.modules.qqbot.qqbot import QQBot
@@ -284,13 +284,15 @@ class AgentImageSupportTest(unittest.TestCase):
"feishu://file/om_audio/file_audio/voice.opus",
]
with patch.object(VoiceHelper, "is_available", return_value=True), patch.object(
with patch.object(
AgentCapabilityManager, "is_audio_input_available", return_value=True
), patch.object(
chain,
"run_module",
side_effect=[b"slack", b"discord", b"qq", b"vocechat", b"synology", b"feishu"],
) as run_module, patch.object(
VoiceHelper,
"transcribe_bytes",
AgentCapabilityManager,
"transcribe_audio",
side_effect=[
"slack text",
"discord text",

View File

@@ -0,0 +1,235 @@
import sys
import unittest
import importlib.util
from base64 import b64encode
from pathlib import Path
from tempfile import TemporaryDirectory
from types import SimpleNamespace
from unittest.mock import Mock, patch
sys.modules.setdefault("psutil", Mock())
sys.modules.setdefault("pyquery", Mock())
from app.core.config import settings
module_path = Path(__file__).resolve().parents[1] / "app" / "agent" / "llm" / "capability.py"
spec = importlib.util.spec_from_file_location("test_agent_llm_capability_module", module_path)
capability_module = importlib.util.module_from_spec(spec)
assert spec and spec.loader
sys.modules[spec.name] = capability_module
spec.loader.exec_module(capability_module)
AgentCapabilityManager = capability_module.AgentCapabilityManager
MiMoAudioProvider = capability_module.MiMoAudioProvider
OpenAIChatAudioProvider = capability_module.OpenAIChatAudioProvider
OpenAIAudioProvider = capability_module.OpenAIAudioProvider
class AgentCapabilityManagerTest(unittest.TestCase):
def test_registered_audio_providers_contains_builtin_providers(self):
self.assertIn("openai", AgentCapabilityManager.get_registered_audio_providers())
self.assertIn(
"openai_chat_audio", AgentCapabilityManager.get_registered_audio_providers()
)
self.assertIn("mimo", AgentCapabilityManager.get_registered_audio_providers())
def test_get_audio_provider_uses_separate_input_and_output_settings(self):
with patch.object(settings, "AUDIO_INPUT_PROVIDER", "openai"), patch.object(
settings, "AUDIO_OUTPUT_PROVIDER", "mimo"
):
self.assertIsInstance(
AgentCapabilityManager.get_audio_provider("input"), OpenAIAudioProvider
)
self.assertIsInstance(
AgentCapabilityManager.get_audio_provider("output"), MiMoAudioProvider
)
def test_chat_audio_provider_keeps_arbitrary_compatible_models(self):
provider = OpenAIChatAudioProvider()
with patch.object(
settings, "AUDIO_INPUT_MODEL", "vendor-omni-audio"
), patch.object(settings, "AUDIO_OUTPUT_MODEL", "vendor-tts-audio"):
self.assertEqual(provider._normalize_stt_model(), "vendor-omni-audio")
self.assertEqual(provider._normalize_tts_model(), "vendor-tts-audio")
def test_chat_audio_provider_uses_openai_audio_payload_shape(self):
provider = OpenAIChatAudioProvider()
fake_client = Mock()
fake_client.chat.completions.create.return_value = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content="你好"))]
)
with patch.object(provider, "_build_client", return_value=fake_client), patch.object(
settings, "AUDIO_INPUT_MODEL", "gpt-4o-audio-preview"
), patch.object(settings, "AUDIO_INPUT_LANGUAGE", "zh"), patch.object(
settings, "AUDIO_INPUT_API_KEY", "sk-test"
), patch.object(settings, "AUDIO_INPUT_BASE_URL", "https://example.com/v1"):
result = provider.transcribe_audio(b"audio-bytes", filename="input.wav")
self.assertEqual(result, "你好")
request = fake_client.chat.completions.create.call_args.kwargs
content = request["messages"][0]["content"]
self.assertEqual(
content[0]["input_audio"],
{"data": b64encode(b"audio-bytes").decode("utf-8"), "format": "wav"},
)
def test_chat_audio_provider_requests_audio_modality_for_tts(self):
provider = OpenAIChatAudioProvider()
fake_client = Mock()
audio_data = b64encode(b"wav-bytes").decode("utf-8")
fake_client.chat.completions.create.return_value = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(audio={"data": audio_data}))]
)
with TemporaryDirectory() as temp_dir, patch.object(
provider, "_build_client", return_value=fake_client
), patch.object(
capability_module,
"settings",
SimpleNamespace(
TEMP_PATH=Path(temp_dir),
AUDIO_OUTPUT_MODEL="gpt-4o-audio-preview",
AUDIO_OUTPUT_VOICE="alloy",
AUDIO_OUTPUT_API_KEY="sk-test",
AUDIO_OUTPUT_BASE_URL="https://example.com/v1",
),
), patch.object(provider, "_convert_wav_to_opus", return_value=None):
output_path = provider.synthesize_speech("你好")
self.assertIsNotNone(output_path)
request = fake_client.chat.completions.create.call_args.kwargs
self.assertEqual(request["messages"][0]["role"], "user")
self.assertEqual(request["modalities"], ["text", "audio"])
self.assertEqual(request["audio"], {"format": "wav", "voice": "alloy"})
def test_audio_input_and_output_switches_are_independent(self):
provider = Mock()
provider.is_available_for_audio_input.return_value = True
provider.is_available_for_audio_output.return_value = True
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT", True
), patch.object(
settings, "LLM_SUPPORT_AUDIO_OUTPUT", False
), patch.object(
AgentCapabilityManager, "get_audio_provider", return_value=provider
):
self.assertTrue(AgentCapabilityManager.is_audio_input_available())
self.assertFalse(AgentCapabilityManager.is_audio_output_available())
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT", False
), patch.object(
settings, "LLM_SUPPORT_AUDIO_OUTPUT", True
), patch.object(
AgentCapabilityManager, "get_audio_provider", return_value=provider
):
self.assertFalse(AgentCapabilityManager.is_audio_input_available())
self.assertTrue(AgentCapabilityManager.is_audio_output_available())
def test_transcribe_audio_routes_to_input_provider(self):
provider = Mock()
provider.is_available_for_audio_input.return_value = True
provider.transcribe_audio.return_value = "你好"
with patch.object(settings, "LLM_SUPPORT_AUDIO_INPUT", True), patch.object(
AgentCapabilityManager, "get_audio_provider", return_value=provider
):
result = AgentCapabilityManager.transcribe_audio(b"audio")
self.assertEqual(result, "你好")
provider.transcribe_audio.assert_called_once()
def test_synthesize_speech_routes_to_output_provider(self):
provider = Mock()
provider.is_available_for_audio_output.return_value = True
provider.synthesize_speech.return_value = Path("/tmp/reply.opus")
with patch.object(settings, "LLM_SUPPORT_AUDIO_OUTPUT", True), patch.object(
AgentCapabilityManager, "get_audio_provider", return_value=provider
):
result = AgentCapabilityManager.synthesize_speech("你好")
self.assertEqual(result, Path("/tmp/reply.opus"))
provider.synthesize_speech.assert_called_once_with(text="你好")
def test_mimo_tts_uses_chat_completions_audio_payload(self):
provider = MiMoAudioProvider()
fake_client = Mock()
audio_data = b64encode(b"wav-bytes").decode("utf-8")
fake_client.chat.completions.create.return_value = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(audio={"data": audio_data}))]
)
with TemporaryDirectory() as temp_dir, patch.object(
provider, "_build_client", return_value=fake_client
), patch.object(
capability_module,
"settings",
SimpleNamespace(
TEMP_PATH=Path(temp_dir),
AUDIO_OUTPUT_MODEL="mimo-v2.5-tts",
AUDIO_OUTPUT_VOICE="冰糖",
AUDIO_OUTPUT_API_KEY="sk-test",
AUDIO_OUTPUT_BASE_URL="https://api.xiaomimimo.com/v1",
),
), patch.object(provider, "_convert_wav_to_opus", return_value=None):
output_path = provider.synthesize_speech("你好")
output_bytes = output_path.read_bytes() if output_path else None
self.assertIsNotNone(output_path)
self.assertEqual(output_bytes, b"wav-bytes")
fake_client.chat.completions.create.assert_called_once()
request = fake_client.chat.completions.create.call_args.kwargs
self.assertEqual(request["model"], "mimo-v2.5-tts")
self.assertEqual(request["messages"][0]["role"], "assistant")
self.assertEqual(request["messages"][0]["content"], "你好")
self.assertEqual(request["audio"], {"format": "wav", "voice": "冰糖"})
def test_mimo_tts_rejects_voice_design_and_clone_models(self):
provider = MiMoAudioProvider()
with patch.object(
settings, "AUDIO_OUTPUT_MODEL", "mimo-v2.5-tts-voiceclone"
), patch.object(provider, "_build_client") as build_client:
result = provider.synthesize_speech("你好")
self.assertIsNone(result)
build_client.assert_not_called()
def test_mimo_stt_rejects_non_audio_mimo_models_by_falling_back(self):
provider = MiMoAudioProvider()
with patch.object(settings, "AUDIO_INPUT_MODEL", "mimo-v2.5-pro"):
self.assertEqual(provider._normalize_stt_model(), "mimo-v2.5")
def test_mimo_stt_uses_base64_audio_input(self):
provider = MiMoAudioProvider()
fake_client = Mock()
fake_client.chat.completions.create.return_value = SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content="你好"))]
)
with patch.object(provider, "_build_client", return_value=fake_client), patch.object(
settings, "AUDIO_INPUT_MODEL", "mimo-v2.5"
), patch.object(settings, "AUDIO_INPUT_LANGUAGE", "zh"), patch.object(
settings, "AUDIO_INPUT_API_KEY", "sk-test"
), patch.object(
settings, "AUDIO_INPUT_BASE_URL", "https://api.xiaomimimo.com/v1"
):
result = provider.transcribe_audio(b"audio-bytes", filename="input.wav")
self.assertEqual(result, "你好")
request = fake_client.chat.completions.create.call_args.kwargs
content = request["messages"][0]["content"]
self.assertEqual(request["model"], "mimo-v2.5")
self.assertTrue(
content[0]["input_audio"]["data"].startswith("data:audio/wav;base64,")
)
self.assertIn("只输出转写结果", content[1]["text"])
if __name__ == "__main__":
unittest.main()

View File

@@ -1,4 +1,5 @@
import asyncio
import importlib.machinery
import sys
import unittest
from types import SimpleNamespace
@@ -18,6 +19,10 @@ def _stub_module(name: str, **attrs):
_stub_module("qbittorrentapi", TorrentFilesList=list)
_stub_module("transmission_rpc", File=object)
_stub_module(
"psutil",
__spec__=importlib.machinery.ModuleSpec("psutil", loader=None),
)
from app.agent.tools.factory import MoviePilotToolFactory
from app.agent import ReplyMode

View File

@@ -1,69 +0,0 @@
import unittest
import sys
from unittest.mock import Mock, patch
sys.modules.setdefault("psutil", Mock())
sys.modules.setdefault("pyquery", Mock())
from app.core.config import settings
from app.helper.voice import VoiceHelper, OpenAIVoiceProvider
class VoiceHelperTest(unittest.TestCase):
def test_registered_providers_contains_openai(self):
self.assertIn("openai", VoiceHelper.get_registered_providers())
def test_get_provider_uses_single_audio_provider_setting(self):
with patch.object(settings, "AI_VOICE_PROVIDER", "openai"):
provider = VoiceHelper.get_provider("stt")
self.assertIsInstance(provider, OpenAIVoiceProvider)
def test_is_available_checks_stt_and_tts_separately(self):
provider = Mock()
provider.is_available_for_stt.return_value = True
provider.is_available_for_tts.return_value = False
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT_OUTPUT", True
), patch.object(VoiceHelper, "get_provider", return_value=provider):
self.assertTrue(VoiceHelper.is_available("stt"))
self.assertFalse(VoiceHelper.is_available("tts"))
def test_is_available_returns_false_when_audio_switch_is_disabled(self):
provider = Mock()
provider.is_available_for_stt.return_value = True
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT_OUTPUT", False
), patch.object(VoiceHelper, "get_provider", return_value=provider):
self.assertFalse(VoiceHelper.is_available("stt"))
self.assertFalse(VoiceHelper.is_available())
def test_transcribe_bytes_routes_to_stt_provider(self):
provider = Mock()
provider.transcribe_bytes.return_value = "你好"
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT_OUTPUT", True
), patch.object(VoiceHelper, "get_provider", return_value=provider):
result = VoiceHelper.transcribe_bytes(b"audio")
self.assertEqual(result, "你好")
provider.transcribe_bytes.assert_called_once()
def test_synthesize_speech_routes_to_tts_provider(self):
provider = Mock()
provider.synthesize_speech.return_value = "/tmp/reply.opus"
with patch.object(
settings, "LLM_SUPPORT_AUDIO_INPUT_OUTPUT", True
), patch.object(VoiceHelper, "get_provider", return_value=provider):
result = VoiceHelper.synthesize_speech("你好")
self.assertEqual(result, "/tmp/reply.opus")
provider.synthesize_speech.assert_called_once_with(text="你好")
if __name__ == "__main__":
unittest.main()