Skip to content

Commit 166ed28

Browse files
Merge branch 'master' into ccampbell/add-sync-unit-tests
2 parents 92a20a9 + a98b105 commit 166ed28

6 files changed

Lines changed: 300 additions & 7 deletions

File tree

‎assemblyai/__version__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
__version__ = "0.64.31"
1+
__version__ = "0.64.32"

‎assemblyai/streaming/v3/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
BeginEvent,
1717
Encoding,
1818
EventMessage,
19+
HeartbeatEvent,
1920
LLMGatewayResponseEvent,
2021
NoiseSuppressionModel,
2122
SessionConfiguration,
@@ -49,6 +50,7 @@
4950
"EnergyVad",
5051
"Encoding",
5152
"EventMessage",
53+
"HeartbeatEvent",
5254
"LLMGatewayResponseEvent",
5355
"NoiseSuppressionModel",
5456
"SpeakerRevisionEvent",

‎assemblyai/streaming/v3/_base.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232
BeginEvent,
3333
ErrorEvent,
3434
EventMessage,
35+
HeartbeatEvent,
3536
LLMGatewayResponseEvent,
3637
SpeakerRevisionEvent,
3738
SpeechStartedEvent,
@@ -241,6 +242,8 @@ def _parse_message(cls, data: Dict[str, Any]) -> Optional[EventMessage]:
241242
return _parse_model(LLMGatewayResponseEvent, data)
242243
elif event_type == StreamingEvents.SpeakerRevision:
243244
return _parse_model(SpeakerRevisionEvent, data)
245+
elif event_type == StreamingEvents.Heartbeat:
246+
return _parse_model(HeartbeatEvent, data)
244247
elif event_type == StreamingEvents.Error:
245248
return _parse_model(ErrorEvent, data)
246249
elif event_type == StreamingEvents.Warning:

‎assemblyai/streaming/v3/models.py‎

Lines changed: 24 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,16 @@ class WarningEvent(BaseModel):
8585
warning: str
8686

8787

88+
class HeartbeatEvent(BaseModel):
89+
type: Literal["Heartbeat"] = "Heartbeat"
90+
total_audio_received_ms: int
91+
total_duration_ms: int
92+
# Unclamped processing speed ratio; may exceed 1.0.
93+
realtime_factor: float
94+
# Highest per-frame speech probability in the interval, range 0-1.
95+
max_speech_probability: float = 0.0
96+
97+
8898
class LLMGatewayResponseEvent(BaseModel):
8999
type: Literal["LLMGatewayResponse"] = "LLMGatewayResponse"
90100
turn_order: int
@@ -124,6 +134,7 @@ class SpeakerRevisionEvent(BaseModel):
124134
SpeechStartedEvent,
125135
ErrorEvent,
126136
WarningEvent,
137+
HeartbeatEvent,
127138
LLMGatewayResponseEvent,
128139
SpeakerRevisionEvent,
129140
]
@@ -157,6 +168,7 @@ class StreamingSessionParameters(BaseModel):
157168
interruption_delay: Optional[int] = None
158169
turn_left_pad_ms: Optional[int] = None
159170
language_codes: Optional[List[str]] = None
171+
session_heartbeat: Optional[bool] = None
160172

161173

162174
class Encoding(str, Enum):
@@ -169,6 +181,9 @@ class Encoding(str, Enum):
169181
# MediaRecorder output). `sample_rate` may be omitted — the Opus stream is
170182
# self-describing and the server ignores it.
171183
ogg_opus = "ogg_opus"
184+
# AAC in an ADTS byte stream. `sample_rate` may be omitted — the ADTS
185+
# headers are self-describing and the server ignores it.
186+
aac = "aac"
172187

173188
def __str__(self):
174189
return self.value
@@ -307,30 +322,32 @@ class StreamingParameters(StreamingSessionParameters):
307322
if pydantic_v2:
308323

309324
@model_validator(mode="after")
310-
def _require_sample_rate_for_non_opus(self):
325+
def _require_sample_rate(self):
311326
if self.sample_rate is None and self.encoding not in (
312327
Encoding.opus,
313328
Encoding.ogg_opus,
329+
Encoding.aac,
314330
):
315331
raise ValueError(
316332
"sample_rate is required; it may only be omitted when "
317-
"encoding is 'opus' or 'ogg_opus' (the Opus stream is "
318-
"self-describing)."
333+
"encoding is 'opus', 'ogg_opus', or 'aac' (these streams "
334+
"are self-describing)."
319335
)
320336
return self
321337

322338
else:
323339

