mirror of
https://github.com/dograh-hq/dograh.git
synced 2026-07-25 12:01:04 +02:00
chore: refactor user turn strategies
This commit is contained in:
parent
01d4d07638
commit
69bb7be38d
10 changed files with 308 additions and 80 deletions
|
|
@ -102,6 +102,8 @@ DEFAULT_USER_TURN_STOP_TIMEOUT = 5.0
|
|||
EXTERNAL_TURN_USER_STOP_TIMEOUT = 30.0
|
||||
DEFAULT_TURN_START_STRATEGY = "default"
|
||||
DEFAULT_TURN_START_MIN_WORDS = 3
|
||||
DEFAULT_PROVISIONAL_VAD_PAUSE_SECS = 1.0
|
||||
DEFAULT_SMART_TURN_STOP_SECS = 2.0
|
||||
|
||||
|
||||
def _resolve_user_turn_stop_timeout(
|
||||
|
|
@ -121,6 +123,17 @@ def _resolve_turn_start_min_words(run_configs: dict) -> int:
|
|||
)
|
||||
|
||||
|
||||
def _resolve_provisional_vad_pause_secs(run_configs: dict) -> float:
|
||||
return max(
|
||||
0.1,
|
||||
float(
|
||||
run_configs.get(
|
||||
"provisional_vad_pause_secs", DEFAULT_PROVISIONAL_VAD_PAUSE_SECS
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _create_non_realtime_user_turn_start_strategies(
|
||||
run_configs: dict, *, uses_external_turns: bool
|
||||
):
|
||||
|
|
@ -138,7 +151,11 @@ def _create_non_realtime_user_turn_start_strategies(
|
|||
]
|
||||
|
||||
if turn_start_strategy == "provisional_vad":
|
||||
return [ProvisionalVADUserTurnStartStrategy()]
|
||||
return [
|
||||
ProvisionalVADUserTurnStartStrategy(
|
||||
pause_secs=_resolve_provisional_vad_pause_secs(run_configs)
|
||||
)
|
||||
]
|
||||
|
||||
if uses_external_turns:
|
||||
# The STT emits its own turn boundaries and owns interruptions. Local
|
||||
|
|
@ -150,6 +167,29 @@ def _create_non_realtime_user_turn_start_strategies(
|
|||
return [VADUserTurnStartStrategy()]
|
||||
|
||||
|
||||
def _create_non_realtime_user_turn_stop_strategies(
|
||||
run_configs: dict, *, uses_external_turns: bool
|
||||
):
|
||||
"""Return user turn stop strategies for non-realtime pipelines."""
|
||||
|
||||
if uses_external_turns:
|
||||
return [ExternalUserTurnStopStrategy()]
|
||||
|
||||
if run_configs.get("turn_stop_strategy") == "turn_analyzer":
|
||||
smart_turn_params = SmartTurnParams(
|
||||
stop_secs=run_configs.get(
|
||||
"smart_turn_stop_secs", DEFAULT_SMART_TURN_STOP_SECS
|
||||
)
|
||||
)
|
||||
return [
|
||||
TurnAnalyzerUserTurnStopStrategy(
|
||||
turn_analyzer=LocalSmartTurnAnalyzerV3(params=smart_turn_params)
|
||||
)
|
||||
]
|
||||
|
||||
return [SpeechTimeoutUserTurnStopStrategy()]
|
||||
|
||||
|
||||
def _create_realtime_user_turn_config(provider: str):
|
||||
"""Return user turn strategies and optional local VAD for realtime providers."""
|
||||
|
||||
|
|
@ -500,8 +540,6 @@ async def _run_pipeline_impl(
|
|||
# Extract configurations from the version's workflow_configurations
|
||||
max_call_duration_seconds = 300 # Default 5 minutes
|
||||
max_user_idle_timeout = 10.0 # Default 10 seconds
|
||||
smart_turn_stop_secs = 2.0 # Default 2 seconds for incomplete turn timeout
|
||||
turn_stop_strategy = "transcription" # Default to transcription-based detection
|
||||
keyterms = None # Dictionary words for STT boosting
|
||||
|
||||
if run_configs:
|
||||
|
|
@ -511,12 +549,6 @@ async def _run_pipeline_impl(
|
|||
if "max_user_idle_timeout" in run_configs:
|
||||
max_user_idle_timeout = run_configs["max_user_idle_timeout"]
|
||||
|
||||
if "smart_turn_stop_secs" in run_configs:
|
||||
smart_turn_stop_secs = run_configs["smart_turn_stop_secs"]
|
||||
|
||||
if "turn_stop_strategy" in run_configs:
|
||||
turn_stop_strategy = run_configs["turn_stop_strategy"]
|
||||
|
||||
if "dictionary" in run_configs:
|
||||
dictionary = run_configs["dictionary"]
|
||||
if dictionary and isinstance(dictionary, str):
|
||||
|
|
@ -780,43 +812,20 @@ async def _run_pipeline_impl(
|
|||
turn_start_strategy = run_configs.get(
|
||||
"turn_start_strategy", DEFAULT_TURN_START_STRATEGY
|
||||
)
|
||||
turn_start_strategy_names = [
|
||||
strategy.__class__.__name__ for strategy in user_turn_start_strategies
|
||||
]
|
||||
turn_start_min_words = (
|
||||
_resolve_turn_start_min_words(run_configs)
|
||||
if turn_start_strategy == "min_words"
|
||||
else None
|
||||
)
|
||||
logger.info(
|
||||
f"[run {workflow_run_id}] Non-realtime interrupt strategy "
|
||||
f"requested={turn_start_strategy} "
|
||||
f"effective={turn_start_strategy_names} "
|
||||
f"min_words={turn_start_min_words} "
|
||||
f"uses_external_turns={uses_external_turns}"
|
||||
)
|
||||
if uses_external_turns:
|
||||
user_turn_strategies = UserTurnStrategies(
|
||||
start=user_turn_start_strategies,
|
||||
stop=[ExternalUserTurnStopStrategy()],
|
||||
)
|
||||
elif turn_stop_strategy == "turn_analyzer":
|
||||
# Smart Turn Analyzer: best for longer responses with natural pauses
|
||||
smart_turn_params = SmartTurnParams(stop_secs=smart_turn_stop_secs)
|
||||
user_turn_strategies = UserTurnStrategies(
|
||||
start=user_turn_start_strategies,
|
||||
stop=[
|
||||
TurnAnalyzerUserTurnStopStrategy(
|
||||
turn_analyzer=LocalSmartTurnAnalyzerV3(params=smart_turn_params)
|
||||
)
|
||||
],
|
||||
)
|
||||
else:
|
||||
# Transcription-based (default): best for short 1-2 word responses
|
||||
user_turn_strategies = UserTurnStrategies(
|
||||
start=user_turn_start_strategies,
|
||||
stop=[SpeechTimeoutUserTurnStopStrategy()],
|
||||
)
|
||||
|
||||
user_turn_stop_strategies = _create_non_realtime_user_turn_stop_strategies(
|
||||
run_configs,
|
||||
uses_external_turns=uses_external_turns,
|
||||
)
|
||||
user_turn_strategies = UserTurnStrategies(
|
||||
start=user_turn_start_strategies,
|
||||
stop=user_turn_stop_strategies,
|
||||
)
|
||||
|
||||
user_turn_stop_timeout = _resolve_user_turn_stop_timeout(
|
||||
run_configs,
|
||||
|
|
|
|||
|
|
@ -10,14 +10,18 @@ from pipecat.turns.user_start.vad_user_turn_start_strategy import (
|
|||
from pipecat.turns.user_stop import (
|
||||
ExternalUserTurnStopStrategy,
|
||||
SpeechTimeoutUserTurnStopStrategy,
|
||||
TurnAnalyzerUserTurnStopStrategy,
|
||||
)
|
||||
|
||||
import api.services.pipecat.run_pipeline as run_pipeline_module
|
||||
from api.services.configuration.registry import ServiceProviders
|
||||
from api.services.pipecat.run_pipeline import (
|
||||
DEFAULT_PROVISIONAL_VAD_PAUSE_SECS,
|
||||
DEFAULT_TURN_START_MIN_WORDS,
|
||||
DEFAULT_USER_TURN_STOP_TIMEOUT,
|
||||
EXTERNAL_TURN_USER_STOP_TIMEOUT,
|
||||
_create_non_realtime_user_turn_start_strategies,
|
||||
_create_non_realtime_user_turn_stop_strategies,
|
||||
_create_realtime_user_turn_config,
|
||||
_resolve_user_turn_stop_timeout,
|
||||
)
|
||||
|
|
@ -181,6 +185,55 @@ def test_non_realtime_can_use_provisional_vad_start_strategy():
|
|||
|
||||
assert len(strategies) == 1
|
||||
assert isinstance(strategies[0], ProvisionalVADUserTurnStartStrategy)
|
||||
assert strategies[0]._pause_secs == DEFAULT_PROVISIONAL_VAD_PAUSE_SECS
|
||||
|
||||
|
||||
def test_non_realtime_provisional_vad_uses_configured_pause_secs():
|
||||
strategies = _create_non_realtime_user_turn_start_strategies(
|
||||
{"turn_start_strategy": "provisional_vad", "provisional_vad_pause_secs": 0.4},
|
||||
uses_external_turns=False,
|
||||
)
|
||||
|
||||
assert len(strategies) == 1
|
||||
assert isinstance(strategies[0], ProvisionalVADUserTurnStartStrategy)
|
||||
assert strategies[0]._pause_secs == 0.4
|
||||
|
||||
|
||||
def test_non_realtime_uses_external_stop_for_external_turn_stt():
|
||||
strategies = _create_non_realtime_user_turn_stop_strategies(
|
||||
{},
|
||||
uses_external_turns=True,
|
||||
)
|
||||
|
||||
assert len(strategies) == 1
|
||||
assert isinstance(strategies[0], ExternalUserTurnStopStrategy)
|
||||
|
||||
|
||||
def test_non_realtime_default_uses_speech_timeout_stop():
|
||||
strategies = _create_non_realtime_user_turn_stop_strategies(
|
||||
{},
|
||||
uses_external_turns=False,
|
||||
)
|
||||
|
||||
assert len(strategies) == 1
|
||||
assert isinstance(strategies[0], SpeechTimeoutUserTurnStopStrategy)
|
||||
|
||||
|
||||
def test_non_realtime_can_use_turn_analyzer_stop_strategy(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
run_pipeline_module,
|
||||
"LocalSmartTurnAnalyzerV3",
|
||||
lambda *, params: params,
|
||||
)
|
||||
|
||||
strategies = _create_non_realtime_user_turn_stop_strategies(
|
||||
{"turn_stop_strategy": "turn_analyzer", "smart_turn_stop_secs": 1.5},
|
||||
uses_external_turns=False,
|
||||
)
|
||||
|
||||
assert len(strategies) == 1
|
||||
assert isinstance(strategies[0], TurnAnalyzerUserTurnStopStrategy)
|
||||
assert strategies[0]._turn_analyzer.stop_secs == 1.5
|
||||
|
||||
|
||||
def test_external_turn_stt_uses_longer_stop_timeout():
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue