|
30 | 30 | ) |
31 | 31 | from assemblyai.streaming.v3._base import _build_uri |
32 | 32 | from assemblyai.streaming.v3.models import ( |
| 33 | + HeartbeatEvent, |
33 | 34 | KeepAlive, |
34 | 35 | TerminateSession, |
| 36 | + UpdateConfiguration, |
35 | 37 | ) |
36 | 38 |
|
37 | 39 |
|
@@ -2047,3 +2049,190 @@ def test_client_connect_retries_disabled(mocker: MockFixture): |
2047 | 2049 | # Then: exactly one attempt is made and the error is reported. |
2048 | 2050 | assert connect_mock.call_count == 1 |
2049 | 2051 | 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