324340
@root_validator(skip_on_failure=True)
325-
def _require_sample_rate_for_non_opus(cls, values):
341+
def _require_sample_rate(cls, values):
326342
if values.get("sample_rate") is None and values.get("encoding") not in (
327343
Encoding.opus,
328344
Encoding.ogg_opus,
345+
Encoding.aac,
329346
):
330347
raise ValueError(
331348
"sample_rate is required; it may only be omitted when "
332-
"encoding is 'opus' or 'ogg_opus' (the Opus stream is "
333-
"self-describing)."
349+
"encoding is 'opus', 'ogg_opus', or 'aac' (these streams "
350+
"are self-describing)."
334351
)
335352
return values
336353

@@ -407,5 +424,6 @@ class StreamingEvents(Enum):
407424
SpeechStarted = "SpeechStarted"
408425
Error = "Error"
409426
Warning = "Warning"
427+
Heartbeat = "Heartbeat"
410428
LLMGatewayResponse = "LLMGatewayResponse"
411429
SpeakerRevision = "SpeakerRevision"

‎tests/unit/test_streaming.py‎

Lines changed: 189 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,8 +30,10 @@
3030
)
3131
from assemblyai.streaming.v3._base import _build_uri
3232
from assemblyai.streaming.v3.models import (
33+
HeartbeatEvent,
3334
KeepAlive,
3435
TerminateSession,
36+
UpdateConfiguration,
3537
)
3638

3739

