Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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}",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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


Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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,
}

Expand Down
Loading