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
9 changes: 5 additions & 4 deletions src/openai/lib/live/_transcript_grouping.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,10 +96,11 @@ def process(self, fragments: list[TranscriptFragment]) -> list[GroupingUpdate]:
return events

def advance(self, time_ms: float, has_incoming_user: bool = False) -> list[GroupingUpdate]:
# Use the same sums as deadline() so fractional deadlines remain reachable.
if (
self._current is not None
and self._buffered is not None
and time_ms - self._current.end_ms >= self._options.min_turn_separation_ms
and time_ms >= self._current.end_ms + self._options.min_turn_separation_ms
and not self._keep_backchannel(time_ms)
):
self._buffered = self._maybe_drop_backchannel(time_ms)
Expand All @@ -109,7 +110,7 @@ def advance(self, time_ms: float, has_incoming_user: bool = False) -> list[Group
self._current is not None
and self._current.speaker == "assistant"
and not has_incoming_user
and time_ms - self._current.end_ms >= self._options.assistant_silence_ms
and time_ms >= self._current.end_ms + self._options.assistant_silence_ms
):
return self._finish_current("inactivity")
return []
Expand Down Expand Up @@ -274,7 +275,7 @@ def _keep_backchannel(self, time_ms: float) -> bool:
and self._buffered.can_drop_as_backchannel
and not self._user_continued()
and not self._recent_assistant()
and time_ms - self._buffered.end_ms < self._options.backchannel_isolation_ms
and time_ms < self._buffered.end_ms + self._options.backchannel_isolation_ms
)

def _maybe_drop_backchannel(
Expand All @@ -292,7 +293,7 @@ def _maybe_drop_backchannel(
):
return self._buffered
if next_fragment is None and (
time_ms is None or time_ms - self._buffered.end_ms < self._options.backchannel_isolation_ms
time_ms is None or time_ms < self._buffered.end_ms + self._options.backchannel_isolation_ms
):
return self._buffered
return None if self._buffered.can_drop_as_backchannel else self._buffered
Expand Down
61 changes: 61 additions & 0 deletions tests/lib/live/test_transcript_grouping.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
from __future__ import annotations

import sys
import itertools
import subprocess
from typing import Callable, Sequence, AsyncIterator
from pathlib import Path

import pytest

Expand Down Expand Up @@ -293,6 +296,64 @@ async def test_zero_disables_suppression(make_transcript: Factory) -> None:
await t.finish(("user", "Tell me more"), ("assistant", "mhm"))


@pytest.mark.parametrize("async_mode", [False, True], ids=["sync", "async"])
@pytest.mark.parametrize(
"options,fragments,expected",
[
pytest.param(
{"assistant_silence_ms": 0.1},
[("assistant", "First", 0, 200), ("assistant", "Second", 201, 300)],
[("assistant", "First", "inactivity"), ("assistant", "Second", "manual")],
id="assistant-silence",
),
pytest.param(
{"min_turn_separation_ms": 0.1},
[("user", "Question", 0, 200), ("assistant", "Answer", 200, 400), ("user", "More", 800, 1000)],
[
("user", "Question", "speaker_change"),
("assistant", "Answer", "speaker_change"),
("user", "More", "manual"),
],
id="speaker-separation",
),
pytest.param(
{"min_turn_separation_ms": 0, "backchannel_isolation_ms": 0.1},
[("user", "Question", 0, 150), ("assistant", "okay", 100, 200), ("user", " More", 800, 1000)],
[("user", "Question More", "manual")],
id="backchannel-isolation",
),
],
)
def test_fractional_source_deadlines(
async_mode: bool,
options: dict[str, float],
fragments: list[tuple[str, str, int, int]],
expected: list[tuple[str, str, str]],
) -> None:
# A synchronous non-progressing loop also blocks asyncio timeouts.
code = f"""
import asyncio
from tests.lib.live.test_transcript_grouping import Transcript, fragment

async def run():
t = Transcript({async_mode!r}, **{options!r})
for speaker, value, start, end in {fragments!r}:
await t.feed(fragment(speaker, value, start, end))
await t.finish(*[(speaker, value) for speaker, value, _ in {expected!r}])
assert [(c.segment.speaker, c.segment.text, c.reason) for c in t.recording.closed] == {expected!r}

asyncio.run(run())
"""
result = subprocess.run(
[sys.executable, "-c", code],
cwd=Path(__file__).resolve().parents[3],
capture_output=True,
text=True,
timeout=5,
)
assert result.returncode == 0, result.stdout + result.stderr


async def test_fractional_timeout_rounds_up(make_transcript: Factory) -> None:
t = make_transcript(assistant_silence_ms=200.5)
await t.feed(fragment("assistant", "Fractional.", 0))
Expand Down
Loading