@@ -2047,3 +2049,190 @@ def test_client_connect_retries_disabled(mocker: MockFixture):
20472049
# Then: exactly one attempt is made and the error is reported.
20482050
assert connect_mock.call_count == 1
20492051
assert len(errors) == 1
2052+
2053+
2054+
def test_encoding_aac_enum():
2055+
# Given/Then: the AAC encoding is a first-class Encoding member whose wire
2056+
# value round-trips through the string constructor and __str__.
2057+
assert Encoding("aac") is Encoding.aac
2058+
assert str(Encoding.aac) == "aac"
2059+
2060+
2061+
def test_client_connect_aac_without_sample_rate(mocker: MockFixture):
2062+
# Given: client + AAC encoding and no sample_rate (the ADTS stream is
2063+
# self-describing, so the parameter may be omitted). Constructing the
2064+
# params must not raise, mirroring the Opus no-sample_rate case.
2065+
actual_url = None
2066+
2067+
def mocked_websocket_connect(
2068+
url: str, additional_headers: dict, open_timeout: float
2069+
):
2070+
nonlocal actual_url
2071+
actual_url = url
2072+
2073+
mocker.patch(
2074+
"assemblyai.streaming.v3.client.websocket_connect",
2075+
new=mocked_websocket_connect,
2076+
)
2077+
_disable_rw_threads(mocker)
2078+
client = StreamingClient(
2079+
StreamingClientOptions(api_key="test", api_host="api.example.com")
2080+
)
2081+
params = StreamingParameters(
2082+
speech_model=SpeechModel.universal_3_5_pro,
2083+
encoding=Encoding.aac,
2084+
)
2085+
2086+
# When: connect
2087+
client.connect(params)
2088+
2089+
# Then: the encoding is forwarded and sample_rate is absent from the URL
2090+
assert "encoding=aac" in actual_url
2091+
assert "sample_rate" not in actual_url
2092+
2093+
2094+
def test_sample_rate_still_required_for_pcm_s16le():
2095+
# Given/When/Then: adding AAC to the self-describing allow-list must not
2096+
# relax the requirement for PCM encodings — pcm_s16le with no sample_rate
2097+
# still fails validation.
2098+
with pytest.raises(ValueError, match="sample_rate is required"):
2099+
StreamingParameters(
2100+
speech_model=SpeechModel.universal_3_5_pro,
2101+
encoding=Encoding.pcm_s16le,
2102+
)
2103+
2104+
2105+
def test_client_connect_with_session_heartbeat(mocker: MockFixture):
2106+
# Given: client + session_heartbeat=True
2107+
actual_url = None
2108+
2109+
def mocked_websocket_connect(
2110+
url: str, additional_headers: dict, open_timeout: float
2111+
):
2112+
nonlocal actual_url
2113+
actual_url = url
2114+
2115+
mocker.patch(
2116+
"assemblyai.streaming.v3.client.websocket_connect",
2117+
new=mocked_websocket_connect,
2118+
)
2119+
_disable_rw_threads(mocker)
2120+
client = StreamingClient(
2121+
StreamingClientOptions(api_key="test", api_host="api.example.com")
2122+
)
2123+
params = StreamingParameters(
2124+
sample_rate=16000,
2125+
speech_model=SpeechModel.universal_streaming_english,
2126+
session_heartbeat=True,
2127+
)
2128+
2129+
# When: connect
2130+
client.connect(params)
2131+
2132+
# Then: the session_heartbeat wire param is forwarded
2133+
assert "session_heartbeat=True" in actual_url
2134+
2135+
2136+
def test_session_heartbeat_defaults_to_none():
2137+
# Given: params/update-config with no session_heartbeat set
2138+
params = StreamingParameters(sample_rate=16000)
2139+
update = UpdateConfiguration()
2140+
2141+
# Then: the field defaults to None on both StreamingSessionParameters
2142+
# subclasses and is omitted from the connection querystring when unset.
2143+
assert params.session_heartbeat is None
2144+
assert update.session_heartbeat is None
2145+
uri = _build_uri("wss://example.com/v3/ws", params)
2146+
assert "session_heartbeat" not in uri
2147+
2148+
2149+
def test_heartbeat_event_parses_from_wire_message():
2150+
# Given: a raw Heartbeat wire message
2151+
data = {
2152+
"type": "Heartbeat",
2153+
"total_audio_received_ms": 45000,
2154+
"total_duration_ms": 45205,
2155+
"realtime_factor": 0.9964,
2156+
"max_speech_probability": 0.999954,
2157+
}
2158+
2159+
# When: routed through the shared inbound-message dispatch
2160+
event = StreamingClient._parse_message(data)
2161+
2162+
# Then: a fully-populated HeartbeatEvent is returned
2163+
assert isinstance(event, HeartbeatEvent)
2164+
assert event.type == "Heartbeat"
2165+
assert event.total_audio_received_ms == 45000
2166+
assert event.total_duration_ms == 45205
2167+
assert event.realtime_factor == 0.9964
2168+
assert event.max_speech_probability == 0.999954
2169+
2170+
2171+
def test_heartbeat_max_speech_probability_defaults_to_zero():
2172+
# Given: a Heartbeat payload without max_speech_probability
2173+
data = {
2174+
"type": "Heartbeat",
2175+
"total_audio_received_ms": 1000,
2176+
"total_duration_ms": 1000,
2177+
"realtime_factor": 1.0,
2178+
}
2179+
2180+
# When: parsed
2181+
event = HeartbeatEvent.parse_obj(data)
2182+
2183+
# Then: max_speech_probability defaults to 0.0
2184+
assert event.max_speech_probability == 0.0
2185+
2186+
2187+
def test_heartbeat_realtime_factor_is_unclamped():
2188+
# Given: a Heartbeat payload with realtime_factor above 1.0
2189+
data = {
2190+
"type": "Heartbeat",
2191+
"total_audio_received_ms": 1000,
2192+
"total_duration_ms": 1500,
2193+
"realtime_factor": 1.5,
2194+
}
2195+
2196+
# When: parsed
2197+
event = HeartbeatEvent.parse_obj(data)
2198+
2199+
# Then: the value passes through unclamped
2200+
assert event.realtime_factor == 1.5
2201+
2202+
2203+
def test_heartbeat_event_dispatched_to_handler(mocker: MockFixture):
2204+
# Given: a Heartbeat frame on the wire and a handler registered
2205+
heartbeat_json = json.dumps(
2206+
{
2207+
"type": "Heartbeat",
2208+
"total_audio_received_ms": 45000,
2209+
"total_duration_ms": 45205,
2210+
"realtime_factor": 0.9964,
2211+
"max_speech_probability": 0.999954,
2212+
}
2213+
)
2214+
fake_ws = _FakeWebSocket(recv_script=[heartbeat_json])
2215+
mocker.patch(
2216+
"assemblyai.streaming.v3.client.websocket_connect",
2217+
return_value=fake_ws,
2218+
)
2219+
received = []
2220+
client = StreamingClient(
2221+
StreamingClientOptions(api_key="test", api_host="api.example.com")
2222+
)
2223+
client.on(StreamingEvents.Heartbeat, lambda _c, event: received.append(event))
2224+
2225+
# When: the client reads the frame
2226+
client.connect(_default_params())
2227+
deadline = time.monotonic() + 2.0
2228+
while time.monotonic() < deadline and not received:
2229+
time.sleep(0.02)
2230+
client.disconnect(terminate=False)
2231+
2232+
# Then: the handler is invoked with a parsed HeartbeatEvent
2233+
assert len(received) == 1
2234+
assert isinstance(received[0], HeartbeatEvent)
2235+
assert received[0].total_audio_received_ms == 45000
2236+
assert received[0].total_duration_ms == 45205
2237+
assert received[0].realtime_factor == 0.9964
2238+
assert received[0].max_speech_probability == 0.999954

0 commit comments

Comments
 (0)