Skip to content
Merged
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 CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -336,7 +336,7 @@ The former `Streaming*` names (`StreamingClient`, `AsyncStreamingClient`, `Strea
- `Warning` → `WarningEvent(warning_code: int, warning: str)`
- `Error` → `RealTimeError` (an `Exception` subclass) with `.code: int | None`; `str(error)` is the message. Server-side errors come through `on_error` rather than being raised, and the payload is a `RealTimeError`, **not** the wire `ErrorEvent` class.
- `LLMGatewayResponse` → `LLMGatewayResponseEvent(turn_order: int, transcript: str, data: Any)`
- `SpeakerRevision` → `SpeakerRevisionEvent(revisions: list[SpeakerRevisionItem])` — diarization-only. Sent once per offline-recluster resolve. Each `SpeakerRevisionItem(turn_order: int, speaker_label: str | None, words: list[Word])` is an earlier Turn whose labels changed (unchanged turns are omitted). For each item, match by `turn_order` against the original Turn and replace its per-word `speaker` (and the turn-level `speaker_label`) with the revision's values. Text and word timestamps are unchanged.
- `SpeakerRevision` → `SpeakerRevisionEvent(revisions: list[SpeakerRevisionItem])` — diarization-only. Sent once per offline-recluster resolve. Each `SpeakerRevisionItem(turn_order: int, speaker_label: str | None, words: list[Word])` is an earlier Turn whose labels changed (unchanged turns are omitted). For each item, match by `turn_order` against the original Turn and replace its per-word `speaker` (and the turn-level `speaker_label`) with the revision's values. Text and word timestamps are unchanged. By default only the end-of-stream revision is sent; set `speaker_labels_revision_interval_ms` on `RealTimeParameters` (ms of audio time; values below 120_000 are raised to 120_000 server-side, larger values honored, 300_000 recommended) to also get them mid-stream — the first can only arrive after ~120 s of speech, and one is sent only when a label actually changed.

**Sync streaming:**
```python
Expand Down
2 changes: 1 addition & 1 deletion assemblyai/__version__.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
__version__ = "1.6.0"
__version__ = "1.6.1"
5 changes: 3 additions & 2 deletions assemblyai/streaming/v3/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,12 @@

try:
# pydantic v2 import
from pydantic import BaseModel, model_validator
from pydantic import BaseModel, Field, model_validator

pydantic_v2 = True
except ImportError:
# pydantic v1 import (fallback for Python < 3.14)
from pydantic import BaseModel, root_validator
from pydantic import BaseModel, Field, root_validator

pydantic_v2 = False

Expand Down Expand Up @@ -332,6 +332,7 @@ class RealTimeParameters(RealTimeSessionParameters):
webhook_auth_header_value: Optional[str] = None
llm_gateway: Optional[LLMGatewayConfig] = None
speaker_labels: Optional[bool] = None
speaker_labels_revision_interval_ms: Optional[int] = Field(None, ge=0)
max_speakers: Optional[int] = None
voice_focus: Optional[NoiseSuppressionModel] = None
voice_focus_threshold: Optional[float] = None
Expand Down
96 changes: 96 additions & 0 deletions tests/unit/test_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -867,6 +867,102 @@ def mocked_websocket_connect(
assert "max_speakers=3" in actual_url


def test_client_connect_with_speaker_labels_revision_interval(mocker: MockFixture):
# Given: speaker_labels plus a mid-stream revision cadence
actual_url = None

def mocked_websocket_connect(
url: str, additional_headers: dict, open_timeout: float
):
nonlocal actual_url
actual_url = url

mocker.patch(
"assemblyai.streaming.v3.client.websocket_connect",
new=mocked_websocket_connect,
)

_disable_rw_threads(mocker)

client = StreamingClient(
StreamingClientOptions(api_key="test", api_host="api.example.com")
)

# When: connecting
client.connect(
StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
speaker_labels_revision_interval_ms=120_000,
)
)

# Then: the GA (un-prefixed) query param carries the value verbatim
assert "speaker_labels=True" in actual_url
assert "speaker_labels_revision_interval_ms=120000" in actual_url
assert "_speaker_labels_revision_interval_ms" not in actual_url


def test_build_uri_keeps_explicit_zero_revision_interval():
# Given: the interval explicitly set to 0 (end-of-stream revision only)
params = StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
speaker_labels_revision_interval_ms=0,
)

# When: the connection URI is built
uri = _build_uri("api.example.com", params)

# Then: 0 is sent, not dropped like an unset value
assert "speaker_labels_revision_interval_ms=0" in uri


def test_build_uri_omits_unset_revision_interval():
# Given: speaker_labels without a revision interval
params = StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
)

# When: the connection URI is built
uri = _build_uri("api.example.com", params)

# Then: the param is absent so the server default applies
assert "speaker_labels_revision_interval_ms" not in uri


def test_negative_revision_interval_is_rejected():
# Given/When: a negative cadence, which the server also rejects at connect
# Then: the SDK refuses it before a connection is attempted
with pytest.raises(ValueError, match="speaker_labels_revision_interval_ms"):
StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
speaker_labels_revision_interval_ms=-1,
)


def test_large_revision_interval_is_not_clamped_client_side():
# Given: a cadence above the server's 300_000 default, which the server
# honors rather than clamps
params = StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
speaker_labels_revision_interval_ms=600_000,
)

# When/Then: the SDK forwards it unchanged
assert "speaker_labels_revision_interval_ms=600000" in _build_uri(
"api.example.com", params
)


def test_client_connect_with_continuous_partials(mocker: MockFixture):
# Given: client + continuous_partials=True (U3-Pro steady-partials mode)
actual_url = None
Expand Down
28 changes: 28 additions & 0 deletions tests/unit/test_streaming_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,34 @@ async def test_client_connect_with_token(mocker: MockFixture):
await client.disconnect()


async def test_client_connect_with_speaker_labels_revision_interval(
mocker: MockFixture,
):
# Given: speaker_labels plus a mid-stream revision cadence
fake_ws = _FakeAsyncWebSocket()
fake_connect = _patch_connect(mocker, fake_ws)

client = AsyncStreamingClient(
StreamingClientOptions(api_key="test", api_host="api.example.com")
)

# When: connecting
await client.connect(
StreamingParameters(
sample_rate=16000,
speech_model=SpeechModel.universal_streaming_english,
speaker_labels=True,
speaker_labels_revision_interval_ms=120_000,
)
)

# Then: the query string carries the GA param name and value
assert "speaker_labels=True" in fake_connect.uri
assert "speaker_labels_revision_interval_ms=120000" in fake_connect.uri

await client.disconnect()


async def test_stream_bytes_writes_to_socket(mocker: MockFixture):
fake_ws = _FakeAsyncWebSocket()
_patch_connect(mocker, fake_ws)
Expand Down
Loading