diff --git a/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_input_without_text_input_end.py b/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_input_without_text_input_end.py index 022d164dbb..f1a7d83126 100644 --- a/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_input_without_text_input_end.py +++ b/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_input_without_text_input_end.py @@ -203,7 +203,10 @@ def _validate_metadata( event_name: str, ) -> bool: """Validate metadata matches expected.""" - if received_metadata != expected_metadata: + if any( + key not in received_metadata or received_metadata[key] != value + for key, value in expected_metadata.items() + ): self._stop_test_with_error( ten_env, f"Metadata mismatch in {event_name}. Expected: {expected_metadata}, Received: {received_metadata}", diff --git a/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_interrupt.py b/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_interrupt.py index 053ac1d128..03b768fef2 100644 --- a/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_interrupt.py +++ b/ai_agents/agents/integration_tests/tts_guarder/tests/test_append_interrupt.py @@ -579,7 +579,11 @@ async def on_data(self, ten_env: AsyncTenEnvTester, data: Data) -> None: if metadata_str: try: received_metadata = json.loads(metadata_str) - if received_metadata != self.sent_flush_metadata: + if any( + key not in received_metadata + or received_metadata[key] != value + for key, value in self.sent_flush_metadata.items() + ): self._stop_test_with_error(ten_env, f"Metadata mismatch in flush_end. Expected: {self.sent_flush_metadata}, Received: {received_metadata}") return except json.JSONDecodeError: diff --git a/ai_agents/agents/integration_tests/tts_guarder/tests/test_connection_status.py b/ai_agents/agents/integration_tests/tts_guarder/tests/test_connection_status.py index 034d0c97e6..3f70833fe9 100644 --- a/ai_agents/agents/integration_tests/tts_guarder/tests/test_connection_status.py +++ b/ai_agents/agents/integration_tests/tts_guarder/tests/test_connection_status.py @@ -158,7 +158,13 @@ def _validate_connection_status_events( def _validate_event_payload( self, ten_env: AsyncTenEnvTester, event: dict[str, Any] ) -> None: - expected_fields = ["id", "module", "vendor", "current", "last"] + expected_fields = [ + "id", + "module", + "vendor_info", + "current", + "last", + ] missing = [field for field in expected_fields if field not in event] if missing: self._stop_test_with_error( @@ -185,8 +191,11 @@ def _validate_event_payload( ) return - if not event.get("vendor"): - self._stop_test_with_error(ten_env, "Missing vendor in event") + vendor_info = event.get("vendor_info") + if not isinstance(vendor_info, dict) or not vendor_info.get("vendor"): + self._stop_test_with_error( + ten_env, "Missing vendor_info.vendor in event" + ) def test_connection_status(extension_name: str, config_dir: str) -> None: diff --git a/ai_agents/agents/integration_tests/tts_guarder/tests/test_flush.py b/ai_agents/agents/integration_tests/tts_guarder/tests/test_flush.py index 20a8686c31..efd1c3a3b3 100644 --- a/ai_agents/agents/integration_tests/tts_guarder/tests/test_flush.py +++ b/ai_agents/agents/integration_tests/tts_guarder/tests/test_flush.py @@ -43,7 +43,7 @@ def __init__( print("đŸŽ¯ Test Objectives:") print(" - Verify flush is generated") print(" - Validate flush_id and metadata consistency in flush_end response") - print(" - Ensure no audio/text data after flush_end for 5 seconds") + print(" - Ensure no audio frames after flush_end for 5 seconds") print("=" * 80) self.session_id: str = session_id @@ -54,7 +54,6 @@ def __init__( self.flush_end_received = False self.flush_id = "test_flush_request_id_1" self.post_flush_end_audio_count = 0 - self.post_flush_end_data_count = 0 self.flush_end_timestamp = None self.sent_flush_metadata = None @@ -184,7 +183,11 @@ async def on_data(self, ten_env: AsyncTenEnvTester, data: Data) -> None: if metadata_str: try: received_metadata = json.loads(metadata_str) - if received_metadata != self.sent_flush_metadata: + if any( + key not in received_metadata + or received_metadata[key] != value + for key, value in self.sent_flush_metadata.items() + ): self._stop_test_with_error(ten_env, f"Metadata mismatch in flush_end. Expected: {self.sent_flush_metadata}, Received: {received_metadata}") return except json.JSONDecodeError: @@ -200,17 +203,13 @@ async def on_data(self, ten_env: AsyncTenEnvTester, data: Data) -> None: self.flush_end_received = True self.flush_end_timestamp = time.time() - # Start a 5-second monitoring task to check if there is any audio/text data after flush_end + # Start a 5-second monitoring task for audio frames after flush_end asyncio.create_task(self._monitor_post_flush_end_data(ten_env)) else: - # Check if any other data is received after flush_end if self.flush_end_received: - # Non-ttfb metrics (e.g. connect_delay, usage) after flush are expected - if name == "metrics": - ten_env.log_info(f"â„šī¸ Received expected metrics data after flush_end") - return - self.post_flush_end_data_count += 1 - ten_env.log_info(f"âš ī¸ Received data '{name}' after flush_end (count: {self.post_flush_end_data_count})") + ten_env.log_info( + f"Received expected non-audio data '{name}' after flush_end" + ) return @@ -249,19 +248,23 @@ async def _send_flush(self, ten_env: AsyncTenEnvTester) -> None: await ten_env.send_data(flush_data) async def _monitor_post_flush_end_data(self, ten_env: AsyncTenEnvTester) -> None: - """Monitor if there is any audio/text data after flush_end for 5 seconds""" - ten_env.log_info("Start monitoring data after flush_end...") + """Monitor for audio frames after flush_end for 5 seconds.""" + ten_env.log_info("Start monitoring audio frames after flush_end...") # Wait for 5 seconds await asyncio.sleep(5.0) - # Check if there is any additional data - if self.post_flush_end_audio_count > 0 or self.post_flush_end_data_count > 0: - error_msg = f"Received additional data after flush_end for 5 seconds: audio frames {self.post_flush_end_audio_count} , other data {self.post_flush_end_data_count} " + if self.post_flush_end_audio_count > 0: + error_msg = ( + "Received audio frames after flush_end for 5 seconds: " + f"{self.post_flush_end_audio_count}" + ) ten_env.log_info(f"❌ {error_msg}") self._stop_test_with_error(ten_env, error_msg) else: - ten_env.log_info("✅ No additional data received after flush_end for 5 seconds, test passed") + ten_env.log_info( + "✅ No audio frames received after flush_end for 5 seconds, test passed" + ) ten_env.stop_test() def test_flush(extension_name: str, config_dir: str) -> None: diff --git a/ai_agents/agents/ten_packages/extension/openai_tts2_python/extension.py b/ai_agents/agents/ten_packages/extension/openai_tts2_python/extension.py index 77a185797f..eec9e3723a 100644 --- a/ai_agents/agents/ten_packages/extension/openai_tts2_python/extension.py +++ b/ai_agents/agents/ten_packages/extension/openai_tts2_python/extension.py @@ -54,13 +54,17 @@ def vendor_metadata(self) -> dict: "Authorization", "", ) or self.config.headers.get("authorization", "") - return { + metadata = { "url": self.config.url or "", "model": self.config.params.get("model", ""), "api_key": self.config.params.get("api_key", ""), "authorization": authorization, "voice": self.config.params.get("voice", ""), } + language = self.config.params.get("language", "") + if language: + metadata["language"] = language + return metadata def synthesize_audio_sample_rate(self) -> int: return 24000 diff --git a/ai_agents/agents/ten_packages/extension/rime_tts/extension.py b/ai_agents/agents/ten_packages/extension/rime_tts/extension.py index 44f42a0bea..116342fe7b 100644 --- a/ai_agents/agents/ten_packages/extension/rime_tts/extension.py +++ b/ai_agents/agents/ten_packages/extension/rime_tts/extension.py @@ -142,6 +142,7 @@ def vendor_metadata(self) -> dict: "key": self.config.api_key, "url": self.config.base_url, "model": self.config.params.get("modelId", ""), + "language": self.config.params.get("lang", ""), "api_key": self.config.api_key, }