Skip to content
Draft
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
2 changes: 1 addition & 1 deletion sdk/voice/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ authors = [{ name = "Speechmatics", email = "support@speechmatics.com" }]
license = "MIT"
requires-python = ">=3.9"
dependencies = [
"speechmatics-rt>=0.5.3",
"speechmatics-rt>=1.0.0",
"pydantic>=2.10.6,<3",
"numpy>=1.26.4,<3"
]
Expand Down
2 changes: 1 addition & 1 deletion tests/voice/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
async def get_client(
api_key: Optional[str] = None,
url: Optional[str] = None,
app: Optional[str] = None,
app: str = "sdk-test",
config: Optional[VoiceAgentConfig] = None,
connect: bool = True,
) -> VoiceAgentClient:
Expand Down
59 changes: 52 additions & 7 deletions tests/voice/test_08_multiple_speakers.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@

# Constants
API_KEY = os.getenv("SPEECHMATICS_API_KEY")
URL = os.getenv("SPEECHMATICS_RT_URL", "wss://eu2.rt.speechmatics.com/v2")
SHOW_LOG = os.getenv("SPEECHMATICS_SHOW_LOG", "0").lower() in ["1", "true"]


Expand Down Expand Up @@ -116,11 +117,17 @@ async def test_multiple_speakers(sample: SpeakerTest):

# Client
client = await get_client(
url=URL,
api_key=API_KEY,
connect=False,
config=config,
)

# Debug
if SHOW_LOG:
print(config.to_json(exclude_none=True, exclude_defaults=True, exclude_unset=True, indent=2))
print(json.dumps(client._transcription_config.to_dict(), indent=2))

# Create an event to track when the callback is called
messages: list[str] = []
bytes_sent: int = 0
Expand Down Expand Up @@ -148,19 +155,35 @@ def log_final_segment(message):
segments: list[SpeakerSegment] = message["segments"]
final_segments.extend(segments)

# Log end of turn
def log_end_of_turn(message):
final_segments.extend([{"speaker_id": "--", "text": "_TURN_"}])

# Add listeners
client.once(AgentServerMessageType.RECOGNITION_STARTED, log_message)
client.once(AgentServerMessageType.INFO, log_message)
client.on(AgentServerMessageType.WARNING, log_message)
client.on(AgentServerMessageType.ERROR, log_message)
client.on(AgentServerMessageType.DIAGNOSTICS, log_message)

# Transcript
client.on(AgentServerMessageType.ADD_PARTIAL_TRANSCRIPT, log_message)
client.on(AgentServerMessageType.ADD_TRANSCRIPT, log_message)
client.on(AgentServerMessageType.ADD_PARTIAL_SEGMENT, log_message)
client.on(AgentServerMessageType.ADD_SEGMENT, log_message)

# Turn events
client.on(AgentServerMessageType.VAD_STATUS, log_message)
client.on(AgentServerMessageType.SPEAKER_STARTED, log_message)
client.on(AgentServerMessageType.SPEAKER_ENDED, log_message)
client.on(AgentServerMessageType.START_OF_TURN, log_message)
client.on(AgentServerMessageType.END_OF_TURN, log_message)
client.on(AgentServerMessageType.END_OF_TURN_PREDICTION, log_message)
client.on(AgentServerMessageType.END_OF_UTTERANCE, log_message)

# Log ADD_SEGMENT
# Log ADD_SEGMENT + END_OF_TURN
client.on(AgentServerMessageType.ADD_SEGMENT, log_final_segment)
client.on(AgentServerMessageType.END_OF_TURN, log_end_of_turn)

# HEADER
if SHOW_LOG:
Expand All @@ -187,22 +210,44 @@ def log_final_segment(message):
progress_callback=log_bytes_sent,
)

# Close session
await client.disconnect()

# FOOTER
if SHOW_LOG:
print("---")
print()
print()

# Print all final_segments
if SHOW_LOG:
print("Final segments:")
for idx, segment in enumerate(final_segments):
print(f"{idx}: [{segment.get('speaker_id')}] {segment.get('text')}")
print()

# Accumulate errors
errors: list[str] = []

# Check number of final segments
if len(final_segments) < len(sample.segment_regex):
errors.append(f"Expected at least {len(sample.segment_regex)} segments, got {len(final_segments)}")

# Check final segments against regex
if SHOW_LOG:
print("Checking final segments against regex:")
for idx, _test in enumerate(sample.segment_regex):
text = final_segments[idx].get("text") if idx < len(final_segments) else None
match = text and re.search(_test, text, flags=re.IGNORECASE | re.MULTILINE)
if SHOW_LOG:
print(f"`{_test}` -> `{final_segments[idx].get('text')}`")
assert re.search(_test, final_segments[idx].get("text"), flags=re.IGNORECASE | re.MULTILINE)
print(f'{idx}: {"✅" if match else "❌"} - `{_test}` -> `{text}`')
if not match:
errors.append(f"Segment {idx}: expected /{_test}/ but got '{text}'")

# Check only speakers present
speakers = [segment.get("speaker_id") for segment in final_segments]
assert set(speakers) == set(sample.speakers_present)
if set(speakers) != set(sample.speakers_present):
errors.append(f"Speakers: expected {set(sample.speakers_present)} but got {set(speakers)}")

# Close session
await client.disconnect()
assert not client._is_connected
# Report all errors
assert not errors, "\n".join(errors)
Loading